Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c5392aa237 | ||
|
|
16a2b55fa1 | ||
|
|
b423b09ff3 | ||
|
|
e66add88cb | ||
|
|
32b22c39cb | ||
|
|
259baac177 | ||
|
|
67ef8c13a8 | ||
|
|
b2bb1876de | ||
|
|
cda4e626e7 | ||
|
|
d835aff0cb | ||
|
|
21e1967c5e | ||
|
|
c9b5baf4d9 | ||
|
|
67c974c96e | ||
|
|
b46ccb03c9 | ||
|
|
0ecf98e08b | ||
|
|
f024034724 | ||
|
|
327a21f009 | ||
|
|
cfc51532b8 | ||
|
|
00988f92e4 | ||
|
|
868c6085a8 | ||
|
|
a47a56bda0 | ||
|
|
3667b42b2f | ||
|
|
7d142d7d62 | ||
|
|
24863e2ed3 | ||
|
|
fe8b526bbb | ||
|
|
6298be393a | ||
|
|
3a7853f9cc | ||
|
|
4a9413c83d | ||
|
|
21b04d62ae | ||
|
|
96929b6d7c | ||
|
|
07712d80a5 | ||
|
|
10c9eff16f | ||
|
|
edd7af986d | ||
|
|
1dc31927e3 | ||
|
|
36ef7d25ef | ||
|
|
b766b8b65d | ||
|
|
6579ff20b4 | ||
|
|
2fbee59c3e | ||
|
|
d3aaa19148 | ||
|
|
e32a3675fc | ||
|
|
b72e7dda08 | ||
|
|
0f77f28a95 | ||
|
|
289f83675b | ||
|
|
36633b4c72 | ||
|
|
4f45457811 | ||
|
|
c39890cd64 | ||
|
|
90f1e49263 | ||
|
|
9a1cf205db | ||
|
|
45eacb6a50 | ||
|
|
6cb2b57463 | ||
|
|
59f654fa39 | ||
|
|
be8ccc1dc4 | ||
|
|
228e5d9183 | ||
|
|
8afe6d0383 | ||
|
|
5f7190b08f | ||
|
|
a70a9b4bb1 | ||
|
|
b796e66890 | ||
|
|
f1a663779a | ||
|
|
b0aa972326 | ||
|
|
ef927a7ed1 | ||
|
|
aa8fc59051 | ||
|
|
ce62204392 | ||
|
|
837f28142d | ||
|
|
60c79c991d | ||
|
|
078aaeb679 | ||
|
|
d9edbd535e | ||
|
|
b4a61b21c3 | ||
|
|
bdc4193ffe | ||
|
|
74fdd6e396 | ||
|
|
b2479ebff2 | ||
|
|
ce2162c764 | ||
|
|
16ffd63c80 | ||
|
|
8faf68348d | ||
|
|
02dbc72856 | ||
|
|
da4dcf92dc | ||
|
|
49b750abcc | ||
|
|
4bb4122628 | ||
|
|
cee54f336e | ||
|
|
e95b3813cc | ||
|
|
6815cfb05e | ||
|
|
b6acbbce35 | ||
|
|
399e74877d | ||
|
|
61083e91a6 | ||
|
|
67b4ec3178 | ||
|
|
0fcb725a7a | ||
|
|
0dbdcdfdc7 | ||
|
|
e426d77353 | ||
|
|
bd15e29f17 | ||
|
|
b323d29567 | ||
|
|
0e54af3356 | ||
|
|
97f12f3bed | ||
|
|
5612047b97 | ||
|
|
2d147a3ae1 | ||
|
|
1a93c0f8e8 | ||
|
|
0a2b64881a | ||
|
|
42e7fe4d93 | ||
|
|
7ada28258c | ||
|
|
bd312afd00 | ||
|
|
d94a8af35b | ||
|
|
078fd10147 | ||
|
|
824e25d77c | ||
|
|
899b887e47 | ||
|
|
e58981d8a3 | ||
|
|
fc41d977a5 | ||
|
|
ab6210e667 | ||
|
|
f41805f053 | ||
|
|
baa809fcd6 | ||
|
|
a38d15e495 | ||
|
|
e97641372a | ||
|
|
9aecc2cb08 | ||
|
|
697667945e | ||
|
|
d908024577 | ||
|
|
a5a656d958 | ||
|
|
ddc3cf05dd | ||
|
|
66ad4b0abd | ||
|
|
a66023adc6 | ||
|
|
7277844128 | ||
|
|
6ef82b1d56 | ||
|
|
8ded4829f3 | ||
|
|
c4b6acb916 | ||
|
|
9beb81c303 | ||
|
|
f8dd4c6efa | ||
|
|
6ce5aa6a3a | ||
|
|
d2efa8a90a | ||
|
|
c6374063e9 | ||
|
|
0846013378 | ||
|
|
c141ba405f | ||
|
|
0320f13a9f | ||
|
|
8adc34be4d | ||
|
|
cb6810d3c1 | ||
|
|
d384f64abf | ||
|
|
ef7035f8ee | ||
|
|
bfcadde5c3 | ||
|
|
496ff41782 | ||
|
|
46f0be5484 | ||
|
|
f3db0131c1 | ||
|
|
83a8d47f51 | ||
|
|
fee0222910 | ||
|
|
1ed7b5511f | ||
|
|
b8f7c31537 | ||
|
|
164791c257 | ||
|
|
8f5e599928 | ||
|
|
7a7aaeb84d | ||
|
|
e2136ab2fc | ||
|
|
c75cb21946 | ||
|
|
bf95218c91 | ||
|
|
8cb4507a5f | ||
|
|
555890d1ba | ||
|
|
e4f54e83b6 | ||
|
|
692c4a709e | ||
|
|
cbd1961459 | ||
|
|
2e31a33ebf | ||
|
|
d16c6137d2 | ||
|
|
0416ab79ec | ||
|
|
fc9a1c62b9 | ||
|
|
5d4567b134 | ||
|
|
ae4a17d271 | ||
|
|
d110a08889 | ||
|
|
e0157293cb | ||
|
|
0d985b3b65 | ||
|
|
a65ade9fda | ||
|
|
874d6c8cb1 | ||
|
|
f70ba2afa3 | ||
|
|
e9f821e578 | ||
|
|
8e488d4b1d | ||
|
|
77201a457d | ||
|
|
076e3b1178 | ||
|
|
6b13fa64dc | ||
|
|
846671a890 | ||
|
|
05b3088b75 | ||
|
|
fe57286959 | ||
|
|
03645bbb33 | ||
|
|
93dba9a399 | ||
|
|
5627ea8073 | ||
|
|
7ba679c9ce | ||
|
|
c7a450e6ce | ||
|
|
beda5156bf | ||
|
|
76a9da7163 | ||
|
|
edd0303f59 | ||
|
|
be6f47a333 | ||
|
|
4cd6a072ca | ||
|
|
743a82efe9 | ||
|
|
9589f28ef7 | ||
|
|
35492c5671 | ||
|
|
db1e695bf3 | ||
|
|
ecc4aec43b | ||
|
|
fc063c2205 | ||
|
|
4d60ce138a | ||
|
|
2afd24f6e4 | ||
|
|
437acd023a | ||
|
|
b00523ae14 | ||
|
|
4405a74993 | ||
|
|
cb16090868 | ||
|
|
396e510dce | ||
|
|
3b9790b969 | ||
|
|
a35d07a7ac | ||
|
|
6d004c61fc | ||
|
|
ffdd06da1b | ||
|
|
f03f34cacb | ||
|
|
0c86ea849e | ||
|
|
0efa4c38c0 | ||
|
|
6092ab7793 | ||
|
|
929def87eb | ||
|
|
be074ccff7 | ||
|
|
3445199393 | ||
|
|
216c7e152e | ||
|
|
cc8bc10690 | ||
|
|
69b4218d60 | ||
|
|
1dd18dc4f8 | ||
|
|
4ccbd999d9 | ||
|
|
fa8d404964 | ||
|
|
30086957c9 | ||
|
|
0e57c620c9 | ||
|
|
3ce1c59a2d | ||
|
|
3337e20b9e | ||
|
|
e816b3626e | ||
|
|
3e0cb0f17a | ||
|
|
41bc606217 | ||
|
|
5a5f4ca49a | ||
|
|
c3a8437cd1 | ||
|
|
8d8a1a392d | ||
|
|
5f93fb5e55 | ||
|
|
d05050d7d8 | ||
|
|
8e9744100d | ||
|
|
1e4e7e287d | ||
|
|
e8f0c73f08 | ||
|
|
e923e28f8d | ||
|
|
5cc75bfa7c | ||
|
|
d6701769b8 | ||
|
|
0ddc67bdab | ||
|
|
38b62b7a68 | ||
|
|
7e726000c7 | ||
|
|
d8dfb292ec | ||
|
|
826975241d | ||
|
|
743637ceaf | ||
|
|
e350c7e31e | ||
|
|
66b1e0ab9f | ||
|
|
7b0374d110 | ||
|
|
e86ef8cbb0 | ||
|
|
8c901c54bc | ||
|
|
408d85691e | ||
|
|
c66cd6901b | ||
|
|
aeadbc4f6d | ||
|
|
224136890e | ||
|
|
3669a1e86d | ||
|
|
d588b5b327 | ||
|
|
b705679098 | ||
|
|
f71a0b0da5 | ||
|
|
ebc2c76b6b | ||
|
|
2e3fff278e | ||
|
|
1f4bc5e089 | ||
|
|
52c38b10dd | ||
|
|
7047aa5456 | ||
|
|
33fe4019f7 | ||
|
|
80b9d97690 | ||
|
|
3c3c92723f | ||
|
|
037bd87006 | ||
|
|
f688310d28 | ||
|
|
c4d65e7a45 | ||
|
|
6f208b710d | ||
|
|
b599faaf85 | ||
|
|
6d991d20dc | ||
|
|
c87e0296f6 | ||
|
|
16cdb4c5b4 | ||
|
|
7631b8924d | ||
|
|
785d307ff3 | ||
|
|
8c713ff35e | ||
|
|
7d80493bef | ||
|
|
bff2760c3d | ||
|
|
a0f8848367 | ||
|
|
5b1cbcd8d5 | ||
|
|
05857a92d5 | ||
|
|
6bdc811286 | ||
|
|
469d50a5b8 | ||
|
|
ef86904bfb | ||
|
|
d4181ea67c | ||
|
|
1c6d17309f | ||
|
|
db293ec41d | ||
|
|
db8d468f29 | ||
|
|
d7d7af7265 | ||
|
|
bd763cadc1 | ||
|
|
22799fc549 | ||
|
|
0f231d1271 | ||
|
|
8c0c911020 | ||
|
|
6cb9df700b | ||
|
|
4fdda537b9 | ||
|
|
cb6f32465a | ||
|
|
c235e36cb4 | ||
|
|
aaca440a94 | ||
|
|
1f57950a29 | ||
|
|
3d8855ec72 | ||
|
|
ac9231d9f3 | ||
|
|
37c0a56d89 | ||
|
|
010915dac4 | ||
|
|
83a8b3b970 | ||
|
|
f130202aa1 | ||
|
|
c8f6800bcd | ||
|
|
078f9f5dd4 | ||
|
|
c0de178c7d | ||
|
|
1c767b538d | ||
|
|
51aab44b5d | ||
|
|
b8a0d4a67b | ||
|
|
cd0dcfbb8c | ||
|
|
fcc9e30eae | ||
|
|
5f66218a43 | ||
|
|
61ef4f9a0f | ||
|
|
f0a8734b42 | ||
|
|
4fe95ef4ec | ||
|
|
2a5148845b | ||
|
|
8fa562caaf | ||
|
|
cab5620cd5 | ||
|
|
be38d36677 | ||
|
|
69236fca89 | ||
|
|
4a4f376bfd | ||
|
|
fd9718fe24 | ||
|
|
26a6e11212 | ||
|
|
de1a669f6e | ||
|
|
e482c9e5c4 | ||
|
|
4e96a77a41 | ||
|
|
3346290e5c | ||
|
|
dcac593efe | ||
|
|
7248d0de02 | ||
|
|
8ad3ce632c | ||
|
|
9398b02562 | ||
|
|
164e4da99d | ||
|
|
0025ea6119 | ||
|
|
eef53a5165 | ||
|
|
9ea066d948 | ||
|
|
4bd900c4a1 | ||
|
|
736cd2bebd | ||
|
|
cd658c2a60 | ||
|
|
9730658f21 | ||
|
|
3fa107acb1 | ||
|
|
7da22179a0 | ||
|
|
fc7b71ee78 | ||
|
|
b79573bf1f | ||
|
|
a01db6f7c0 | ||
|
|
117d58c58e | ||
|
|
ff6626ed89 | ||
|
|
10c798a440 | ||
|
|
9097d87819 | ||
|
|
d901f503d1 | ||
|
|
65f2b6ce6f | ||
|
|
d949fe8bf1 | ||
|
|
15cfb48550 | ||
|
|
447dc6d4c3 | ||
|
|
c92e43b920 | ||
|
|
be6c32b0e0 | ||
|
|
3d2062e810 | ||
|
|
dd816e95cd | ||
|
|
6d1b51890d | ||
|
|
11f03ec99a | ||
|
|
bd192f43e7 | ||
|
|
36e4b11983 | ||
|
|
97397ba8c2 | ||
|
|
052eee4111 | ||
|
|
e319496044 | ||
|
|
b83b63c362 | ||
|
|
4d6b1675bb | ||
|
|
13110fab39 | ||
|
|
74ea509848 | ||
|
|
192bff9d2c | ||
|
|
44ed8812dc | ||
|
|
200696ba21 | ||
|
|
51cf3b0c04 | ||
|
|
6a4831c83b | ||
|
|
c5e7ed95a3 | ||
|
|
42a97fa4d9 | ||
|
|
45240d0012 | ||
|
|
8ed085febd | ||
|
|
37803ea61b | ||
|
|
acd416952c | ||
|
|
6ec46cbc44 | ||
|
|
a9d971e476 | ||
|
|
9fe064675d | ||
|
|
c84fa467d0 | ||
|
|
3d7a55f6d3 | ||
|
|
d49baa1540 | ||
|
|
5c686af842 | ||
|
|
22425b5bc6 | ||
|
|
b6d9b338d2 | ||
|
|
a191a13751 | ||
|
|
3eccdbcc9b | ||
|
|
50063903f9 | ||
|
|
95a1b70533 | ||
|
|
c3679ac90b | ||
|
|
29e48eb6a2 | ||
|
|
71d02e9651 | ||
|
|
f5193f3eec | ||
|
|
480c4d6919 | ||
|
|
1c5e030540 | ||
|
|
27e83a5908 | ||
|
|
d3cbf8fa8d | ||
|
|
74b1f8129b | ||
|
|
56ed513cfd | ||
|
|
5f412371c4 | ||
|
|
c6b0b67585 | ||
|
|
16d18e681a | ||
|
|
1fe99f33b2 | ||
|
|
a2ece25ac0 | ||
|
|
4865f4d148 | ||
|
|
41e88824cf | ||
|
|
0961ab138e | ||
|
|
fe8271a12f | ||
|
|
f3866ede89 | ||
|
|
d938adf3cc | ||
|
|
9908cff64b | ||
|
|
1b9b0bb4e6 | ||
|
|
03acd9bea5 | ||
|
|
4c42949023 | ||
|
|
2e4d9836e5 | ||
|
|
43c6b58354 | ||
|
|
df637e8196 | ||
|
|
8e78f9786c | ||
|
|
f4130f06ed | ||
|
|
b766e714a4 | ||
|
|
1100a90be3 | ||
|
|
33aaf80c82 | ||
|
|
b5861dbc24 | ||
|
|
93731416fc | ||
|
|
a305e736ca | ||
|
|
a046ebbafb | ||
|
|
af65e96723 | ||
|
|
c9eb0ab5f0 | ||
|
|
9802e841a8 | ||
|
|
b896df8d54 | ||
|
|
720b8c237b | ||
|
|
371f9f813f | ||
|
|
33e229c41c | ||
|
|
2c33c0d801 | ||
|
|
d64fee5954 | ||
|
|
0217678c8c | ||
|
|
2959a9c31f | ||
|
|
1928a18992 | ||
|
|
a168171009 | ||
|
|
6f767f9700 | ||
|
|
e5459f63fd | ||
|
|
d0ab85a8c4 | ||
|
|
b7b86fe8c4 | ||
|
|
c3bcf6907a | ||
|
|
330fed867b | ||
|
|
45bbc31dc1 | ||
|
|
9dc45239c4 | ||
|
|
b6cfb30908 | ||
|
|
8efd94cc76 | ||
|
|
5d864e7ea2 | ||
|
|
7a1c91a2d5 | ||
|
|
1339c8a1f4 | ||
|
|
27025579d6 | ||
|
|
473d70f818 | ||
|
|
907a5d8d7e | ||
|
|
28c349fd7e | ||
|
|
df9545015f | ||
|
|
07b0b94e46 | ||
|
|
92f85b6d9c | ||
|
|
6919eadb21 | ||
|
|
be76280fd2 | ||
|
|
6b673bdd44 | ||
|
|
c6591db45e | ||
|
|
5813ebaa1c | ||
|
|
743e752b0f | ||
|
|
707cd28cb7 | ||
|
|
ecdb687c25 | ||
|
|
ba68d2dd45 | ||
|
|
d8259de52f | ||
|
|
07cec3566b | ||
|
|
dd8c531889 | ||
|
|
686ebcfd8b | ||
|
|
51aaba39cf | ||
|
|
d7c6632499 | ||
|
|
1b0ea06876 | ||
|
|
8ebe88629b | ||
|
|
e226992703 | ||
|
|
11a8394d69 | ||
|
|
9f54a1b91a | ||
|
|
35d11061e9 | ||
|
|
65d8b490ca | ||
|
|
0fef12c3b1 | ||
|
|
8516bff224 | ||
|
|
1bcc501352 | ||
|
|
f492b17fbe | ||
|
|
48ae90f80e | ||
|
|
9321ccbc48 | ||
|
|
402cd01e1a | ||
|
|
7b2d0e29c6 | ||
|
|
52d38c401a | ||
|
|
a34dd61076 | ||
|
|
a53a3e772a | ||
|
|
8f24c294a7 | ||
|
|
acc3f76654 | ||
|
|
8c977fb442 | ||
|
|
aa20a2de67 | ||
|
|
5fcb154d89 | ||
|
|
0980129f4e | ||
|
|
1258746886 | ||
|
|
ed128b0ad6 | ||
|
|
f0db08acd6 | ||
|
|
7c655e3080 | ||
|
|
1b9871c3df | ||
|
|
7568aaf243 | ||
|
|
a76be8450d | ||
|
|
5564ee1246 | ||
|
|
a6e9251521 | ||
|
|
0bee093916 | ||
|
|
13a9878823 | ||
|
|
acd35d50f8 | ||
|
|
5c0d99e72d | ||
|
|
e37af93be3 | ||
|
|
244c1700e1 | ||
|
|
a3a15473ba | ||
|
|
d734b5077c | ||
|
|
5430072b19 | ||
|
|
465aebaed4 | ||
|
|
ee0b16c2ea | ||
|
|
21d5eacb41 | ||
|
|
c06688eb0b | ||
|
|
b74bbcd279 | ||
|
|
64d8d9b05d | ||
|
|
6ab60f281b | ||
|
|
e35be3b2fa | ||
|
|
0a0c27ac96 | ||
|
|
fb249e84eb | ||
|
|
76ad86fcae | ||
|
|
037614d227 | ||
|
|
4a50e445fd | ||
|
|
6f3c1c4393 | ||
|
|
c6a9b4b592 | ||
|
|
e915ac4eca | ||
|
|
a857793f63 | ||
|
|
31914f7510 | ||
|
|
0be859f0ee | ||
|
|
14b9c3697b | ||
|
|
96b66a57bb | ||
|
|
f13701c489 | ||
|
|
31515b810e | ||
|
|
29e84e08a4 | ||
|
|
0ac9ad9757 | ||
|
|
9a432e0608 | ||
|
|
c83ba5fe7f | ||
|
|
1b55c743ea | ||
|
|
a93579376c | ||
|
|
eba49f3c68 | ||
|
|
3a3da49c69 | ||
|
|
3d68e48219 | ||
|
|
4351fa6a0e | ||
|
|
77222d2808 | ||
|
|
36db7e5a9a | ||
|
|
3572368f16 | ||
|
|
0755dc1462 | ||
|
|
94f81b7102 | ||
|
|
08f8fe3d7e | ||
|
|
a6cd383d67 | ||
|
|
10face6ab0 | ||
|
|
a6259ff600 | ||
|
|
d7d9e6cbfe | ||
|
|
8263609470 | ||
|
|
fc2367de76 | ||
|
|
9a4f2ebc70 | ||
|
|
812879610a | ||
|
|
c5e521ccc1 | ||
|
|
bc1c8fa351 | ||
|
|
7271fcf9c1 | ||
|
|
ace3b7707b | ||
|
|
9eb65cc4ee | ||
|
|
c5c2bc779c | ||
|
|
ae1751d9c0 | ||
|
|
d17583ef7d | ||
|
|
a363713ae0 | ||
|
|
1566165bd4 | ||
|
|
fe065fa318 | ||
|
|
202d5cf071 | ||
|
|
b785a9dc5b | ||
|
|
73bc658b2f | ||
|
|
0b94216138 | ||
|
|
86eec2b4cc | ||
|
|
c99b531d28 | ||
|
|
74a4338cb5 | ||
|
|
d691c52e49 | ||
|
|
bd3e9e4b3c | ||
|
|
3c3ca5fb9c | ||
|
|
e4ff4fce1c | ||
|
|
2918d4b07d | ||
|
|
c40e49be46 | ||
|
|
4ccda20975 | ||
|
|
09957617d3 | ||
|
|
c703aa7058 | ||
|
|
64d366d323 | ||
|
|
333e0a2faa | ||
|
|
a677d95bc8 | ||
|
|
aafd87e84b | ||
|
|
efa3bae54b | ||
|
|
8e4362689d | ||
|
|
b86634284e | ||
|
|
b3293ddccd | ||
|
|
4b40831b83 | ||
|
|
e426b0521c | ||
|
|
1777bf6e06 | ||
|
|
dafc892f0f | ||
|
|
fea0cfd5dd | ||
|
|
a7e158db6d | ||
|
|
455ac4abd3 | ||
|
|
778dfa2cf5 | ||
|
|
329f2e6f81 | ||
|
|
63ad6d97d7 | ||
|
|
8db56db7cf | ||
|
|
bd542f1e0b | ||
|
|
a162e53dea | ||
|
|
a8a4c848ed | ||
|
|
bddd38996a | ||
|
|
9f084eae94 | ||
|
|
c54c635161 | ||
|
|
fa8b42e05e | ||
|
|
9c3c323884 | ||
|
|
c8a46439be | ||
|
|
b7ec701259 | ||
|
|
a2cd0e0a38 | ||
|
|
7bccc0e236 | ||
|
|
a3a649a79f | ||
|
|
2a14d30552 | ||
|
|
313bef0609 | ||
|
|
960a80aeca | ||
|
|
20e6d50a98 | ||
|
|
507c3417d6 | ||
|
|
d1adb8d4ed | ||
|
|
915ff12747 | ||
|
|
fbade79137 | ||
|
|
a240a677d0 | ||
|
|
ab64cf31f6 | ||
|
|
67f2e32dae | ||
|
|
0dc40fe052 | ||
|
|
b7225be552 | ||
|
|
629e00ec94 | ||
|
|
6d033c9314 | ||
|
|
4f8926ed00 | ||
|
|
1daa1a4603 | ||
|
|
062773d929 | ||
|
|
57decadaef | ||
|
|
ec804ab7c9 | ||
|
|
3ee7533098 | ||
|
|
0d383ccc1f | ||
|
|
e2f2257c34 | ||
|
|
d62b9fc4c6 | ||
|
|
be7ad0c7fb | ||
|
|
156864cc8b | ||
|
|
7240e496cc | ||
|
|
8acdf4018d | ||
|
|
1ea7c3e203 | ||
|
|
e54aeb6125 | ||
|
|
d59f51fbcf | ||
|
|
a667eb6982 | ||
|
|
8d72732247 | ||
|
|
f0f3b30a62 | ||
|
|
d99fe24542 | ||
|
|
ea4c7381bd | ||
|
|
e900d20641 | ||
|
|
8a46647d8c | ||
|
|
968178bf57 | ||
|
|
406a255db0 | ||
|
|
574557810e | ||
|
|
998a02c3a4 | ||
|
|
c6f964c921 | ||
|
|
efb0e147c5 | ||
|
|
9cf7356f98 | ||
|
|
af05c43174 | ||
|
|
380c68ff2b | ||
|
|
068b00b99f | ||
|
|
cd6a42ab64 | ||
|
|
f115abec92 | ||
|
|
f0e23cf878 | ||
|
|
a94f11d809 | ||
|
|
38972bea5f | ||
|
|
a761ff552a | ||
|
|
dbeb84ea9a | ||
|
|
0a4938f39a | ||
|
|
b2182c716d | ||
|
|
8253be73f6 | ||
|
|
20318e296e | ||
|
|
cea1b69286 | ||
|
|
ab8aa69389 | ||
|
|
681491f1d0 | ||
|
|
c2fb815074 | ||
|
|
fffa14dc44 | ||
|
|
df37166d42 | ||
|
|
25fa3a8f6a | ||
|
|
b11507c5e6 | ||
|
|
f06d02489f | ||
|
|
c203af2f71 | ||
|
|
c715155a70 | ||
|
|
8163133294 | ||
|
|
765be5dab4 | ||
|
|
7c1523389d | ||
|
|
7a2b1ba166 | ||
|
|
45b4dcfcd0 | ||
|
|
3b2e535566 | ||
|
|
db556d13a3 | ||
|
|
a987063c68 | ||
|
|
4ce30ef899 | ||
|
|
d988282d98 | ||
|
|
695fdf7ceb | ||
|
|
0befe164cc | ||
|
|
d506c68a80 | ||
|
|
c59c429b75 | ||
|
|
c726b6e4a2 | ||
|
|
96075ad4e1 | ||
|
|
7a8dc07a8a | ||
|
|
1fdac0bc09 | ||
|
|
804b942a36 | ||
|
|
6335d4378b | ||
|
|
cfc2189616 | ||
|
|
e06e032701 | ||
|
|
c9499c2c79 | ||
|
|
2bf43541bb | ||
|
|
0386c7266d | ||
|
|
7aa6ed9a5d | ||
|
|
b71325afa5 | ||
|
|
2cc29bdf77 | ||
|
|
b5d602abc4 | ||
|
|
33c45637ac | ||
|
|
662d4478d0 | ||
|
|
31d3809572 | ||
|
|
024ff4a309 | ||
|
|
9ae8d30b6b | ||
|
|
44349c10b0 | ||
|
|
506520a3c4 | ||
|
|
50f7020977 | ||
|
|
02a27a03cc | ||
|
|
4ad6bacf7b | ||
|
|
26ecc0fa44 | ||
|
|
2619befca6 | ||
|
|
ea24f52b13 | ||
|
|
4abfc47346 | ||
|
|
3b95010d06 | ||
|
|
b3c1b96088 | ||
|
|
a9f1326873 | ||
|
|
75a696fb64 | ||
|
|
24aaacba6d | ||
|
|
e3cd7d5f91 | ||
|
|
2e228e8db5 | ||
|
|
5c9dd80370 | ||
|
|
727f5f2e48 | ||
|
|
1c238f7697 | ||
|
|
119d7cce15 | ||
|
|
4afc8f6083 | ||
|
|
a9ec3af066 | ||
|
|
9edae81fee | ||
|
|
b941b12f12 | ||
|
|
3069de188a | ||
|
|
a968f08abd | ||
|
|
4153d3e5ff | ||
|
|
9d9c1a6c84 | ||
|
|
16ef10a4d9 | ||
|
|
b3766e440a | ||
|
|
6359c3f70f | ||
|
|
efe73fb965 | ||
|
|
c45a962fcc | ||
|
|
f98a03e2e9 | ||
|
|
5b6257814d | ||
|
|
69a445d4ed | ||
|
|
e82c786b8a | ||
|
|
eec2225c89 | ||
|
|
f7355e0b71 | ||
|
|
6c6a99cfe4 | ||
|
|
b4634e2e0d | ||
|
|
3f4cba0612 | ||
|
|
38db99cc75 | ||
|
|
4d5906394b | ||
|
|
2fc212b156 | ||
|
|
53fbb5b027 | ||
|
|
4f24721450 | ||
|
|
83043727b5 | ||
|
|
2d336afb85 | ||
|
|
4d309435c8 | ||
|
|
099ce9cdfd | ||
|
|
8914e60cb8 | ||
|
|
dbd30a40e9 | ||
|
|
c9a598fd59 | ||
|
|
e331e588cf | ||
|
|
2011557771 | ||
|
|
f0ba45d14e | ||
|
|
8352a521b7 | ||
|
|
aa3d4d79f8 | ||
|
|
4f650d760c | ||
|
|
ea4b792627 | ||
|
|
883605239a | ||
|
|
5b8cab920c | ||
|
|
8d3d327335 | ||
|
|
32574050c4 | ||
|
|
8d45a90d9b | ||
|
|
f66862a422 | ||
|
|
6a56be3a9b | ||
|
|
ebf6395de2 | ||
|
|
5df9fbf50d | ||
|
|
ff961155c9 | ||
|
|
14838d06a8 | ||
|
|
99def24dd8 | ||
|
|
f5b210d142 | ||
|
|
a137a23b48 | ||
|
|
6bbf06d9e9 | ||
|
|
0b614b40cf | ||
|
|
27673561bd | ||
|
|
6e2070410d | ||
|
|
7ccf21f74f | ||
|
|
3b2710f285 | ||
|
|
c4b277235b | ||
|
|
a55318add1 | ||
|
|
b57123a4fe | ||
|
|
04dcc00670 | ||
|
|
746a02b49f | ||
|
|
bdbe3db2a9 | ||
|
|
27ae99ad86 | ||
|
|
38add89547 | ||
|
|
a9612fbb2f | ||
|
|
429cc29b5b | ||
|
|
8eca94e405 | ||
|
|
ad71daafb6 | ||
|
|
c936d83688 | ||
|
|
c6684d680f | ||
|
|
fe358b0e13 | ||
|
|
897f259a2a | ||
|
|
f3302c1b3a | ||
|
|
573feeaaab | ||
|
|
019c98ecc1 | ||
|
|
ad6a51a4b5 | ||
|
|
f94278776e | ||
|
|
315885cb0b | ||
|
|
7e605f8228 | ||
|
|
aed70435f9 | ||
|
|
1fb1728ede | ||
|
|
497c4fe5a3 | ||
|
|
09ad7764b8 | ||
|
|
bc3d24fddf | ||
|
|
2d634d628a | ||
|
|
790c22d919 | ||
|
|
c275a56806 | ||
|
|
82c3c7addd | ||
|
|
b464d85c04 | ||
|
|
078618b4cf | ||
|
|
561805a417 | ||
|
|
7bb4324365 | ||
|
|
0608653d35 | ||
|
|
fa3472cdc5 | ||
|
|
7d553b6fcf | ||
|
|
677627630e | ||
|
|
9f23172b22 | ||
|
|
14ceaa472d | ||
|
|
1188b9d3bc | ||
|
|
424c9a9423 | ||
|
|
8f60c81ef3 | ||
|
|
dd648d1ae5 | ||
|
|
02e839e272 | ||
|
|
ec8c56707b | ||
|
|
1e8d317ee8 | ||
|
|
e81df111a7 | ||
|
|
90e55ffe14 | ||
|
|
4e73b1d3fc | ||
|
|
21dca1e34f | ||
|
|
c2292850bb | ||
|
|
04791c92e2 | ||
|
|
d9462b6d8a | ||
|
|
be9b83559e | ||
|
|
958889afee | ||
|
|
dce677035f | ||
|
|
d98a8855ad | ||
|
|
7cd0587a64 | ||
|
|
bb3967f11e | ||
|
|
36ee203bd7 | ||
|
|
914919ba75 | ||
|
|
a4514e8565 | ||
|
|
116e983c50 | ||
|
|
fb667b6c42 | ||
|
|
b0a090cf14 | ||
|
|
110b470d33 | ||
|
|
c517c4d015 | ||
|
|
35ed4f9101 | ||
|
|
18c723c2e3 | ||
|
|
631223602c | ||
|
|
a6cc907de2 | ||
|
|
f847dcccf4 | ||
|
|
ff3f8f52d0 | ||
|
|
d7d46682fc |
@@ -0,0 +1,25 @@
|
||||
name: Publish to Comfy registry
|
||||
on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
if: ${{ github.repository_owner == 'shadowcz007' }}
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
@@ -1,4 +1,8 @@
|
||||
__pycache__/
|
||||
https/
|
||||
nodes/config.json
|
||||
workflow/my_workflow.json
|
||||
workflow/my_workflow.json
|
||||
workflow/my_workflow_app.json
|
||||
workflow/prompt_result.json
|
||||
app/*
|
||||
workflow/prompt_result.json
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2024 shadow
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
@@ -1,36 +1,351 @@
|
||||
##
|
||||

|
||||
|
||||
> 适配了最新版 comfyui 的 py3.11 ,torch 2.3.1+cu121
|
||||
> [Mixlab nodes discord](https://discord.gg/cXs9vZSqeK)
|
||||
|
||||
商务合作请联系 389570357@qq.com
|
||||
For business cooperation, please contact email 389570357@qq.com
|
||||
|
||||

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

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

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

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

|
||||
|
||||

|
||||
|
||||

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

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

|
||||
|
||||
|
||||
### ScreenShareNode & FloatingVideoNode
|
||||
> Now comfyui supports capturing screen pixel streams from any software and can be used for LCM-Lora integration. Let's get started with implementation and design! 💻🌐
|
||||
|
||||
https://github.com/shadowcz007/comfyui-mixlab-nodes/assets/12645064/e7e77f90-e43e-410a-ab3a-1952b7b4e7da
|
||||
|
||||
|
||||
<!-- [ScreenShareNode](./workflow/2-screeshare.json) -->
|
||||
|
||||
[ScreenShareNode & FloatingVideoNode](./workflow/3-FloatVideo-workflow.json)
|
||||
|
||||
!! Please use the address with HTTPS (https://127.0.0.1).
|
||||
|
||||
### SpeechRecognition & SpeechSynthesis
|
||||
|
||||
### LoadImagesFromLocal
|
||||
> Monitor changes to images in a local folder, and trigger real-time execution of workflows, supporting common image formats, especially PSD format, in conjunction with Photoshop. Q: Translate into English
|
||||

|
||||
|
||||
[Voice + Real-time Face Swap Workflow](./workflow/语音+实时换脸workflow.json)
|
||||
|
||||
- Preview Audio
|
||||
|
||||
[text-to-audio](./workflow/text-to-audio-base-workflow.json)
|
||||
|
||||
### GPT
|
||||
|
||||
> Support for calling multiple GPTs.Local LLM 、 ChatGPT、ChatGLM3 、ChatGLM4 , Some code provided by rui. If you are using OpenAI's service, fill in https://api.openai.com/v1 . If you are using a local LLM service, fill in http://127.0.0.1:xxxx/v1 . Azure OpenAI:https://xxxx.openai.azure.com
|
||||
|
||||
[LLM_base_workflow](./workflow/LLM_base_workflow.json)
|
||||
|
||||
- SiliconflowLLM
|
||||
- ChatGPTOpenAI
|
||||
|
||||
<!-- 最新:ChatGPT 节点支持 Local LLM(llama.cpp),Phi3、llama3 都可以直接一个节点运行了。
|
||||
|
||||
Model download,move to :`models/llamafile/`
|
||||
|
||||
强烈推荐:[Phi-3-mini-4k-instruct-GGUF](https://huggingface.co/lmstudio-community/Phi-3-mini-4k-instruct-GGUF/tree/main)
|
||||
|
||||
备选:[llama3_if_ai_sdpromptmkr_q2k](https://hf-mirror.com/impactframes/llama3_if_ai_sdpromptmkr_q2k/tree/main)
|
||||
|
||||
> 如果碰到安装失败,可以尝试手动安装
|
||||
|
||||
```
|
||||
../../../python_embeded/python.exe -s -m pip install llama-cpp-python --extra-index-url https://abetlen.github.io/llama-cpp-python/whl/cu121
|
||||
|
||||
../../../python_embeded/python.exe -s -m pip install llama-cpp-python[server]
|
||||
|
||||
```
|
||||
|
||||
> [Mac](https://llama-cpp-python.readthedocs.io/en/latest/install/macos/)
|
||||
|
||||
```
|
||||
pip uninstall llama-cpp-python -y
|
||||
CMAKE_ARGS="-DLLAMA_METAL=on" pip install -U llama-cpp-python --no-cache-dir
|
||||
pip install 'llama-cpp-python[server]'
|
||||
```
|
||||
|
||||
```
|
||||
pip install llama-cpp-python \
|
||||
--extra-index-url https://abetlen.github.io/llama-cpp-python/whl/metal
|
||||
``` -->
|
||||
|
||||
## Prompt
|
||||
|
||||
> PromptSlide
|
||||
> 
|
||||
|
||||
<!--  -->
|
||||
|
||||
> randomPrompt
|
||||
|
||||

|
||||
|
||||
> ClipInterrogator
|
||||
|
||||
[add clip-interrogator](https://github.com/pharmapsychotic/clip-interrogator)
|
||||
|
||||
> PromptImage & PromptSimplification,Assist in simplifying prompt words, comparing images and prompt word nodes.
|
||||
|
||||
> ChinesePrompt && PromptGenerate,中文 prompt 节点,直接用中文书写你的 prompt
|
||||
|
||||

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

|
||||
|
||||

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

|
||||

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

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

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

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

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

|
||||

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

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

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

|
||||
|
||||

|
||||
|
||||
|
||||

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

|
||||
|
||||
|
||||
### Models
|
||||
|
||||
- [Download TripoSR](https://huggingface.co/stabilityai/TripoSR/blob/main/model.ckpt) and place it in `models/triposr`
|
||||
|
||||
- [Download facebook/dino-vitb16](https://huggingface.co/facebook/dino-vitb16/tree/main) and place it in `models/triposr/facebook/dino-vitb16`
|
||||
|
||||
[Download rembg Models](https://github.com/danielgatis/rembg/tree/main#Models),move to:`models/rembg`
|
||||
|
||||
[Download lama](https://github.com/enesmsahin/simple-lama-inpainting/releases/download/v0.1.0/big-lama.pt), move to : `models/lama`
|
||||
|
||||
[Download Salesforce/blip-image-captioning-base](https://huggingface.co/Salesforce/blip-image-captioning-base), move to :`models/clip_interrogator/Salesforce/blip-image-captioning-base`
|
||||
|
||||
[Download succinctly/text2image-prompt-generator](https://huggingface.co/succinctly/text2image-prompt-generator/tree/main),move to:`models/prompt_generator/text2image-prompt-generator`
|
||||
|
||||
[Download Helsinki-NLP/opus-mt-zh-en](https://huggingface.co/Helsinki-NLP/opus-mt-zh-en/tree/main),move to:`models/prompt_generator/opus-mt-zh-en`
|
||||
|
||||
## Installation
|
||||
|
||||
@@ -46,77 +361,51 @@ git clone https://github.com/shadowcz007/comfyui-mixlab-nodes.git
|
||||
Install the requirements:
|
||||
|
||||
run directly:
|
||||
|
||||
```
|
||||
cd ComfyUI_Mixlab
|
||||
cd ComfyUI/custom_nodes/comfyui-mixlab-nodes
|
||||
install.bat
|
||||
```
|
||||
|
||||
or install the requirements using:
|
||||
|
||||
```
|
||||
../../../python_embeded/python.exe -s -m pip install -r requirements.txt
|
||||
```
|
||||
|
||||
If you are using a venv, make sure you have it activated before installation and use:
|
||||
|
||||
```
|
||||
pip3 install -r requirements.txt
|
||||
```
|
||||
|
||||
#### Chinese community
|
||||
|
||||
访问 [www.mixcomfy.com](https://www.mixcomfy.com),获得更多内测功能,关注微信公众号:Mixlab 无界社区
|
||||
|
||||
## Nodes
|
||||
####
|
||||
|
||||

|
||||

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

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

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

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

|
||||
|
||||
|
||||
|
||||
### Improvement
|
||||
An improvement has been made to directly redirect to GitHub to search for missing nodes when loading the graph.
|
||||
|
||||

|
||||
|
||||
|
||||
### Models
|
||||
[Download CLIPSeg](https://huggingface.co/CIDAS/clipseg-rd64-refined/tree/main), move to : model/clipseg
|
||||
|
||||
<!-- ### Workflow
|
||||
[Workflow](./workflow.md) -->
|
||||
|
||||
#### Thanks:
|
||||
[ComfyUI-CLIPSeg](https://github.com/biegert/ComfyUI-CLIPSeg/tree/main)
|
||||
File / LoadImagesFromPath SaveImageToLocal LoadImagesFromURL
|
||||
|
||||
#### discussions:
|
||||
|
||||
[discussions](https://github.com/shadowcz007/comfyui-mixlab-nodes/discussions)
|
||||
|
||||
### TODO:
|
||||
- vector https://github.com/GeorgLegato/stable-diffusion-webui-vectorstudio
|
||||
|
||||
<picture>
|
||||
<source
|
||||
media="(prefers-color-scheme: dark)"
|
||||
srcset="
|
||||
https://api.star-history.com/svg?repos=shadowcz007/comfyui-mixlab-nodes&type=Date&theme=dark
|
||||
"
|
||||
/>
|
||||
<source
|
||||
media="(prefers-color-scheme: light)"
|
||||
srcset="
|
||||
https://api.star-history.com/svg?repos=shadowcz007/comfyui-mixlab-nodes&type=Date
|
||||
"
|
||||
/>
|
||||
<img
|
||||
alt="Star History Chart"
|
||||
src="https://api.star-history.com/svg?repos=shadowcz007/comfyui-mixlab-nodes&type=Date"
|
||||
/>
|
||||
</picture>
|
||||
|
||||
|
After Width: | Height: | Size: 101 KiB |
|
After Width: | Height: | Size: 537 KiB |
|
After Width: | Height: | Size: 340 KiB |
|
After Width: | Height: | Size: 29 KiB |
|
After Width: | Height: | Size: 366 KiB |
|
After Width: | Height: | Size: 135 KiB |
|
After Width: | Height: | Size: 210 KiB |
|
After Width: | Height: | Size: 450 KiB |
|
After Width: | Height: | Size: 2.2 MiB |
|
After Width: | Height: | Size: 2.4 MiB |
|
After Width: | Height: | Size: 1.1 MiB |
|
After Width: | Height: | Size: 11 KiB |
|
After Width: | Height: | Size: 240 KiB |
|
After Width: | Height: | Size: 254 KiB |
|
After Width: | Height: | Size: 73 KiB |
|
Before Width: | Height: | Size: 784 KiB |
|
After Width: | Height: | Size: 255 KiB |
|
After Width: | Height: | Size: 7.1 MiB |
|
Before Width: | Height: | Size: 35 KiB After Width: | Height: | Size: 94 KiB |
|
After Width: | Height: | Size: 51 KiB |
|
After Width: | Height: | Size: 9.9 MiB |
|
After Width: | Height: | Size: 75 KiB |
|
After Width: | Height: | Size: 63 KiB |
|
After Width: | Height: | Size: 477 KiB |
|
After Width: | Height: | Size: 965 KiB |
@@ -0,0 +1,30 @@
|
||||
Jony Ive
|
||||
Dieter Rams
|
||||
Philippe Starck
|
||||
Karim Rashid
|
||||
Yves Béhar
|
||||
Marc Newson
|
||||
Naoto Fukasawa
|
||||
Jonathan Adler
|
||||
Patricia Urquiola
|
||||
Ross Lovegrove
|
||||
Tom Dixon
|
||||
Jasper Morrison
|
||||
Charles Eames
|
||||
Ray Eames
|
||||
Achille Castiglioni
|
||||
Ron Arad
|
||||
Konstantin Grcic
|
||||
Marcel Wanders
|
||||
Maarten Baas
|
||||
Stefan Sagmeister
|
||||
Ingo Maurer
|
||||
Hella Jongerius
|
||||
Sam Hecht
|
||||
Kim Colin
|
||||
Jaime Hayon
|
||||
Michael Anastassiades
|
||||
Nendo
|
||||
Oki Sato
|
||||
Matali Crasset
|
||||
Tokujin Yoshioka
|
||||
@@ -0,0 +1,10 @@
|
||||
Chibi Anime Style
|
||||
Gakuen Anime Style
|
||||
Gekiga Anime Style
|
||||
Jidaimono Anime Style
|
||||
Kawaii Anime Style
|
||||
Mecha Anime Style
|
||||
Realistic Anime Style
|
||||
Semi-Realistic Anime Style
|
||||
Shoji Anime Style
|
||||
Kemonomimi Anime Style
|
||||
@@ -0,0 +1,23 @@
|
||||
GoPro
|
||||
Drone
|
||||
polaroid
|
||||
black and white film
|
||||
Kodachrome
|
||||
shot on 8mm
|
||||
shot on 16mm
|
||||
shot on 35mm
|
||||
Microscopic
|
||||
Fisheye Lens
|
||||
Wide Angle
|
||||
Ultra-Wide Angle
|
||||
Panorama
|
||||
Short Exposure
|
||||
Long Exposure
|
||||
Double Exposure
|
||||
f2.8
|
||||
Depth of Field
|
||||
Soft Focus
|
||||
Deep Focus
|
||||
Shallow Focus
|
||||
Vanishing Point
|
||||
Vantage Point
|
||||
@@ -0,0 +1,30 @@
|
||||
Elegant evening gown
|
||||
Casual jeans and t-shirt
|
||||
Formal black suit
|
||||
Stylish leather jacket
|
||||
Flowy bohemian dress
|
||||
Sporty tracksuit
|
||||
Chic little black dress
|
||||
Trendy ripped jeans
|
||||
Classic white button-down shirt
|
||||
Cozy oversized sweater
|
||||
Sophisticated tailored blazer
|
||||
Quirky patterned leggings
|
||||
Striped sailor top
|
||||
Polished knee-length skirt
|
||||
Vintage-inspired floral dress
|
||||
Edgy motorcycle jacket
|
||||
Preppy polo shirt
|
||||
Boho maxi skirt
|
||||
Professional pinstripe suit
|
||||
Relaxed denim shorts
|
||||
Glamorous sequined dress
|
||||
Athletic running shoes
|
||||
Formal bow tie
|
||||
Casual baseball cap
|
||||
Stylish fedora hat
|
||||
Warm woolen scarf
|
||||
Comfortable cotton socks
|
||||
Trendy ankle boots
|
||||
Cute summer sandals
|
||||
Cozy pajama set
|
||||
@@ -0,0 +1,30 @@
|
||||
Happy
|
||||
Sad
|
||||
Angry
|
||||
Surprised
|
||||
Excited
|
||||
Worried
|
||||
Confused
|
||||
Disgusted
|
||||
Amused
|
||||
Bored
|
||||
Curious
|
||||
Embarrassed
|
||||
Frustrated
|
||||
Nervous
|
||||
Pleased
|
||||
Relieved
|
||||
Shy
|
||||
Tired
|
||||
Serious
|
||||
Silly
|
||||
Proud
|
||||
Grumpy
|
||||
Smug
|
||||
Sarcastic
|
||||
Flirty
|
||||
Skeptical
|
||||
Shocked
|
||||
Blissful
|
||||
Envious
|
||||
Mischievous
|
||||
@@ -0,0 +1,10 @@
|
||||
[
|
||||
{
|
||||
"keyword":"Dog",
|
||||
"imgurl":"http://127.0.0.1:8188/view?filename=1709966910233.png&type=input&subfolder=&rand=0.2734446552394221"
|
||||
},
|
||||
{
|
||||
"keyword":"x",
|
||||
"imgurl":"http://127.0.0.1:8188/view?filename=image%20(33).png&type=input&subfolder=pasted&rand=0.6984318219852814"
|
||||
}
|
||||
]
|
||||
@@ -0,0 +1,16 @@
|
||||
Mood Lighting
|
||||
Moody Lighting
|
||||
Studio Lighting
|
||||
Cove Lighting
|
||||
Soft Lighting
|
||||
Hard Lighting
|
||||
Volumetric Lighting
|
||||
Low-Key Lighting
|
||||
High-Key Lighting
|
||||
Epic Light
|
||||
Rembrandt Lighting
|
||||
Contre-Jour
|
||||
Veiling Flare
|
||||
Crepuscular Rays
|
||||
Rays of Shimmering Light
|
||||
Godrays
|
||||
@@ -0,0 +1 @@
|
||||
{}
|
||||
@@ -0,0 +1,132 @@
|
||||
Aaron Siskind
|
||||
Alessio Albi
|
||||
Alfred Eisenstaedt
|
||||
Alfred Stieglitz
|
||||
Alyssa Monks
|
||||
André Kertész
|
||||
Andreas Gursky
|
||||
Andrew Wyeth
|
||||
Anne Geddes
|
||||
Annie Leibovitz
|
||||
Ansel Adams
|
||||
Arnold Newman
|
||||
August Sander
|
||||
Balthus
|
||||
Berenice Abbott
|
||||
Bill Brandt
|
||||
Bill Henson
|
||||
Brassaï (Gyula Halász)
|
||||
Brooke Shaden
|
||||
Bruce Davidson
|
||||
Bruce Weber
|
||||
Bunny Yeager
|
||||
Carleton Watkins
|
||||
Carrie Mae Weems
|
||||
Chuck Close
|
||||
Cindy Sherman
|
||||
Clarence H. White
|
||||
Claude Cahun
|
||||
Danny Lyon
|
||||
David LaChapelle
|
||||
Dawoud Bey
|
||||
Diane Arbus
|
||||
Don McCullin
|
||||
Dora Maar
|
||||
Dorothea Lange
|
||||
Duane Michals
|
||||
Eadweard Muybridge
|
||||
Edward Burtynsky
|
||||
Edward Curtis
|
||||
Edward Ruscha
|
||||
Edward Steichen
|
||||
Edward Weston
|
||||
Elliott Erwitt
|
||||
Ernst Haas
|
||||
Eugene Atget
|
||||
Fan Ho
|
||||
Francesca Woodman
|
||||
Frans Lanting
|
||||
Garry Winogrand
|
||||
Georges Melies
|
||||
Gerda Taro
|
||||
Gertrude Käsebier
|
||||
Gordon Parks
|
||||
Graciela Iturbide
|
||||
Gregory Crewdson
|
||||
Harold Edgerton
|
||||
Helen Levitt
|
||||
Helmut Newton
|
||||
Hendrik Kerstens
|
||||
Henri Cartier-Bresson
|
||||
Hugh Kretschmer
|
||||
Irving Penn
|
||||
Jacques Henri Lartigue
|
||||
James Nachtwey
|
||||
James Van Der Zee
|
||||
Jay Maisel
|
||||
Jerry Uelsmann
|
||||
Joel Peter Witkin
|
||||
Joel Sartore
|
||||
John Frederick William Herschel
|
||||
Josef Sudek
|
||||
Julia Margaret Cameron
|
||||
Karl Blossfeldt
|
||||
Larry Burrows
|
||||
László Moholy-Nagy (photography)
|
||||
Lee Jeffries
|
||||
Lewis Hine
|
||||
Lorna Simpson
|
||||
Lynsey Addario
|
||||
Margaret Bourke-White
|
||||
Mario Testino
|
||||
Martin Parr
|
||||
Martin Schoeller
|
||||
Mary Ellen Mark
|
||||
Mathew B. Brady
|
||||
Méret Oppenheim
|
||||
Meryl McMaster
|
||||
Mick Rock
|
||||
Miles Aldridge
|
||||
Minor Martin White
|
||||
Nan Goldin
|
||||
Nathan Wirth
|
||||
Olive Cotton
|
||||
Olivier Rousteing
|
||||
Patrick Demarchelier
|
||||
Paul Nicklen
|
||||
Paul Outerbridge
|
||||
Paul Strand
|
||||
Pete Souza
|
||||
Peter Dombrovskis
|
||||
Peter Henry Emerson
|
||||
Peter Lik
|
||||
Peter Lindbergh
|
||||
Philip-Lorca diCorcia
|
||||
Philippe Halsman
|
||||
Ralph Gibson
|
||||
Richard Avedon
|
||||
Robert Adams
|
||||
Robert Bechtle
|
||||
Robert Capa
|
||||
Robert Frank
|
||||
Robert Mapplethorpe
|
||||
Roger Fenton
|
||||
Ruth Bernhard
|
||||
Sally Mann
|
||||
Sebastião Salgado
|
||||
Shirin Neshat
|
||||
Stefan Gesell
|
||||
Steven Meisel
|
||||
Susan Meiselas
|
||||
Vivian Maier
|
||||
Vivian Maier
|
||||
Viviane Sassen
|
||||
Walker Evans
|
||||
Wes Anderson
|
||||
William Eggleston
|
||||
William Eugene Smith
|
||||
William Henry Fox Talbot
|
||||
Yinka Shonibare
|
||||
Yousuf Karsh
|
||||
Man Ray
|
||||
Robert Mapplethorpe
|
||||
@@ -0,0 +1,101 @@
|
||||
Doctor
|
||||
Teacher
|
||||
Engineer
|
||||
Lawyer
|
||||
Accountant
|
||||
Nurse
|
||||
Architect
|
||||
Chef
|
||||
Pilot
|
||||
Scientist
|
||||
Artist
|
||||
Writer
|
||||
Musician
|
||||
Actor
|
||||
Photographer
|
||||
Police officer
|
||||
Firefighter
|
||||
Dentist
|
||||
Pharmacist
|
||||
Veterinarian
|
||||
Electrician
|
||||
Plumber
|
||||
Carpenter
|
||||
Mechanic
|
||||
Farmer
|
||||
Astronaut
|
||||
Athlete
|
||||
Journalist
|
||||
Politician
|
||||
Economist
|
||||
Psychologist
|
||||
Social worker
|
||||
Librarian
|
||||
Translator
|
||||
Salesperson
|
||||
Entrepreneur
|
||||
Financial advisor
|
||||
Graphic designer
|
||||
Web developer
|
||||
Marketing manager
|
||||
Human resources manager
|
||||
Project manager
|
||||
Event planner
|
||||
Fashion designer
|
||||
Interior decorator
|
||||
Real estate agent
|
||||
Archaeologist
|
||||
Biologist
|
||||
Chemist
|
||||
Geologist
|
||||
Physicist
|
||||
Mathematician
|
||||
Historian
|
||||
Geographer
|
||||
Economist
|
||||
Sociologist
|
||||
Anthropologist
|
||||
Archaeologist
|
||||
Linguist
|
||||
Philosopher
|
||||
Economist
|
||||
Sociologist
|
||||
Anthropologist
|
||||
Archaeologist
|
||||
Linguist
|
||||
Philosopher
|
||||
Geographer
|
||||
Historian
|
||||
Economist
|
||||
Sociologist
|
||||
Anthropologist
|
||||
Archaeologist
|
||||
Linguist
|
||||
Philosopher
|
||||
Geographer
|
||||
Historian
|
||||
Economist
|
||||
Sociologist
|
||||
Anthropologist
|
||||
Archaeologist
|
||||
Linguist
|
||||
Philosopher
|
||||
Geographer
|
||||
Historian
|
||||
Economist
|
||||
Sociologist
|
||||
Anthropologist
|
||||
Archaeologist
|
||||
Linguist
|
||||
Philosopher
|
||||
Geographer
|
||||
Historian
|
||||
Economist
|
||||
Sociologist
|
||||
Anthropologist
|
||||
Archaeologist
|
||||
Linguist
|
||||
Philosopher
|
||||
Geographer
|
||||
Historian
|
||||
#MixCopilot
|
||||
@@ -0,0 +1,58 @@
|
||||
Residential space
|
||||
Apartment building
|
||||
Villa
|
||||
Bungalow
|
||||
Condominium
|
||||
Commercial space
|
||||
Shopping mall
|
||||
Supermarket
|
||||
Restaurant
|
||||
Store
|
||||
Market
|
||||
Office space
|
||||
Office building
|
||||
Office
|
||||
Meeting room
|
||||
Co-working space
|
||||
Educational space
|
||||
School
|
||||
University
|
||||
Training institution
|
||||
Library
|
||||
Laboratory
|
||||
Medical space
|
||||
Hospital
|
||||
Clinic
|
||||
Pharmacy
|
||||
Nursing home
|
||||
Rehabilitation center
|
||||
Cultural space
|
||||
Museum
|
||||
Library
|
||||
Theater
|
||||
Concert hall
|
||||
Gallery
|
||||
Sports space
|
||||
Sports stadium
|
||||
Gym
|
||||
Swimming pool
|
||||
Basketball court
|
||||
Football field
|
||||
Transportation space
|
||||
Airport
|
||||
Train station
|
||||
Subway station
|
||||
Bus stop
|
||||
Parking lot
|
||||
Public space
|
||||
Park
|
||||
Square
|
||||
Street
|
||||
Pedestrian street
|
||||
Community center
|
||||
Industrial space
|
||||
Factory
|
||||
Warehouse
|
||||
Production workshop
|
||||
Mine
|
||||
Power plant
|
||||
@@ -0,0 +1,135 @@
|
||||
Vintage
|
||||
Grain
|
||||
Sepia
|
||||
High Key
|
||||
Low Key
|
||||
High Dynamic Range
|
||||
Cross Process
|
||||
Radial Blur
|
||||
Infrared
|
||||
Lomo
|
||||
Photocopy
|
||||
Pencil Sketch
|
||||
Pop Art
|
||||
Orton
|
||||
Mosaic
|
||||
Selective Black and White
|
||||
Torn Paper
|
||||
Tilt-Shift
|
||||
Double Exposure
|
||||
Polaroid
|
||||
Liquid Ink
|
||||
Color Splash
|
||||
Sketch
|
||||
Water Drops
|
||||
Polarizer
|
||||
Chinese Painting
|
||||
Water Droplets
|
||||
Polarization
|
||||
Color Inversion
|
||||
Fish-eye
|
||||
Soft Focus
|
||||
Solarization
|
||||
Posterize
|
||||
Comic Book
|
||||
Duotone
|
||||
Gradient Map
|
||||
Edge Detection
|
||||
Oil Painting
|
||||
Reflection
|
||||
Mirror
|
||||
ASCII Art
|
||||
Glitch
|
||||
Time-Lapse
|
||||
Day to Night
|
||||
Surreal
|
||||
Black and White
|
||||
Sepia Tone
|
||||
Vintage Film
|
||||
Grainy Texture
|
||||
High Key Lighting
|
||||
Low Key Lighting
|
||||
Cross Processed Film
|
||||
Infrared Photography
|
||||
Photocopy
|
||||
Pencil Drawing
|
||||
Pop Art Filter
|
||||
Mosaic Filter
|
||||
Selective Desaturation
|
||||
Torn Paper
|
||||
Tilt-Shift Photography
|
||||
Double Exposure
|
||||
Polaroid Style Frame
|
||||
Water Drops Texture
|
||||
Polarizer
|
||||
Chinese Painting
|
||||
Water Droplets Texture
|
||||
Polarization
|
||||
Color Inversion
|
||||
Fish-eye Lens
|
||||
Soft Focus
|
||||
Solarize Filter
|
||||
Edge Detection
|
||||
Oil Painting
|
||||
Reflection
|
||||
Mirror Image
|
||||
Time-Lapse Photography
|
||||
Day to Night Transition
|
||||
Surreal Art Style
|
||||
Abstract Expressionism
|
||||
Acrylic Painting
|
||||
Anime
|
||||
Art Deco
|
||||
Biomorphic Abstraction
|
||||
Black and White Photograph
|
||||
Cartoon
|
||||
Charcoal Sketch
|
||||
Chibi Anime
|
||||
Chinese Painting
|
||||
Classicist Painting
|
||||
Collage
|
||||
Concept Art
|
||||
Cyberpunk
|
||||
Dada Art
|
||||
Digital Art
|
||||
Fantasy Art
|
||||
Fashion Art
|
||||
Fashion Sketch
|
||||
Fish-Eye lens Photograph
|
||||
Goth Art
|
||||
Graffiti
|
||||
Harlem Renaissance
|
||||
High Key Photograph
|
||||
Hyperrealist Pencil Sketch
|
||||
Impressionist Painting
|
||||
Josei Anime
|
||||
Long Exposure Photograph
|
||||
Low Key Photograph
|
||||
Macro Photograph
|
||||
Manga
|
||||
Metal Sculpture
|
||||
Mid Century Modern Illustration
|
||||
Mixed Media
|
||||
Modern Art
|
||||
Moe Anime
|
||||
Nihonga
|
||||
Origami
|
||||
Paper Mache
|
||||
Pen and Ink
|
||||
Pencil Sketch
|
||||
Photograph
|
||||
Photorealism
|
||||
Pinup Art
|
||||
Romanticist Painting
|
||||
Sci-Fi Art
|
||||
Semi Realistic Fantasy Art
|
||||
Semi Realistic Cyberpunk Art
|
||||
Shallow Depth of Field Photograph
|
||||
Steam Punk Art
|
||||
Stone Sculpture
|
||||
Superhero Comic
|
||||
Surrealist Art
|
||||
Tempura Painting
|
||||
Underground Comic
|
||||
Watercolor Painting
|
||||
Zulu Urban Art
|
||||
@@ -10,6 +10,12 @@ if exist "%python_exec%" (
|
||||
for /f "delims=" %%i in (%requirements_txt%) do (
|
||||
%python_exec% -s -m pip install "%%i" -i https://pypi.tuna.tsinghua.edu.cn/simple
|
||||
)
|
||||
|
||||
@REM %python_exec% -s -m pip install --upgrade --force llama-cpp-python --extra-index-url https://abetlen.github.io/llama-cpp-python/whl/cu121
|
||||
|
||||
@REM %python_exec% -s -m pip install --upgrade --force llama-cpp-python[server]
|
||||
|
||||
|
||||
) else (
|
||||
echo Installing with system Python
|
||||
for /f "delims=" %%i in (%requirements_txt%) do (
|
||||
|
||||
@@ -0,0 +1,215 @@
|
||||
|
||||
import os
|
||||
import folder_paths
|
||||
import torchaudio
|
||||
|
||||
class AnyType(str):
|
||||
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
|
||||
|
||||
def __ne__(self, __value: object) -> bool:
|
||||
return False
|
||||
|
||||
any_type = AnyType("*")
|
||||
|
||||
|
||||
def analyze_audio_data(audio_data):
|
||||
total_duration = 0
|
||||
total_gap_duration = 0
|
||||
emotion_counts = {}
|
||||
audio_types = set()
|
||||
languages = set()
|
||||
|
||||
for i, entry in enumerate(audio_data):
|
||||
# Calculate the duration of each audio segment
|
||||
start_time = entry['start_time']
|
||||
end_time = entry['end_time']
|
||||
duration = end_time - start_time
|
||||
total_duration += duration
|
||||
|
||||
# Count the emotions
|
||||
if "emotion" in entry:
|
||||
emotion = entry['emotion']
|
||||
if emotion in emotion_counts:
|
||||
emotion_counts[emotion] += 1
|
||||
else:
|
||||
emotion_counts[emotion] = 1
|
||||
|
||||
# Collect the audio types
|
||||
if "audio_type" in entry:
|
||||
audio_types.add(entry['audio_type'])
|
||||
|
||||
if "language" in entry:
|
||||
languages.add(entry['language'])
|
||||
|
||||
# Calculate gap duration if not the last entry
|
||||
if i < len(audio_data) - 1:
|
||||
next_start_time = audio_data[i + 1]['start_time']
|
||||
gap_duration = next_start_time - end_time
|
||||
if gap_duration > 0:
|
||||
total_gap_duration += gap_duration
|
||||
|
||||
# Get the most frequent emotion
|
||||
if len(emotion_counts.keys())>0:
|
||||
most_frequent_emotion = max(emotion_counts, key=emotion_counts.get)
|
||||
else:
|
||||
most_frequent_emotion=None
|
||||
|
||||
# Convert audio_types set to list for better readability
|
||||
audio_types = list(audio_types)
|
||||
|
||||
languages=list(languages)
|
||||
|
||||
# Print the results
|
||||
print(f"Total Effective Duration: {total_duration:.2f} seconds")
|
||||
print(f"Total Gap Duration: {total_gap_duration:.2f} seconds")
|
||||
print(f"Emotion Changes: {emotion_counts}")
|
||||
print(f"Most Frequent Emotion: {most_frequent_emotion}")
|
||||
print(f"Audio Types: {audio_types}")
|
||||
|
||||
|
||||
return {
|
||||
"total_duration": total_duration,
|
||||
"total_gap_duration": total_gap_duration,
|
||||
"emotion_changes": emotion_counts,
|
||||
"most_frequent_emotion": most_frequent_emotion,
|
||||
"audio_types": audio_types,
|
||||
"languages":languages
|
||||
}
|
||||
|
||||
|
||||
# 分析音频数据
|
||||
class AnalyzeAudioNone:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"json":(any_type,),},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any_type,)
|
||||
RETURN_NAMES = ("result",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio"
|
||||
|
||||
def run(self,json):
|
||||
result=analyze_audio_data(json)
|
||||
return (result,)
|
||||
|
||||
|
||||
|
||||
class SpeechRecognition:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"upload":("AUDIOINPUTMIX",), },
|
||||
"optional":{
|
||||
"start_by":("INT", {
|
||||
"default": 0,
|
||||
"min": 0, #Minimum value
|
||||
"max": 2048, #Maximum value
|
||||
"step": 1, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("prompt",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
def run(self,upload,start_by):
|
||||
return {"ui": {"start_by": [start_by]}, "result": (upload,)}
|
||||
|
||||
|
||||
class SpeechSynthesis:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"forceInput": True}),
|
||||
}
|
||||
}
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "run"
|
||||
OUTPUT_NODE = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio"
|
||||
|
||||
def run(self, text):
|
||||
# print(session_history)
|
||||
return {"ui": {"text": text}, "result": (text,)}
|
||||
|
||||
|
||||
|
||||
class AudioPlayNode:
|
||||
def __init__(self):
|
||||
self.output_dir = folder_paths.get_temp_directory()
|
||||
self.type = "temp"
|
||||
self.prefix_append = ""
|
||||
self.compress_level = 4
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"audio": ("AUDIO",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = ()
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def run(self,audio):
|
||||
|
||||
# 判断是否是 Tensor 类型
|
||||
is_tensor = not isinstance(audio, dict)
|
||||
# print('#判断是否是 Tensor 类型',is_tensor,audio)
|
||||
if not is_tensor and 'waveform' in audio and 'sample_rate' in audio:
|
||||
# {'waveform': tensor([], size=(1, 1, 0)), 'sample_rate': 44100}
|
||||
is_tensor=True
|
||||
|
||||
if is_tensor and (not 'audio_path' in audio):
|
||||
filename_prefix=""
|
||||
# 保存
|
||||
filename_prefix += self.prefix_append
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, self.output_dir)
|
||||
results = list()
|
||||
|
||||
filename_with_batch_num = filename.replace("%batch_num%", str(1))
|
||||
file = f"{filename_with_batch_num}_{counter:05}_.wav"
|
||||
|
||||
torchaudio.save(os.path.join(full_output_folder, file), audio['waveform'].squeeze(0), audio["sample_rate"])
|
||||
results.append({
|
||||
"filename": file,
|
||||
"subfolder": subfolder,
|
||||
"type": self.type
|
||||
})
|
||||
|
||||
else:
|
||||
results=[{
|
||||
"filename": audio['filename'],
|
||||
"subfolder":audio['subfolder'],
|
||||
"type": audio['type'],
|
||||
"audio_path":audio['audio_path']
|
||||
}]
|
||||
|
||||
|
||||
# print(audio)
|
||||
return {"ui": {"audio":results}}
|
||||
@@ -0,0 +1,283 @@
|
||||
import os,sys
|
||||
import folder_paths
|
||||
|
||||
from PIL import Image
|
||||
import importlib.util
|
||||
|
||||
import comfy.utils
|
||||
import numpy as np
|
||||
import json
|
||||
import torch
|
||||
import random
|
||||
|
||||
|
||||
# from clip_interrogator import Config, Interrogator
|
||||
|
||||
|
||||
global _available
|
||||
_available=False
|
||||
|
||||
def is_installed(package):
|
||||
try:
|
||||
spec = importlib.util.find_spec(package)
|
||||
except ModuleNotFoundError:
|
||||
return False
|
||||
return spec is not None
|
||||
|
||||
|
||||
try:
|
||||
if is_installed('clip_interrogator')==False:
|
||||
import subprocess
|
||||
|
||||
# 安装
|
||||
print('#pip install clip-interrogator==0.6.0')
|
||||
|
||||
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'clip-interrogator==0.6.0'], capture_output=True, text=True)
|
||||
|
||||
#检查命令执行结果
|
||||
if result.returncode == 0:
|
||||
print("#install success")
|
||||
from clip_interrogator import Config, Interrogator
|
||||
_available=True
|
||||
else:
|
||||
print("#install error")
|
||||
|
||||
else:
|
||||
from clip_interrogator import Config, Interrogator
|
||||
_available=True
|
||||
|
||||
except:
|
||||
_available=False
|
||||
|
||||
try:
|
||||
from transformers import AutoProcessor, BlipForConditionalGeneration
|
||||
except:
|
||||
_available=False
|
||||
print('pls check transformers.__version__>=4.36.0:: AutoProcessor, BlipForConditionalGeneration')
|
||||
|
||||
|
||||
|
||||
def load_caption_model(model_path,config,t='blip-base'):
|
||||
dtype=torch.float16 if config.device == 'cuda' else torch.float32
|
||||
caption_model = BlipForConditionalGeneration.from_pretrained(model_path, torch_dtype=dtype)
|
||||
|
||||
caption_processor = AutoProcessor.from_pretrained(model_path)
|
||||
|
||||
caption_model.eval()
|
||||
if not config.caption_offload:
|
||||
caption_model = caption_model.to(config.device)
|
||||
|
||||
return (caption_model,caption_processor)
|
||||
|
||||
|
||||
def get_clip_interrogator_path():
|
||||
try:
|
||||
return folder_paths.get_folder_paths('clip_interrogator')[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, "clip_interrogator")
|
||||
|
||||
|
||||
cache_path=get_clip_interrogator_path()
|
||||
|
||||
caption_model_path=os.path.join(cache_path, "Salesforce","blip-image-captioning-base")
|
||||
if not os.path.exists(caption_model_path):
|
||||
print(f"## clip_interrogator_model not found: {caption_model_path}, pls download from https://huggingface.co/Salesforce/blip-image-captioning-base")
|
||||
caption_model_path='Salesforce/blip-image-captioning-base'
|
||||
|
||||
|
||||
# Tensor to PIL
|
||||
def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
# Convert PIL to Tensor
|
||||
def pil2tensor(image):
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
|
||||
def image_analysis_fn(ci,image):
|
||||
image = image.convert('RGB')
|
||||
image_features = ci.image_to_features(image)
|
||||
|
||||
top_mediums = ci.mediums.rank(image_features, 5)
|
||||
top_artists = ci.artists.rank(image_features, 5)
|
||||
top_movements = ci.movements.rank(image_features, 5)
|
||||
top_trendings = ci.trendings.rank(image_features, 5)
|
||||
top_flavors = ci.flavors.rank(image_features, 5)
|
||||
|
||||
medium_ranks = {medium: sim for medium, sim in zip(top_mediums, ci.similarities(image_features, top_mediums))}
|
||||
artist_ranks = {artist: sim for artist, sim in zip(top_artists, ci.similarities(image_features, top_artists))}
|
||||
movement_ranks = {movement: sim for movement, sim in zip(top_movements, ci.similarities(image_features, top_movements))}
|
||||
trending_ranks = {trending: sim for trending, sim in zip(top_trendings, ci.similarities(image_features, top_trendings))}
|
||||
flavor_ranks = {flavor: sim for flavor, sim in zip(top_flavors, ci.similarities(image_features, top_flavors))}
|
||||
|
||||
return {
|
||||
"medium_ranks":medium_ranks,
|
||||
"artist_ranks":artist_ranks,
|
||||
"movement_ranks":movement_ranks,
|
||||
"trending_ranks":trending_ranks,
|
||||
"flavor_ranks":flavor_ranks
|
||||
}
|
||||
|
||||
|
||||
def generate_sentences(data):
|
||||
sentences = []
|
||||
|
||||
# Get the length of data
|
||||
data_length = len(data)
|
||||
|
||||
# Use a recursive function to handle variable-length data
|
||||
def generate_recursive(index, current_sentence, current_score):
|
||||
# Check if recursion is complete
|
||||
if index == data_length:
|
||||
sentences.append({"sentence": current_sentence, "score": current_score})
|
||||
return
|
||||
|
||||
# Get the current level data
|
||||
current_data = data[index]
|
||||
|
||||
# Iterate through the current level data
|
||||
for phrase in current_data:
|
||||
sentence = current_sentence + ("," if current_sentence.strip() else "") + phrase
|
||||
score = current_score + current_data[phrase]
|
||||
generate_recursive(index + 1, sentence, score)
|
||||
|
||||
# Start recursive generation of sentences
|
||||
generate_recursive(0, "", 0)
|
||||
|
||||
# Sort the generated sentences by score in descending order
|
||||
sentences.sort(key=lambda x: x["score"], reverse=True)
|
||||
|
||||
def get_random_elements(elements, num):
|
||||
return random.sample(elements, num)
|
||||
|
||||
ps = get_random_elements(sentences, 5)
|
||||
ps = [s["sentence"] for s in sorted(ps, key=lambda x: x["score"], reverse=True)]
|
||||
|
||||
return ps
|
||||
|
||||
|
||||
|
||||
|
||||
def image_to_prompt(ci,image, mode):
|
||||
ci.config.chunk_size = 2048 if ci.config.clip_model_name == "ViT-L-14/openai" else 1024
|
||||
ci.config.flavor_intermediate_count = 2048 if ci.config.clip_model_name == "ViT-L-14/openai" else 1024
|
||||
image = image.convert('RGB')
|
||||
if mode == 'best':
|
||||
return ci.interrogate(image)
|
||||
elif mode == 'classic':
|
||||
return ci.interrogate_classic(image)
|
||||
elif mode == 'fast':
|
||||
return ci.interrogate_fast(image)
|
||||
elif mode == 'negative':
|
||||
return ci.interrogate_negative(image)
|
||||
|
||||
# image = Image.open(image_path).convert('RGB')
|
||||
# ci = Interrogator(Config(clip_model_name="ViT-L-14/openai"))
|
||||
# print(ci.interrogate(image))
|
||||
|
||||
|
||||
class ClipInterrogator:
|
||||
|
||||
global _available
|
||||
available=_available
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"image": ("IMAGE",),
|
||||
"prompt_mode": (['fast','classic','best','negative'],),
|
||||
"image_analysis": (["off","on"],),
|
||||
},
|
||||
|
||||
# "optional":{
|
||||
# "output":("CLIPINTERROGATOR", {"multiline": True,"default": "", "dynamicPrompts": False})
|
||||
# },
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING","STRING",)
|
||||
RETURN_NAMES = ("prompt","random_samples",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,True,)
|
||||
global ci
|
||||
ci = None
|
||||
def run(self,image,prompt_mode,image_analysis):
|
||||
global ci
|
||||
|
||||
prompt_mode=prompt_mode[0]
|
||||
analysis=image_analysis[0]
|
||||
|
||||
prompt_result=[]
|
||||
analysis_result=[]
|
||||
|
||||
# 进度条
|
||||
pbar = comfy.utils.ProgressBar(len(image)*(2 if analysis=='on' else 1))
|
||||
|
||||
if ci==None:
|
||||
config=Config(
|
||||
clip_model_name="ViT-L-14/openai",
|
||||
device="cuda" if torch.cuda.is_available() else "cpu",
|
||||
download_cache=True,
|
||||
clip_model_path=cache_path,
|
||||
cache_path=cache_path
|
||||
)
|
||||
config.apply_low_vram_defaults()
|
||||
|
||||
caption_model,caption_processor=load_caption_model(caption_model_path,config)
|
||||
|
||||
config.caption_model= caption_model
|
||||
config.caption_processor= caption_processor
|
||||
|
||||
ci = Interrogator(config)
|
||||
# else:
|
||||
# simple_lama.model.to("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
for i in range(len(image)):
|
||||
im=image[i]
|
||||
|
||||
im=tensor2pil(im)
|
||||
im=im.convert('RGB')
|
||||
|
||||
if analysis=='on':
|
||||
analysis_res=image_analysis_fn(ci,im)
|
||||
analysis_result.append( analysis_res )
|
||||
pbar.update(1)
|
||||
|
||||
prompt=image_to_prompt(ci,im,prompt_mode)
|
||||
pbar.update(1)
|
||||
prompt_result.append(prompt)
|
||||
|
||||
|
||||
# result.save("inpainted.png")
|
||||
if ci.config.clip_offload and not ci.clip_offloaded:
|
||||
ci.clip_model = ci.clip_model.to('cpu')
|
||||
ci.clip_offloaded = True
|
||||
|
||||
if ci.config.caption_offload and not ci.caption_offloaded:
|
||||
ci.caption_model = ci.caption_model.to('cpu')
|
||||
ci.caption_offloaded = True
|
||||
|
||||
# analysis_result=[]
|
||||
# items = app.graph.getNodeById(31).widgets[2].value["items"]
|
||||
|
||||
random_samples=[]
|
||||
|
||||
for r in analysis_result:
|
||||
random_sample = generate_sentences([r['medium_ranks'], r['artist_ranks'],r['movement_ranks'],r['trending_ranks'],r['flavor_ranks']])
|
||||
for s in random_sample:
|
||||
random_samples.append(s)
|
||||
# print(len(random_samples))
|
||||
# print('-----')
|
||||
# print( random_samples)
|
||||
return {
|
||||
"ui":{
|
||||
"prompt": prompt_result,
|
||||
"analysis":analysis_result,
|
||||
"random_samples":random_samples
|
||||
},
|
||||
"result": (prompt_result,random_samples,)}
|
||||
@@ -1,258 +0,0 @@
|
||||
from transformers import CLIPSegProcessor, CLIPSegForImageSegmentation
|
||||
|
||||
from PIL import Image
|
||||
import torch
|
||||
import torchvision.transforms as T
|
||||
import numpy as np
|
||||
|
||||
from torchvision.transforms.functional import to_pil_image
|
||||
import matplotlib.pyplot as plt
|
||||
import matplotlib.cm as cm
|
||||
|
||||
|
||||
import cv2
|
||||
|
||||
from scipy.ndimage import gaussian_filter
|
||||
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import warnings,os
|
||||
warnings.filterwarnings("ignore", category=UserWarning, module="torch")
|
||||
warnings.filterwarnings("ignore", category=UserWarning, module="safetensors")
|
||||
|
||||
import folder_paths
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger('CLIPSeg nodes')
|
||||
|
||||
clipseg_model_dir = os.path.join(folder_paths.models_dir, "clipseg")
|
||||
|
||||
if not os.path.exists(clipseg_model_dir):
|
||||
clipseg_model_dir='CIDAS/clipseg-rd64-refined'
|
||||
|
||||
"""Helper methods for CLIPSeg nodes"""
|
||||
|
||||
def tensor_to_numpy(tensor: torch.Tensor) -> np.ndarray:
|
||||
"""Convert a tensor to a numpy array and scale its values to 0-255."""
|
||||
array = tensor.numpy().squeeze()
|
||||
return (array * 255).astype(np.uint8)
|
||||
|
||||
def numpy_to_tensor(array: np.ndarray) -> torch.Tensor:
|
||||
"""Convert a numpy array to a tensor and scale its values from 0-255 to 0-1."""
|
||||
array = array.astype(np.float32) / 255.0
|
||||
return torch.from_numpy(array)[None,]
|
||||
|
||||
def apply_colormap(mask: torch.Tensor, colormap) -> np.ndarray:
|
||||
"""Apply a colormap to a tensor and convert it to a numpy array."""
|
||||
colored_mask = colormap(mask.numpy())[:, :, :3]
|
||||
return (colored_mask * 255).astype(np.uint8)
|
||||
|
||||
def resize_image(image: np.ndarray, dimensions: Tuple[int, int]) -> np.ndarray:
|
||||
"""Resize an image to the given dimensions using linear interpolation."""
|
||||
return cv2.resize(image, dimensions, interpolation=cv2.INTER_LINEAR)
|
||||
|
||||
def overlay_image(background: np.ndarray, foreground: np.ndarray, alpha: float) -> np.ndarray:
|
||||
"""Overlay the foreground image onto the background with a given opacity (alpha)."""
|
||||
return cv2.addWeighted(background, 1 - alpha, foreground, alpha, 0)
|
||||
|
||||
def dilate_mask(mask: torch.Tensor, dilation_factor: float) -> torch.Tensor:
|
||||
"""Dilate a mask using a square kernel with a given dilation factor."""
|
||||
kernel_size = int(dilation_factor * 2) + 1
|
||||
kernel = np.ones((kernel_size, kernel_size), np.uint8)
|
||||
mask_dilated = cv2.dilate(mask.numpy(), kernel, iterations=1)
|
||||
return torch.from_numpy(mask_dilated)
|
||||
|
||||
|
||||
|
||||
class CLIPSeg:
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
"""
|
||||
Return a dictionary which contains config for all input fields.
|
||||
Some types (string): "MODEL", "VAE", "CLIP", "CONDITIONING", "LATENT", "IMAGE", "INT", "STRING", "FLOAT".
|
||||
Input types "INT", "STRING" or "FLOAT" are special values for fields on the node.
|
||||
The type can be a list for selection.
|
||||
|
||||
Returns: `dict`:
|
||||
- Key input_fields_group (`string`): Can be either required, hidden or optional. A node class must have property `required`
|
||||
- Value input_fields (`dict`): Contains input fields config:
|
||||
* Key field_name (`string`): Name of a entry-point method's argument
|
||||
* Value field_config (`tuple`):
|
||||
+ First value is a string indicate the type of field or a list for selection.
|
||||
+ Secound value is a config for type "INT", "STRING" or "FLOAT".
|
||||
"""
|
||||
return {"required":
|
||||
{
|
||||
"image": ("IMAGE",),
|
||||
"text": ("STRING", {"multiline": False}),
|
||||
|
||||
},
|
||||
"optional":
|
||||
{
|
||||
"blur": ("FLOAT", {"min": 0, "max": 15, "step": 0.1, "default": 7}),
|
||||
"threshold": ("FLOAT", {"min": 0, "max": 1, "step": 0.05, "default": 0.4}),
|
||||
"dilation_factor": ("INT", {"min": 0, "max": 10, "step": 1, "default": 4}),
|
||||
}
|
||||
}
|
||||
|
||||
CATEGORY = "Mixlab/mask"
|
||||
RETURN_TYPES = ("MASK", "IMAGE", "IMAGE",)
|
||||
RETURN_NAMES = ("Mask","Heatmap Mask", "BW Mask")
|
||||
|
||||
# INPUT_IS_LIST = True
|
||||
# OUTPUT_IS_LIST = (True,)
|
||||
|
||||
FUNCTION = "segment_image"
|
||||
def segment_image(self, image: torch.Tensor, text: str, blur: float, threshold: float, dilation_factor: int) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Create a segmentation mask from an image and a text prompt using CLIPSeg.
|
||||
|
||||
Args:
|
||||
image (torch.Tensor): The image to segment.
|
||||
text (str): The text prompt to use for segmentation.
|
||||
blur (float): How much to blur the segmentation mask.
|
||||
threshold (float): The threshold to use for binarizing the segmentation mask.
|
||||
dilation_factor (int): How much to dilate the segmentation mask.
|
||||
|
||||
Returns:
|
||||
Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: The segmentation mask, the heatmap mask, and the binarized mask.
|
||||
"""
|
||||
|
||||
# Convert the Tensor to a PIL image
|
||||
image_np = image.numpy().squeeze() # Remove the first dimension (batch size of 1)
|
||||
# Convert the numpy array back to the original range (0-255) and data type (uint8)
|
||||
image_np = (image_np * 255).astype(np.uint8)
|
||||
# Create a PIL image from the numpy array
|
||||
i = Image.fromarray(image_np, mode="RGB")
|
||||
|
||||
processor = CLIPSegProcessor.from_pretrained(clipseg_model_dir)
|
||||
model = CLIPSegForImageSegmentation.from_pretrained(clipseg_model_dir)
|
||||
|
||||
prompt = text
|
||||
|
||||
input_prc = processor(text=prompt, images=i, padding="max_length", return_tensors="pt")
|
||||
|
||||
# Predict the segemntation mask
|
||||
with torch.no_grad():
|
||||
outputs = model(**input_prc)
|
||||
|
||||
tensor = torch.sigmoid(outputs[0]) # get the mask
|
||||
|
||||
# Apply a threshold to the original tensor to cut off low values
|
||||
thresh = threshold
|
||||
tensor_thresholded = torch.where(tensor > thresh, tensor, torch.tensor(0, dtype=torch.float))
|
||||
|
||||
# Apply Gaussian blur to the thresholded tensor
|
||||
sigma = blur
|
||||
tensor_smoothed = gaussian_filter(tensor_thresholded.numpy(), sigma=sigma)
|
||||
tensor_smoothed = torch.from_numpy(tensor_smoothed)
|
||||
|
||||
# Normalize the smoothed tensor to [0, 1]
|
||||
mask_normalized = (tensor_smoothed - tensor_smoothed.min()) / (tensor_smoothed.max() - tensor_smoothed.min())
|
||||
|
||||
# Dilate the normalized mask
|
||||
mask_dilated = dilate_mask(mask_normalized, dilation_factor)
|
||||
|
||||
# Convert the mask to a heatmap and a binary mask
|
||||
heatmap = apply_colormap(mask_dilated, cm.viridis)
|
||||
binary_mask = apply_colormap(mask_dilated, cm.Greys_r)
|
||||
|
||||
# Overlay the heatmap and binary mask on the original image
|
||||
dimensions = (image_np.shape[1], image_np.shape[0])
|
||||
heatmap_resized = resize_image(heatmap, dimensions)
|
||||
binary_mask_resized = resize_image(binary_mask, dimensions)
|
||||
|
||||
alpha_heatmap, alpha_binary = 0.5, 1
|
||||
overlay_heatmap = overlay_image(image_np, heatmap_resized, alpha_heatmap)
|
||||
overlay_binary = overlay_image(image_np, binary_mask_resized, alpha_binary)
|
||||
|
||||
# Convert the numpy arrays to tensors
|
||||
image_out_heatmap = numpy_to_tensor(overlay_heatmap)
|
||||
image_out_binary = numpy_to_tensor(overlay_binary)
|
||||
|
||||
# Save or display the resulting binary mask
|
||||
binary_mask_image = Image.fromarray(binary_mask_resized[..., 0])
|
||||
|
||||
# convert PIL image to numpy array
|
||||
tensor_bw = binary_mask_image.convert("RGB")
|
||||
tensor_bw = np.array(tensor_bw).astype(np.float32) / 255.0
|
||||
tensor_bw = torch.from_numpy(tensor_bw)[None,]
|
||||
tensor_bw = tensor_bw.squeeze(0)[..., 0]
|
||||
|
||||
return tensor_bw, image_out_heatmap, image_out_binary
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
class CombineMasks:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{
|
||||
"input_image": ("IMAGE", ),
|
||||
"mask_1": ("MASK", ),
|
||||
"mask_2": ("MASK", ),
|
||||
},
|
||||
"optional":
|
||||
{
|
||||
"mask_3": ("MASK",),
|
||||
},
|
||||
}
|
||||
|
||||
CATEGORY = "Mixlab/mask"
|
||||
RETURN_TYPES = ("MASK", "IMAGE", "IMAGE",)
|
||||
RETURN_NAMES = ("Combined Mask","Heatmap Mask", "BW Mask")
|
||||
|
||||
FUNCTION = "combine_masks"
|
||||
|
||||
def combine_masks(self, input_image: torch.Tensor, mask_1: torch.Tensor, mask_2: torch.Tensor, mask_3: Optional[torch.Tensor] = None) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""A method that combines two or three masks into one mask. Takes in tensors and returns the mask as a tensor, as well as the heatmap and binary mask as tensors."""
|
||||
|
||||
# Combine masks
|
||||
if mask_1 is not None:
|
||||
mask_1 = mask_1.squeeze()
|
||||
if mask_2 is not None:
|
||||
mask_2 = mask_2.squeeze()
|
||||
if mask_3 is not None:
|
||||
mask_3 = mask_3.squeeze()
|
||||
|
||||
print(mask_1.shape,mask_2.shape , mask_3.shape)
|
||||
combined_mask = mask_1 + mask_2 + mask_3 if mask_3 is not None else mask_1 + mask_2
|
||||
# print(combined_mask)
|
||||
|
||||
# Convert image and masks to numpy arrays
|
||||
image_np = tensor_to_numpy(input_image)
|
||||
heatmap = apply_colormap(combined_mask, cm.viridis)
|
||||
binary_mask = apply_colormap(combined_mask, cm.Greys_r)
|
||||
|
||||
# Resize heatmap and binary mask to match the original image dimensions
|
||||
dimensions = (image_np.shape[1], image_np.shape[0])
|
||||
print('heatmap',heatmap)
|
||||
if dimensions is None or dimensions[0] == 0 or dimensions[1] == 0:
|
||||
raise ValueError("Invalid dimensions")
|
||||
|
||||
heatmap_resized = resize_image(heatmap, dimensions)
|
||||
binary_mask_resized = resize_image(binary_mask, dimensions)
|
||||
|
||||
# Overlay the heatmap and binary mask onto the original image
|
||||
alpha_heatmap, alpha_binary = 0.5, 1
|
||||
overlay_heatmap = overlay_image(image_np, heatmap_resized, alpha_heatmap)
|
||||
overlay_binary = overlay_image(image_np, binary_mask_resized, alpha_binary)
|
||||
|
||||
# Convert overlays to tensors
|
||||
image_out_heatmap = numpy_to_tensor(overlay_heatmap)
|
||||
image_out_binary = numpy_to_tensor(overlay_binary)
|
||||
|
||||
return combined_mask, image_out_heatmap, image_out_binary
|
||||
|
||||
# A dictionary that contains all nodes you want to export with their names
|
||||
# NOTE: names should be globally unique
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"CLIPSeg": CLIPSeg,
|
||||
"CombineSegMasks": CombineMasks,
|
||||
}
|
||||
@@ -0,0 +1,332 @@
|
||||
# 修改自 https://github.com/gokayfem/ComfyUI-fal-API/blob/main/nodes/video_node.py
|
||||
# image-to-video all in one
|
||||
|
||||
import os,sys
|
||||
import torch
|
||||
from PIL import Image
|
||||
import tempfile
|
||||
import numpy as np
|
||||
import requests
|
||||
import cv2
|
||||
import subprocess
|
||||
import importlib.util
|
||||
python = sys.executable
|
||||
|
||||
def is_installed(package, package_overwrite=None,auto_install=True):
|
||||
is_has=False
|
||||
try:
|
||||
spec = importlib.util.find_spec(package)
|
||||
is_has=spec is not None
|
||||
except ModuleNotFoundError:
|
||||
pass
|
||||
|
||||
package = package_overwrite or package
|
||||
|
||||
if spec is None:
|
||||
if auto_install==True:
|
||||
print(f"Installing {package}...")
|
||||
# 清华源 -i https://pypi.tuna.tsinghua.edu.cn/simple
|
||||
command = f'"{python}" -m pip install {package}'
|
||||
|
||||
result = subprocess.run(command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, shell=True, env=os.environ)
|
||||
|
||||
is_has=True
|
||||
|
||||
if result.returncode != 0:
|
||||
print(f"Couldn't install\nCommand: {command}\nError code: {result.returncode}")
|
||||
is_has=False
|
||||
else:
|
||||
print(package+'## OK')
|
||||
|
||||
return is_has
|
||||
|
||||
|
||||
try:
|
||||
if is_installed('fal_client','fal-client')==True:
|
||||
from fal_client import submit, upload_file
|
||||
except:
|
||||
print("#install fal-client error")
|
||||
|
||||
|
||||
def upload_image(image):
|
||||
try:
|
||||
# Convert the image tensor to a numpy array
|
||||
if isinstance(image, torch.Tensor):
|
||||
image_np = image.cpu().numpy()
|
||||
else:
|
||||
image_np = np.array(image)
|
||||
|
||||
# Ensure the image is in the correct format (H, W, C)
|
||||
if image_np.ndim == 4:
|
||||
image_np = image_np.squeeze(0) # Remove batch dimension if present
|
||||
if image_np.ndim == 2:
|
||||
image_np = np.stack([image_np] * 3, axis=-1) # Convert grayscale to RGB
|
||||
elif image_np.shape[0] == 3:
|
||||
image_np = np.transpose(image_np, (1, 2, 0)) # Change from (C, H, W) to (H, W, C)
|
||||
|
||||
# Normalize the image data to 0-255 range
|
||||
if image_np.dtype == np.float32 or image_np.dtype == np.float64:
|
||||
image_np = (image_np * 255).astype(np.uint8)
|
||||
|
||||
# Convert to PIL Image
|
||||
pil_image = Image.fromarray(image_np)
|
||||
|
||||
# Save the image to a temporary file
|
||||
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as temp_file:
|
||||
pil_image.save(temp_file, format="PNG")
|
||||
temp_file_path = temp_file.name
|
||||
|
||||
# Upload the temporary file
|
||||
image_url = upload_file(temp_file_path)
|
||||
return image_url
|
||||
except Exception as e:
|
||||
print(f"Error uploading image: {str(e)}")
|
||||
return None
|
||||
finally:
|
||||
# Clean up the temporary file
|
||||
if 'temp_file_path' in locals():
|
||||
os.unlink(temp_file_path)
|
||||
|
||||
|
||||
class VideoGenKlingNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"duration": (["5", "10"], {"default": "5"}),
|
||||
"aspect_ratio": (["16:9", "9:16", "1:1"], {"default": "16:9"}),
|
||||
"mode": (["standard", "pro"], {"default": "standard"}),
|
||||
"fal_key":("STRING", {"forceInput": True,}),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "generate_video"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
def generate_video(self, prompt, duration, aspect_ratio,mode,fal_key, image=None):
|
||||
arguments = {
|
||||
"prompt": prompt,
|
||||
"duration": duration,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
}
|
||||
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
|
||||
api_url="fal-ai/kling-video/v1/"+mode
|
||||
|
||||
try:
|
||||
if image is not None:
|
||||
image_url = upload_image(image)
|
||||
if image_url:
|
||||
arguments["image_url"] = image_url
|
||||
handler = submit(api_url+"/image-to-video", arguments=arguments)
|
||||
else:
|
||||
return ("Error: Unable to upload image.",)
|
||||
else:
|
||||
handler = submit(api_url+"/text-to-video", arguments=arguments)
|
||||
|
||||
result = handler.get()
|
||||
video_url = result["video"]["url"]
|
||||
return (video_url,)
|
||||
except Exception as e:
|
||||
print(f"Error generating video: {str(e)}")
|
||||
return ("Error: Unable to generate video.",)
|
||||
|
||||
|
||||
class VideoGenRunwayGen3Node:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"image": ("IMAGE",),
|
||||
"duration": (["5", "10"], {"default": "5"}),
|
||||
"aspect_ratio": (["16:9", "9:16"], {"default": "16:9"}),
|
||||
"fal_key":("STRING", {"forceInput": True,}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "generate_video"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
def generate_video(self, prompt, image, duration,aspect_ratio,fal_key):
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
try:
|
||||
image_url = upload_image(image)
|
||||
if not image_url:
|
||||
return ("Error: Unable to upload image.",)
|
||||
|
||||
arguments = {
|
||||
"prompt": prompt,
|
||||
"image_url": image_url,
|
||||
"duration": duration,
|
||||
"ratio":aspect_ratio
|
||||
}
|
||||
|
||||
handler = submit("fal-ai/runway-gen3/turbo/image-to-video", arguments=arguments)
|
||||
result = handler.get()
|
||||
video_url = result["video"]["url"]
|
||||
return (video_url,)
|
||||
except Exception as e:
|
||||
print(f"Error generating video: {str(e)}")
|
||||
return ("Error: Unable to generate video.",)
|
||||
|
||||
class VideoGenLumaDreamMachineNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
"aspect_ratio": (["16:9", "9:16", "4:3", "3:4", "21:9", "9:21"], {"default": "16:9"}),
|
||||
"fal_key":("STRING", {"forceInput": True,}),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE",),
|
||||
"loop": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "generate_video"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
def generate_video(self, prompt, aspect_ratio,fal_key, image=None, loop=True):
|
||||
|
||||
os.environ["FAL_KEY"] = fal_key
|
||||
|
||||
arguments = {
|
||||
"prompt": prompt,
|
||||
"aspect_ratio": aspect_ratio,
|
||||
"loop": loop,
|
||||
}
|
||||
|
||||
try:
|
||||
if image is not None:
|
||||
image_url = upload_image(image)
|
||||
if not image_url:
|
||||
return ("Error: Unable to upload image.",)
|
||||
arguments["image_url"] = image_url
|
||||
endpoint = "fal-ai/luma-dream-machine/image-to-video"
|
||||
else:
|
||||
endpoint = "fal-ai/luma-dream-machine"
|
||||
|
||||
handler = submit(endpoint, arguments=arguments)
|
||||
result = handler.get()
|
||||
video_url = result["video"]["url"]
|
||||
return (video_url,)
|
||||
except Exception as e:
|
||||
print(f"Error generating video: {str(e)}")
|
||||
return ("Error: Unable to generate video.",)
|
||||
|
||||
class LoadVideoFromURL:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"url": ("STRING", {"default": "https://example.com/video.mp4"}),
|
||||
"force_rate": ("INT", {"default": 0, "min": 0, "max": 60, "step": 1}),
|
||||
"force_size": (["Disabled", "Custom Height", "Custom Width", "Custom", "256x?", "?x256", "256x256", "512x?", "?x512", "512x512"],),
|
||||
"custom_width": ("INT", {"default": 512, "min": 0, "max": 8192, "step": 8}),
|
||||
"custom_height": ("INT", {"default": 512, "min": 0, "max": 8192, "step": 8}),
|
||||
"frame_load_cap": ("INT", {"default": 0, "min": 0, "max": 1000000, "step": 1}),
|
||||
"skip_first_frames": ("INT", {"default": 0, "min": 0, "max": 1000000, "step": 1}),
|
||||
"select_every_nth": ("INT", {"default": 1, "min": 1, "max": 1000000, "step": 1}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "INT", "VHS_VIDEOINFO")
|
||||
RETURN_NAMES = ("frames", "frame_count", "video_info")
|
||||
FUNCTION = "load_video_from_url"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
def load_video_from_url(self, url, force_rate, force_size, custom_width, custom_height, frame_load_cap, skip_first_frames, select_every_nth):
|
||||
# Download the video to a temporary file
|
||||
with tempfile.NamedTemporaryFile(delete=False, suffix=".mp4") as temp_file:
|
||||
response = requests.get(url, stream=True)
|
||||
for chunk in response.iter_content(chunk_size=8192):
|
||||
temp_file.write(chunk)
|
||||
temp_file_path = temp_file.name
|
||||
|
||||
# Load the video using OpenCV
|
||||
cap = cv2.VideoCapture(temp_file_path)
|
||||
|
||||
# Get video properties
|
||||
fps = cap.get(cv2.CAP_PROP_FPS)
|
||||
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
duration = total_frames / fps
|
||||
|
||||
# Calculate target size
|
||||
if force_size != "Disabled":
|
||||
if force_size == "Custom Width":
|
||||
new_height = int(height * (custom_width / width))
|
||||
new_width = custom_width
|
||||
elif force_size == "Custom Height":
|
||||
new_width = int(width * (custom_height / height))
|
||||
new_height = custom_height
|
||||
elif force_size == "Custom":
|
||||
new_width, new_height = custom_width, custom_height
|
||||
else:
|
||||
target_width, target_height = map(int, force_size.replace("?", "0").split("x"))
|
||||
if target_width == 0:
|
||||
new_width = int(width * (target_height / height))
|
||||
new_height = target_height
|
||||
else:
|
||||
new_height = int(height * (target_width / width))
|
||||
new_width = target_width
|
||||
else:
|
||||
new_width, new_height = width, height
|
||||
|
||||
frames = []
|
||||
frame_count = 0
|
||||
|
||||
for i in range(total_frames):
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
break
|
||||
|
||||
if i < skip_first_frames:
|
||||
continue
|
||||
|
||||
if (i - skip_first_frames) % select_every_nth != 0:
|
||||
continue
|
||||
|
||||
if force_size != "Disabled":
|
||||
frame = cv2.resize(frame, (new_width, new_height))
|
||||
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frame = torch.from_numpy(frame).float() / 255.0
|
||||
frames.append(frame)
|
||||
|
||||
frame_count += 1
|
||||
|
||||
if frame_load_cap > 0 and frame_count >= frame_load_cap:
|
||||
break
|
||||
|
||||
cap.release()
|
||||
os.unlink(temp_file_path)
|
||||
|
||||
frames = torch.stack(frames)
|
||||
|
||||
video_info = {
|
||||
"source_fps": fps,
|
||||
"source_frame_count": total_frames,
|
||||
"source_duration": duration,
|
||||
"source_width": width,
|
||||
"source_height": height,
|
||||
"loaded_fps": fps if force_rate == 0 else force_rate,
|
||||
"loaded_frame_count": frame_count,
|
||||
"loaded_duration": frame_count / (fps if force_rate == 0 else force_rate),
|
||||
"loaded_width": new_width,
|
||||
"loaded_height": new_height,
|
||||
}
|
||||
|
||||
return (frames, frame_count, video_info)
|
||||
|
||||
@@ -0,0 +1,250 @@
|
||||
# 修改自 https://github.com/AnyaCoder/ComfyUI-fish-speech/
|
||||
|
||||
import torch,os
|
||||
from pathlib import Path
|
||||
from .fish_speech.llama_utils import load_model as load_llama_model
|
||||
from .fish_speech.vqgan_utils import load_model as load_vqgan_model
|
||||
from .fish_speech.vqgan_utils import audio2prompt, semantic2audio
|
||||
from .fish_speech.llama_utils import prompt2semantic
|
||||
|
||||
import folder_paths
|
||||
|
||||
def get_checkpoints_path():
|
||||
try:
|
||||
return folder_paths.get_folder_paths('fish_speech')[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, "fish_speech")
|
||||
|
||||
current_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
configs_dir=os.path.join(current_directory,"fish_speech","configs")
|
||||
|
||||
CKPTS_FOLDER = Path(get_checkpoints_path())
|
||||
|
||||
CONFIGS_FOLDER = Path(configs_dir)
|
||||
|
||||
|
||||
class LoadVQGAN:
|
||||
def __init__(self):
|
||||
self.vqgan = None
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"config": ([str(c.relative_to(CONFIGS_FOLDER)) for c in CONFIGS_FOLDER.glob("*vq*.yaml")], {"default": "firefly_gan_vq.yaml"}),
|
||||
"model": ([str(p.relative_to(CKPTS_FOLDER)) for p in CKPTS_FOLDER.glob("*vq*.pth")], ),
|
||||
"device": (["cuda", "cpu"], {"default": "cuda"}),
|
||||
},
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, model):
|
||||
return ""
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(s, model):
|
||||
return True
|
||||
|
||||
RETURN_TYPES = ("VQGAN", )
|
||||
RETURN_NAMES = ("vqgan", )
|
||||
|
||||
FUNCTION = "load_vqgan"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/FishSpeech"
|
||||
|
||||
def load_vqgan(self, config, model, device):
|
||||
config = config.rsplit(".", 1)[0]
|
||||
model = str(CKPTS_FOLDER / model)
|
||||
if self.vqgan is None:
|
||||
self.vqgan = load_vqgan_model(config,model, device=device)
|
||||
return (self.vqgan, )
|
||||
|
||||
|
||||
|
||||
class AudioToPrompt:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"vqgan": ("VQGAN", ),
|
||||
"audio": ("AUDIO", ),
|
||||
"device": (["cuda", "cpu"], {"default": "cuda"}),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("AUDIO", "NUMPY")
|
||||
RETURN_NAMES = ("restored_audio", "prompt_tokens")
|
||||
|
||||
FUNCTION = "encode"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/FishSpeech"
|
||||
|
||||
def encode(self, vqgan, audio, device):
|
||||
return audio2prompt(vqgan, audio, device)
|
||||
|
||||
|
||||
|
||||
class Prompt2Semantic:
|
||||
|
||||
def __init__(self):
|
||||
self.llama = None
|
||||
self.decode_func = None
|
||||
pass
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"multiline": True}),
|
||||
"prompt_text": ("STRING", {"multiline": True}),
|
||||
"prompt_tokens": ("NUMPY", ),
|
||||
"max_new_tokens": ("INT", {
|
||||
"default": 1024,
|
||||
"min": 0,
|
||||
"max": 2048,
|
||||
"step": 8,
|
||||
"display": "number",
|
||||
}),
|
||||
"top_p": ("FLOAT", {
|
||||
"default": 0.7,
|
||||
"min": 0.6,
|
||||
"max": 0.9,
|
||||
"step": 0.01,
|
||||
"display": "number",
|
||||
}),
|
||||
"repetition_penalty": ("FLOAT", {
|
||||
"default": 1.2,
|
||||
"min": 1.0,
|
||||
"max": 1.5,
|
||||
"step": 0.01,
|
||||
"display": "number",
|
||||
}),
|
||||
"temperature": ("FLOAT", {
|
||||
"default": 0.7,
|
||||
"min": 0.6,
|
||||
"max": 0.9,
|
||||
"step": 0.01,
|
||||
"display": "number",
|
||||
}),
|
||||
|
||||
"seed": ("INT", {
|
||||
"default": 42,
|
||||
"min": 0,
|
||||
"max": 4294967295,
|
||||
"step": 1,
|
||||
"display": "number",
|
||||
}),
|
||||
"iterative_prompt": (["yes", "no"], {"default": "yes"}),
|
||||
"chunk_length": ("INT", {
|
||||
"default": 100,
|
||||
"min": 0,
|
||||
"max": 500,
|
||||
"step": 8,
|
||||
"display": "number",
|
||||
}),
|
||||
|
||||
"compile": (["yes", "no"], {"default": "no"}),
|
||||
"precision": (["bf16", "half"], {"default": "bf16"}),
|
||||
|
||||
# "decode_func": ("DECODE_FUNC", ),
|
||||
"device": (["cuda", "cpu"], {"default": "cuda"}),
|
||||
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("NUMPY", )
|
||||
RETURN_NAMES = ("codes", )
|
||||
|
||||
FUNCTION = "decode"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/FishSpeech"
|
||||
|
||||
def decode(
|
||||
self,
|
||||
|
||||
text: str,
|
||||
prompt_text: str,
|
||||
prompt_tokens,
|
||||
max_new_tokens: int,
|
||||
top_p: float,
|
||||
repetition_penalty: float,
|
||||
temperature: float,
|
||||
|
||||
seed: int,
|
||||
iterative_prompt: str,
|
||||
chunk_length: int,
|
||||
|
||||
compile: str,
|
||||
precision,
|
||||
device: str,
|
||||
):
|
||||
|
||||
model = get_checkpoints_path()
|
||||
precision = torch.bfloat16 if precision == "bf16" else torch.half
|
||||
compile=True if compile == "yes" else False
|
||||
if self.llama is None or self.decode_func is None:
|
||||
self.llama, self.decode_func = load_llama_model(model, device, precision, compile)
|
||||
|
||||
|
||||
return prompt2semantic(
|
||||
self.llama,
|
||||
self.decode_func,
|
||||
text,
|
||||
[prompt_text,],
|
||||
[prompt_tokens,],
|
||||
max_new_tokens,
|
||||
top_p,
|
||||
repetition_penalty,
|
||||
temperature,
|
||||
device,
|
||||
compile=True if compile == "yes" else False,
|
||||
seed=seed,
|
||||
iterative_prompt=True if iterative_prompt == "yes" else False,
|
||||
chunk_length=chunk_length,
|
||||
)
|
||||
|
||||
|
||||
|
||||
class Semantic2Audio:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"vqgan": ("VQGAN", ),
|
||||
"codes": ("NUMPY", ),
|
||||
"device": (["cuda", "cpu"], {"default": "cuda"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("AUDIO", )
|
||||
RETURN_NAMES = ("generated_audio", )
|
||||
|
||||
FUNCTION = "generate"
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/FishSpeech"
|
||||
|
||||
def generate(self, vqgan, codes, device):
|
||||
return semantic2audio(vqgan, codes, device)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,123 @@
|
||||
import os,sys
|
||||
import folder_paths
|
||||
|
||||
from PIL import Image
|
||||
import importlib.util
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
|
||||
global _available
|
||||
_available=False
|
||||
|
||||
def is_installed(package):
|
||||
try:
|
||||
spec = importlib.util.find_spec(package)
|
||||
except ModuleNotFoundError:
|
||||
return False
|
||||
return spec is not None
|
||||
|
||||
if is_installed('simple_lama_inpainting')==False:
|
||||
import subprocess
|
||||
from packaging import version
|
||||
|
||||
if version.parse(torch.__version__)>=version.parse('2.1'):
|
||||
# 安装
|
||||
print('#pip install simple_lama_inpainting')
|
||||
|
||||
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'simple_lama_inpainting'], capture_output=True, text=True)
|
||||
|
||||
#检查命令执行结果
|
||||
if result.returncode == 0:
|
||||
print("#install success")
|
||||
from simple_lama_inpainting import SimpleLama
|
||||
_available=True
|
||||
else:
|
||||
print("#install error")
|
||||
else:
|
||||
print('#pls check your torch version >= 2.1')
|
||||
|
||||
else:
|
||||
from simple_lama_inpainting import SimpleLama
|
||||
_available=True
|
||||
|
||||
|
||||
def get_lama_path():
|
||||
try:
|
||||
return folder_paths.get_folder_paths('lama')[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, "lama")
|
||||
|
||||
llma_model_path=os.path.join(get_lama_path(), "big-lama.pt")
|
||||
if not os.path.exists(llma_model_path):
|
||||
os.environ['LAMA_MODEL']=''
|
||||
print(f"## lama torchscript model not found: {llma_model_path},pls download from https://github.com/enesmsahin/simple-lama-inpainting/releases/download/v0.1.0/big-lama.pt")
|
||||
else:
|
||||
os.environ['LAMA_MODEL'] = llma_model_path
|
||||
|
||||
# Tensor to PIL
|
||||
def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
# Convert PIL to Tensor
|
||||
def pil2tensor(image):
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
|
||||
# simple_lama = SimpleLama()
|
||||
|
||||
# img_path = "image.png"
|
||||
# mask_path = "mask.png"
|
||||
|
||||
# image = Image.open(img_path)
|
||||
# mask = Image.open(mask_path).convert('L')
|
||||
|
||||
# result = simple_lama(image, mask)
|
||||
# result.save("inpainted.png")
|
||||
|
||||
|
||||
class LaMaInpainting:
|
||||
global _available
|
||||
available=_available
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"image": ("IMAGE",),
|
||||
"mask": ("MASK",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Image"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
global simple_lama
|
||||
simple_lama = None
|
||||
def run(self,image,mask):
|
||||
global simple_lama
|
||||
|
||||
result=[]
|
||||
if simple_lama==None:
|
||||
simple_lama = SimpleLama()
|
||||
else:
|
||||
simple_lama.model.to("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
for i in range(len(image)):
|
||||
im=image[i]
|
||||
ma=mask[i]
|
||||
im=tensor2pil(im)
|
||||
ma=tensor2pil(ma)
|
||||
ma =ma.convert('L')
|
||||
|
||||
res = simple_lama(im, ma)
|
||||
res=pil2tensor(res)
|
||||
result.append(res)
|
||||
# result.save("inpainted.png")
|
||||
if simple_lama.device=='cuda':
|
||||
simple_lama.model.to('cpu')
|
||||
|
||||
return (result,)
|
||||
@@ -0,0 +1,274 @@
|
||||
|
||||
import scipy.ndimage
|
||||
import torch
|
||||
|
||||
import numpy as np
|
||||
# from PIL import Image, ImageDraw
|
||||
from PIL import Image, ImageOps
|
||||
|
||||
from comfy.cli_args import args
|
||||
import cv2,os
|
||||
from nodes import MAX_RESOLUTION, SaveImage, common_ksampler
|
||||
import folder_paths,random
|
||||
|
||||
# Tensor to PIL
|
||||
def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
# Convert PIL to Tensor
|
||||
def pil2tensor(image):
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
|
||||
def add_masks(mask1, mask2):
|
||||
mask1 = mask1.cpu()
|
||||
mask2 = mask2.cpu()
|
||||
cv2_mask1 = np.array(mask1) * 255
|
||||
cv2_mask2 = np.array(mask2) * 255
|
||||
|
||||
if cv2_mask1.shape == cv2_mask2.shape:
|
||||
cv2_mask = cv2.add(cv2_mask1, cv2_mask2)
|
||||
return torch.clamp(torch.from_numpy(cv2_mask) / 255.0, min=0, max=1)
|
||||
else:
|
||||
return mask1
|
||||
|
||||
|
||||
def grow(mask, expand, tapered_corners):
|
||||
c = 0 if tapered_corners else 1
|
||||
kernel = np.array([[c, 1, c],
|
||||
[1, 1, 1],
|
||||
[c, 1, c]])
|
||||
mask = mask.reshape((-1, mask.shape[-2], mask.shape[-1]))
|
||||
out = []
|
||||
for m in mask:
|
||||
output = m.numpy()
|
||||
for _ in range(abs(expand)):
|
||||
if expand < 0:
|
||||
output = scipy.ndimage.grey_erosion(output, footprint=kernel)
|
||||
else:
|
||||
output = scipy.ndimage.grey_dilation(output, footprint=kernel)
|
||||
output = torch.from_numpy(output)
|
||||
out.append(output)
|
||||
return torch.stack(out, dim=0)
|
||||
|
||||
def combine(destination, source, x, y):
|
||||
output = destination.reshape((-1, destination.shape[-2], destination.shape[-1])).clone()
|
||||
source = source.reshape((-1, source.shape[-2], source.shape[-1]))
|
||||
|
||||
left, top = (x, y,)
|
||||
right, bottom = (min(left + source.shape[-1], destination.shape[-1]), min(top + source.shape[-2], destination.shape[-2]))
|
||||
visible_width, visible_height = (right - left, bottom - top,)
|
||||
|
||||
source_portion = source[:, :visible_height, :visible_width]
|
||||
destination_portion = destination[:, top:bottom, left:right]
|
||||
|
||||
#operation == "subtract":
|
||||
output[:, top:bottom, left:right] = destination_portion - source_portion
|
||||
|
||||
output = torch.clamp(output, 0.0, 1.0)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class PreviewMask_(SaveImage):
|
||||
def __init__(self):
|
||||
self.output_dir = folder_paths.get_temp_directory()
|
||||
self.type = "temp"
|
||||
self.prefix_append =''.join(random.choice("abcdehijklmnopqrstupvxyzfg") for x in range(5))
|
||||
self.compress_level = 4
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"mask": ("MASK",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/Mask"
|
||||
|
||||
# 运行的函数
|
||||
def run(self, mask ):
|
||||
img=tensor2pil(mask)
|
||||
img=img.convert('RGB')
|
||||
img=pil2tensor(img)
|
||||
return self.save_images(img, 'temp_', None, None)
|
||||
|
||||
|
||||
class OutlineMask:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"mask": ("MASK",),
|
||||
"outline_width":("INT", {"default": 10,"min": 1, "max": MAX_RESOLUTION, "step": 1}),
|
||||
"tapered_corners": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('MASK',)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Mask"
|
||||
|
||||
# 运行的函数
|
||||
def run(self, mask, outline_width, tapered_corners):
|
||||
|
||||
m1=grow(mask,outline_width,tapered_corners)
|
||||
m2=grow(mask,-outline_width,tapered_corners)
|
||||
|
||||
m3=combine(m1,m2,0,0)
|
||||
|
||||
return (m3,)
|
||||
|
||||
|
||||
class MaskListReplace:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"masks": ("MASK",),
|
||||
"mask_replace": ("MASK",),
|
||||
"start_index":("INT", {"default": 0, "min": 0, "step": 1}),
|
||||
"end_index":("INT", {"default": 0, "min": 0, "step": 1}),
|
||||
"invert": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MASK",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/Video"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
def run(self, masks,mask_replace,start_index,end_index,invert):
|
||||
mask_replace=mask_replace[0]
|
||||
start_index=start_index[0]
|
||||
end_index=end_index[0]
|
||||
invert=invert[0]
|
||||
|
||||
new_masks=[]
|
||||
for i in range(len(masks)):
|
||||
if i>=start_index and i<=end_index:
|
||||
if invert:
|
||||
new_masks.append(masks[i])
|
||||
else:
|
||||
new_masks.append(mask_replace)
|
||||
else:
|
||||
if invert:
|
||||
new_masks.append(mask_replace)
|
||||
else:
|
||||
new_masks.append(masks[i])
|
||||
|
||||
return (new_masks,)
|
||||
|
||||
|
||||
class MaskListMerge:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"masks": ("MASK",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MASK",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/Mask"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
def run(self, masks):
|
||||
mask=masks[0]
|
||||
if isinstance(masks, list):
|
||||
for m in masks:
|
||||
# print(m.shape)
|
||||
mask = add_masks(mask, m)
|
||||
return (mask,)
|
||||
|
||||
|
||||
class FeatheredMask:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"mask": ("MASK",),
|
||||
"start_offset":("INT", {"default": 1,
|
||||
"min": -150,
|
||||
"max": 150,
|
||||
"step": 1,
|
||||
"display": "slider"}),
|
||||
"feathering_weight":("FLOAT", {"default": 0.1,
|
||||
"min": 0.0,
|
||||
"max": 1,
|
||||
"step": 0.1,
|
||||
"display": "slider"})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ('MASK',)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Mask"
|
||||
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
# 运行的函数
|
||||
def run(self,mask,start_offset, feathering_weight):
|
||||
# print(mask.shape,mask.size())
|
||||
|
||||
num,_,_=mask.size()
|
||||
|
||||
masks=[]
|
||||
|
||||
for i in range(num):
|
||||
mm=mask[i]
|
||||
image=tensor2pil(mm)
|
||||
|
||||
# Open the image using PIL
|
||||
image = image.convert("L")
|
||||
if start_offset>0:
|
||||
image=ImageOps.invert(image)
|
||||
|
||||
# Convert the image to a numpy array
|
||||
image_np = np.array(image)
|
||||
|
||||
# Use Canny edge detection to get black contours
|
||||
edges = cv2.Canny(image_np, 30, 150)
|
||||
|
||||
for i in range(0,abs(start_offset)):
|
||||
# int(100*feathering_weight)
|
||||
a=int(abs(start_offset)*0.1*i)
|
||||
# Dilate the black contours to make them wider
|
||||
kernel = np.ones((a, a), np.uint8)
|
||||
|
||||
dilated_edges = cv2.dilate(edges, kernel, iterations=1)
|
||||
# dilated_edges = cv2.erode(edges, kernel, iterations=1)
|
||||
# Smooth the dilated edges using Gaussian blur
|
||||
smoothed_edges = cv2.GaussianBlur(dilated_edges, (5, 5), 0)
|
||||
|
||||
# Adjust the feathering weight
|
||||
feathering_weight = max(0, min(feathering_weight, 1))
|
||||
|
||||
# Blend the smoothed edges with the original image to achieve feathering effect
|
||||
image_np = cv2.addWeighted(image_np, 1, smoothed_edges, feathering_weight, feathering_weight)
|
||||
|
||||
# Convert the result back to PIL image
|
||||
result_image = Image.fromarray(np.uint8(image_np))
|
||||
result_image=result_image.convert("L")
|
||||
|
||||
if start_offset>0:
|
||||
result_image=ImageOps.invert(result_image)
|
||||
|
||||
result_image=result_image.convert("L")
|
||||
mt=pil2tensor(result_image)
|
||||
masks.append(mt)
|
||||
|
||||
# print( mt.size())
|
||||
return (masks,)
|
||||
@@ -0,0 +1,157 @@
|
||||
# Referenced some code:https://github.com/IuvenisSapiens/ComfyUI_MiniCPM-V-2_6-int4
|
||||
# https://github.com/CY-CHENYUE/ComfyUI-MiniCPM-Plus
|
||||
|
||||
import os
|
||||
import torch
|
||||
import folder_paths
|
||||
from transformers import AutoTokenizer, AutoModel
|
||||
from torchvision.transforms.v2 import ToPILImage
|
||||
# from decord import VideoReader, cpu # pip install decord
|
||||
# from PIL import Image
|
||||
|
||||
def get_model_path(n=""):
|
||||
try:
|
||||
return folder_paths.get_folder_paths(n)[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, n)
|
||||
|
||||
|
||||
class MiniCPM_VQA_Simple:
|
||||
def __init__(self):
|
||||
self.model_checkpoint = None
|
||||
self.tokenizer = None
|
||||
self.model = None
|
||||
self.device = (
|
||||
torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
|
||||
)
|
||||
self.bf16_support = (
|
||||
torch.cuda.is_available()
|
||||
and torch.cuda.get_device_capability(self.device)[0] >= 8
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"text": ("STRING", {"default": "", "multiline": True}),
|
||||
"seed": ("INT", {"default": -1}), # add seed parameter, default is -1
|
||||
"extract_keywords":("BOOLEAN", {"default": False}),
|
||||
"temperature": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.7,
|
||||
},
|
||||
),
|
||||
"keep_model_loaded": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING","STRING",)
|
||||
RETURN_NAMES = ("result","keywords",)
|
||||
|
||||
FUNCTION = "inference"
|
||||
CATEGORY = "♾️Mixlab/Image"
|
||||
|
||||
def inference(
|
||||
self,
|
||||
images,
|
||||
text,
|
||||
seed, # add seed parameter, default is -1
|
||||
extract_keywords,
|
||||
temperature,
|
||||
keep_model_loaded,
|
||||
):
|
||||
if seed != -1:
|
||||
torch.manual_seed(seed)
|
||||
model_id = "openbmb/MiniCPM-V-2_6-int4"
|
||||
|
||||
self.model_checkpoint = os.path.join( get_model_path("prompt_generator"), os.path.basename(model_id))
|
||||
|
||||
if not os.path.exists(self.model_checkpoint):
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
snapshot_download(
|
||||
repo_id=model_id,
|
||||
local_dir=self.model_checkpoint,
|
||||
local_dir_use_symlinks=False,
|
||||
endpoint='https://hf-mirror.com'
|
||||
)
|
||||
|
||||
if self.tokenizer is None:
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(
|
||||
self.model_checkpoint,
|
||||
trust_remote_code=True,
|
||||
low_cpu_mem_usage=True,
|
||||
)
|
||||
|
||||
if self.model is None:
|
||||
self.model = AutoModel.from_pretrained(
|
||||
self.model_checkpoint,
|
||||
trust_remote_code=True,
|
||||
low_cpu_mem_usage=True,
|
||||
attn_implementation="sdpa",
|
||||
torch_dtype=torch.bfloat16 if self.bf16_support else torch.float16,
|
||||
)
|
||||
|
||||
|
||||
with torch.no_grad():
|
||||
images = images.permute([0, 3, 1, 2])
|
||||
images = [ToPILImage()(img).convert("RGB") for img in images]
|
||||
msgs = [{"role": "user", "content": images + [text]}]
|
||||
|
||||
params = {"use_image_id": False, }
|
||||
|
||||
# offload model to CPU
|
||||
# self.model = self.model.to(torch.device("cpu"))
|
||||
# self.model.eval()
|
||||
|
||||
result = self.model.chat(
|
||||
image=None,
|
||||
msgs=msgs,
|
||||
tokenizer=self.tokenizer,
|
||||
sampling=True,
|
||||
# top_k=top_k,
|
||||
# top_p=top_p,
|
||||
temperature=temperature,
|
||||
# repetition_penalty=repetition_penalty,
|
||||
# max_new_tokens=max_new_tokens,
|
||||
**params,
|
||||
)
|
||||
|
||||
keyword_result=""
|
||||
|
||||
if extract_keywords:#extract_keywords
|
||||
keyword_prompt = f"""Please extract keywords from the following text, including all occurrences of language (e.g. Chinese, English, etc.):
|
||||
[[[{result}]]]
|
||||
Please list the keywords extracted, separated by commas. Make sure to include all important words, no matter what language. For English words, please keep the original case."""
|
||||
|
||||
keyword_msgs = [{'role': 'user', 'content': keyword_prompt}]
|
||||
keyword_result = self.model.chat(
|
||||
image=None,
|
||||
msgs=keyword_msgs,
|
||||
tokenizer=self.tokenizer,
|
||||
sampling=True,
|
||||
# top_k=top_k,
|
||||
# top_p=top_p,
|
||||
temperature=temperature,
|
||||
# repetition_penalty=repetition_penalty,
|
||||
# max_new_tokens=max_new_tokens,
|
||||
**params,
|
||||
)
|
||||
print("keyword_result",keyword_result)
|
||||
|
||||
|
||||
# offload model to GPU
|
||||
# self.model = self.model.to(torch.device("cpu"))
|
||||
# self.model.eval()
|
||||
if not keep_model_loaded:
|
||||
del self.tokenizer # release tokenizer memory
|
||||
del self.model # release model memory
|
||||
self.tokenizer = None # set tokenizer to None
|
||||
self.model = None # set model to None
|
||||
torch.cuda.empty_cache() # release GPU memory
|
||||
torch.cuda.ipc_collect()
|
||||
# print(result)
|
||||
return (result,keyword_result,)
|
||||
@@ -0,0 +1,104 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image,ImageSequence,ImageOps
|
||||
import base64
|
||||
import io
|
||||
import comfy.utils
|
||||
import folder_paths
|
||||
import node_helpers
|
||||
|
||||
|
||||
# Tensor to PIL
|
||||
def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
# Convert PIL to Tensor
|
||||
def pil2tensor(image):
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
def load_image_to_tensor( image):
|
||||
image_path = folder_paths.get_annotated_filepath(image)
|
||||
|
||||
img = node_helpers.pillow(Image.open, image_path)
|
||||
|
||||
output_images = []
|
||||
output_masks = []
|
||||
w, h = None, None
|
||||
|
||||
excluded_formats = ['MPO']
|
||||
|
||||
for i in ImageSequence.Iterator(img):
|
||||
i = node_helpers.pillow(ImageOps.exif_transpose, i)
|
||||
|
||||
if i.mode == 'I':
|
||||
i = i.point(lambda i: i * (1 / 255))
|
||||
image = i.convert("RGB")
|
||||
|
||||
if len(output_images) == 0:
|
||||
w = image.size[0]
|
||||
h = image.size[1]
|
||||
|
||||
if image.size[0] != w or image.size[1] != h:
|
||||
continue
|
||||
|
||||
image = np.array(image).astype(np.float32) / 255.0
|
||||
image = torch.from_numpy(image)[None,]
|
||||
if 'A' in i.getbands():
|
||||
mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0
|
||||
mask = 1. - torch.from_numpy(mask)
|
||||
else:
|
||||
mask = torch.zeros((64,64), dtype=torch.float32, device="cpu")
|
||||
output_images.append(image)
|
||||
output_masks.append(mask.unsqueeze(0))
|
||||
|
||||
if len(output_images) > 1 and img.format not in excluded_formats:
|
||||
output_image = torch.cat(output_images, dim=0)
|
||||
output_mask = torch.cat(output_masks, dim=0)
|
||||
else:
|
||||
output_image = output_images[0]
|
||||
output_mask = output_masks[0]
|
||||
|
||||
return (output_image, output_mask)
|
||||
|
||||
|
||||
|
||||
class P5Input:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"frames":("IMAGEBASE64",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("frames",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Input"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
def run(self, frames):
|
||||
ims=[]
|
||||
for im in frames['images']:
|
||||
# print(im)
|
||||
if 'type' in im and (not f"[{im['type']}]" in im['name']):
|
||||
im['name']=im['name']+" "+f"[{im['type']}]"
|
||||
|
||||
output_image, output_mask = load_image_to_tensor(im['name'])
|
||||
ims.append(output_image)
|
||||
|
||||
if len(ims)==0:
|
||||
image1 = Image.new('RGB', (512, 512), color='black')
|
||||
return (pil2tensor(image1),)
|
||||
image1 = ims[0]
|
||||
for image2 in ims[1:]:
|
||||
if image1.shape[1:] != image2.shape[1:]:
|
||||
image2 = comfy.utils.common_upscale(image2.movedim(-1, 1), image1.shape[2], image1.shape[1], "bilinear", "center").movedim(1, -1)
|
||||
image1 = torch.cat((image1, image2), dim=0)
|
||||
|
||||
# 用于节点提示:p5节点提示有多少帧
|
||||
return {"ui": {"_info": [len(frames['images'])]}, "result": (image1,)}
|
||||
@@ -1,15 +1,97 @@
|
||||
import random
|
||||
import comfy.utils
|
||||
import json
|
||||
import os
|
||||
import numpy as np
|
||||
from urllib import request, parse
|
||||
import folder_paths
|
||||
from PIL import Image, ImageOps,ImageFilter,ImageEnhance,ImageDraw,ImageSequence, ImageFont
|
||||
from PIL.PngImagePlugin import PngInfo
|
||||
|
||||
import hashlib
|
||||
import requests
|
||||
import json
|
||||
|
||||
|
||||
# def queue_prompt(prompt_workflow):
|
||||
# p = {"prompt": prompt_workflow}
|
||||
# data = json.dumps(p).encode('utf-8')
|
||||
# req = request.Request("http://127.0.0.1:8188/prompt", data=data)
|
||||
# request.urlopen(req)
|
||||
|
||||
def get_model_path(n=""):
|
||||
try:
|
||||
return folder_paths.get_folder_paths(n)[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, n)
|
||||
|
||||
embeddings_path=get_model_path("embeddings")
|
||||
|
||||
def get_files_with_extension(directory, extension):
|
||||
|
||||
file_list = []
|
||||
for root, dirs, files in os.walk(directory):
|
||||
for file in files:
|
||||
if file.endswith(extension):
|
||||
file_name = os.path.splitext(file)[0]
|
||||
file_list.append(file_name)
|
||||
return file_list
|
||||
|
||||
def join_with_(text_list,delimiter):
|
||||
joined_text = delimiter.join(text_list)
|
||||
return joined_text
|
||||
|
||||
|
||||
|
||||
def queue_prompt(prompt_workflow):
|
||||
p = {"prompt": prompt_workflow}
|
||||
data = json.dumps(p).encode('utf-8')
|
||||
req = request.Request("http://127.0.0.1:8188/prompt", data=data)
|
||||
request.urlopen(req)
|
||||
def load_json(file_path):
|
||||
try:
|
||||
with open(file_path, 'r') as json_file:
|
||||
data = json.load(json_file)
|
||||
return data
|
||||
except FileNotFoundError:
|
||||
print(f"File not found: {file_path}")
|
||||
return None
|
||||
except json.JSONDecodeError:
|
||||
print(f"Error decoding JSON in file: {file_path}")
|
||||
return None
|
||||
|
||||
def save_json(data_dict, file_path):
|
||||
try:
|
||||
with open(file_path, 'w') as json_file:
|
||||
json.dump(data_dict, json_file, indent=4)
|
||||
print(f"Data saved to {file_path}")
|
||||
except Exception as e:
|
||||
print(f"Error saving JSON to file: {e}")
|
||||
|
||||
# pysss的lora加载器
|
||||
# def get_model_version_info(hash_value):
|
||||
# # http://127.0.0.1:1082
|
||||
# proxies = {'http': 'http://127.0.0.1:1082', 'https': 'https://127.0.0.1:1082'}
|
||||
# api_url = f"https://civitai.com/api/v1/model-versions/by-hash/{hash_value}"
|
||||
# print(api_url)
|
||||
# response = requests.get(api_url,proxies=proxies, verify=False)
|
||||
|
||||
# if response.status_code == 200:
|
||||
# return response.json()
|
||||
# else:
|
||||
# return None
|
||||
|
||||
# def calculate_sha256(file_path):
|
||||
# sha256_hash = hashlib.sha256()
|
||||
# with open(file_path, "rb") as f:
|
||||
# for chunk in iter(lambda: f.read(4096), b""):
|
||||
# sha256_hash.update(chunk)
|
||||
# return sha256_hash.hexdigest()
|
||||
|
||||
|
||||
|
||||
|
||||
class AnyType(str):
|
||||
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
|
||||
|
||||
def __ne__(self, __value: object) -> bool:
|
||||
return False
|
||||
|
||||
any_type = AnyType("*")
|
||||
|
||||
|
||||
default_prompt1='''Swing
|
||||
@@ -45,6 +127,252 @@ default_prompt1='''Swing
|
||||
'''
|
||||
default_prompt1="\n".join([p.strip() for p in default_prompt1.split('\n') if p.strip()!=''])
|
||||
|
||||
|
||||
def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
def addWeight(text, weight=1):
|
||||
if weight == 1:
|
||||
return text
|
||||
else:
|
||||
return f"({text}:{round(weight,3)})"
|
||||
|
||||
def prompt_delete_words(sentence, new_words_length):
|
||||
# 使用逗号分割句子,并去除空格
|
||||
words = [word.strip() for word in sentence.split(",")]
|
||||
|
||||
# 计算需要删除的单词数量
|
||||
num_to_delete = len(words) - new_words_length
|
||||
|
||||
words_to=[w for w in words]
|
||||
|
||||
# 逐个删除单词并存储在新列表中
|
||||
new_words = []
|
||||
for i in range(len(words)):
|
||||
if num_to_delete > 0:
|
||||
num_to_delete -= 1
|
||||
else:
|
||||
words_to.pop()
|
||||
if len(words_to)>0:
|
||||
new_words.append(", ".join(words_to))
|
||||
|
||||
return new_words
|
||||
|
||||
# # 测试方法
|
||||
# sentence = "a computer, a glass tablet with a keyboard on a dark background, 3d illustration, reflection, cgi 8k, clear glass, archaic, cut-away, white outline"
|
||||
# new_words_length = 5
|
||||
# result = prompt_delete_words(sentence, new_words_length)
|
||||
# print(result)
|
||||
|
||||
class PromptImage:
|
||||
def __init__(self):
|
||||
self.output_dir = folder_paths.get_output_directory()
|
||||
self.type = "output"
|
||||
self.prefix_append = "PromptImage"
|
||||
self.compress_level = 4
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"prompts": ("STRING",
|
||||
{
|
||||
"multiline": True,
|
||||
"default": '',
|
||||
"dynamicPrompts": False
|
||||
}),
|
||||
|
||||
"images": ("IMAGE",{"default": None}),
|
||||
"save_to_image": (["enable", "disable"],),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("json_str",)
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Output"
|
||||
|
||||
# 运行的函数
|
||||
def run(self,prompts,images,save_to_image):
|
||||
filename_prefix="mixlab_"
|
||||
filename_prefix += self.prefix_append
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(
|
||||
filename_prefix,self.output_dir, images[0].shape[1], images[0].shape[0])
|
||||
|
||||
full_output_folder=os.path.join(full_output_folder,'PromptImage')
|
||||
subfolder='PromptImage'
|
||||
|
||||
results = list()
|
||||
|
||||
save_to_image=save_to_image[0]=='enable'
|
||||
|
||||
#保存到本地的json文件,记录图片和prompt的对应关系
|
||||
output_images=[]
|
||||
output_prompt=[]
|
||||
|
||||
for index in range(len(images)):
|
||||
res=[]
|
||||
imgs=images[index]
|
||||
|
||||
for image in imgs:
|
||||
img=tensor2pil(image)
|
||||
|
||||
prompt_text=prompts[index]
|
||||
|
||||
metadata = None
|
||||
if save_to_image:
|
||||
metadata = PngInfo()
|
||||
if prompt_text is not None:
|
||||
metadata.add_text("prompt_text", prompt_text)
|
||||
|
||||
file = f"{filename}_{index}_{counter:05}_.png"
|
||||
fp=os.path.join(full_output_folder,file)
|
||||
img.save(fp, pnginfo=metadata, compress_level=self.compress_level)
|
||||
res.append({
|
||||
"filename": file,
|
||||
"subfolder": subfolder,
|
||||
"type": self.type
|
||||
})
|
||||
output_images.append(fp)
|
||||
output_prompt.append(prompt_text)
|
||||
counter += 1
|
||||
results.append(res)
|
||||
|
||||
# if save_to_image:
|
||||
# # 保存为本地文件
|
||||
# with open(os.path.join(full_output_folder,'PromptImage.json'), 'w') as file:
|
||||
# json.dump(output_dict, file, ensure_ascii=False, indent=4)
|
||||
|
||||
return { "ui": { "_images": results,"prompts":prompts },"result":(json.dumps({
|
||||
"images":output_images,
|
||||
"prompts":output_prompt
|
||||
}),) }
|
||||
|
||||
|
||||
|
||||
|
||||
class PromptSimplification:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING",
|
||||
{
|
||||
"multiline": True,
|
||||
"default": '',
|
||||
"dynamicPrompts": False
|
||||
}),
|
||||
|
||||
"length":("INT", {"default": 5, "min": 1,"max":100, "step": 1, "display": "number"}),
|
||||
|
||||
# "min_value":("FLOAT", {
|
||||
# "default": -2,
|
||||
# "min": -10,
|
||||
# "max": 0xffffffffffffffff,
|
||||
# "step": 0.01,
|
||||
# "display": "number"
|
||||
# }),
|
||||
# "max_value":("FLOAT", {
|
||||
# "default": 2,
|
||||
# "min": -10,
|
||||
# "max": 0xffffffffffffffff,
|
||||
# "step": 0.01,
|
||||
# "display": "number"
|
||||
# }),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("prompts",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
OUTPUT_NODE = True
|
||||
|
||||
# 运行的函数
|
||||
def run(self,prompt,length):
|
||||
length=length[0]
|
||||
result=[]
|
||||
for p in prompt:
|
||||
nps=prompt_delete_words(p,length)
|
||||
for n in nps:
|
||||
result.append(n)
|
||||
|
||||
result= [elem.strip() for elem in result if elem.strip()]
|
||||
|
||||
return {"ui": {"prompts": result}, "result": (result,)}
|
||||
|
||||
|
||||
|
||||
class PromptSlide:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
|
||||
"prompt_keyword": ("STRING",
|
||||
{
|
||||
"multiline": False,
|
||||
"default": '',
|
||||
"dynamicPrompts": False
|
||||
}),
|
||||
|
||||
"weight":("FLOAT", {"default": 1, "min": -3,"max": 3,
|
||||
"step": 0.01,
|
||||
"display": "slider"}),
|
||||
|
||||
# "min_value":("FLOAT", {
|
||||
# "default": -2,
|
||||
# "min": -10,
|
||||
# "max": 0xffffffffffffffff,
|
||||
# "step": 0.01,
|
||||
# "display": "number"
|
||||
# }),
|
||||
# "max_value":("FLOAT", {
|
||||
# "default": 2,
|
||||
# "min": -10,
|
||||
# "max": 0xffffffffffffffff,
|
||||
# "step": 0.01,
|
||||
# "display": "number"
|
||||
# }),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("prompt",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
OUTPUT_NODE = False
|
||||
|
||||
# 运行的函数
|
||||
def run(self,prompt_keyword,weight):
|
||||
# if weight < min_value:
|
||||
# weight= min_value
|
||||
# elif weight > max_value:
|
||||
# weight= max_value
|
||||
p=addWeight(prompt_keyword,weight)
|
||||
return (p,)
|
||||
|
||||
|
||||
|
||||
|
||||
class RandomPrompt:
|
||||
|
||||
'''
|
||||
@@ -70,6 +398,10 @@ class RandomPrompt:
|
||||
"default": 'sticker, Cartoon, ``'
|
||||
}),
|
||||
"random_sample": (["enable", "disable"],),
|
||||
# "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "step": 1}),
|
||||
},
|
||||
"optional":{
|
||||
"seed": (any_type, {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -79,15 +411,15 @@ class RandomPrompt:
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "Mixlab/prompt"
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
OUTPUT_NODE = True
|
||||
|
||||
|
||||
# 运行的函数
|
||||
def run(self,max_count,mutable_prompt,immutable_prompt,random_sample):
|
||||
print('#运行的函数',mutable_prompt,immutable_prompt,max_count,random_sample)
|
||||
def run(self,max_count,mutable_prompt,immutable_prompt,random_sample,seed=0):
|
||||
# print('#运行的函数',mutable_prompt,immutable_prompt,max_count,random_sample)
|
||||
|
||||
# Split the text into an array of words
|
||||
words1 = mutable_prompt.split("\n")
|
||||
@@ -106,6 +438,11 @@ class RandomPrompt:
|
||||
w1=w1.strip()
|
||||
for w2 in words2:
|
||||
w2=w2.strip()
|
||||
if '``' not in w2:
|
||||
if w2=="":
|
||||
w2='``'
|
||||
else:
|
||||
w2=w2+',``'
|
||||
if w1!='' and w2!='':
|
||||
prompts.append(w2.replace('``', w1))
|
||||
pbar.update(1)
|
||||
@@ -119,69 +456,258 @@ class RandomPrompt:
|
||||
else:
|
||||
prompts = prompts[:min(max_count,len(prompts))]
|
||||
|
||||
prompts= [elem.strip() for elem in prompts if elem.strip()]
|
||||
|
||||
# return (new_prompt)
|
||||
return {"ui": {"prompts": prompts}, "result": (prompts,)}
|
||||
|
||||
|
||||
# class LoraPrompt:
|
||||
# @classmethod
|
||||
# def INPUT_TYPES(s):
|
||||
# return {
|
||||
# "required": {
|
||||
# "lora_name":(sorted(folder_paths.get_filename_list("loras"), key=str.lower),),
|
||||
# "weight": ("FLOAT", {"default": 1, "min": -2, "max": 2,"step":0.01 ,"display": "slider"}),
|
||||
# "force_update": ("BOOLEAN", {"default": False}),
|
||||
# },
|
||||
|
||||
# }
|
||||
|
||||
# RETURN_TYPES = ("STRING","STRING",any_type)
|
||||
# RETURN_NAMES = ("lora_name","prompt","tags",)
|
||||
|
||||
# FUNCTION = "run"
|
||||
|
||||
# CATEGORY = "♾️Mixlab/Prompt"
|
||||
|
||||
# OUTPUT_IS_LIST = (False,False,True,)
|
||||
# # OUTPUT_NODE = True
|
||||
|
||||
# # 运行的函数
|
||||
# def run(self,lora_name,weight,force_update=False):
|
||||
|
||||
# # print('##LoraPrompt',__file__)
|
||||
# # 从本地数据库读取
|
||||
# json_tags_path = os.path.join(os.path.dirname(os.path.dirname(__file__)),r'data/loras_tags.json')
|
||||
|
||||
# if not os.path.exists(json_tags_path):
|
||||
# save_json({},json_tags_path)
|
||||
|
||||
# lora_tags = load_json(json_tags_path)
|
||||
# output_tags = lora_tags.get(lora_name, None) if lora_tags is not None else None
|
||||
# if output_tags is not None:
|
||||
# output_tags = ",".join(output_tags)
|
||||
# print("trainedWords:",output_tags)
|
||||
# else:
|
||||
# output_tags = ""
|
||||
|
||||
|
||||
# lora_path = folder_paths.get_full_path("loras", lora_name)
|
||||
# if output_tags == "" or force_update:
|
||||
# print("calculating lora hash")
|
||||
# LORAsha256 = calculate_sha256(lora_path)
|
||||
# print("requesting infos")
|
||||
# model_info = get_model_version_info(LORAsha256)
|
||||
# if model_info is not None:
|
||||
# if "trainedWords" in model_info:
|
||||
# print("tags found!")
|
||||
# if lora_tags is None:
|
||||
# lora_tags = {}
|
||||
# lora_tags[lora_name] = model_info["trainedWords"]
|
||||
# save_json(lora_tags,json_tags_path)
|
||||
# output_tags = ",".join(model_info["trainedWords"])
|
||||
# print("trainedWords:",output_tags)
|
||||
# else:
|
||||
# print("No informations found.")
|
||||
# if lora_tags is None:
|
||||
# lora_tags = {}
|
||||
# lora_tags[lora_name] = []
|
||||
# save_json(lora_tags,json_tags_path)
|
||||
|
||||
|
||||
# weight = round(weight, 3)
|
||||
# prompt=[]
|
||||
# for p in output_tags.split(','):
|
||||
|
||||
# if weight!=1:
|
||||
# prompt.append('('+p+':'+str(weight)+')')
|
||||
# else:
|
||||
# prompt.append(p)
|
||||
|
||||
# prompt=",".join(prompt)
|
||||
|
||||
# return (lora_name,prompt,output_tags.split(','),)
|
||||
|
||||
|
||||
|
||||
class RunWorkflow:
|
||||
class EmbeddingPrompt:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"workflow": ("STRING", {
|
||||
"multiline": False,
|
||||
"default": ''
|
||||
}),
|
||||
"prompt": ("STRING", {
|
||||
"multiline": False,
|
||||
"default": ''
|
||||
}),
|
||||
"image": ("IMAGE",),
|
||||
"input_node": ("STRING", {
|
||||
"multiline": False,
|
||||
"default": ''
|
||||
}),
|
||||
"output_node": ("STRING", {
|
||||
"multiline": False,
|
||||
"default": ''
|
||||
}),
|
||||
"embedding":(folder_paths.get_filename_list("embeddings"),),
|
||||
"weight": ("FLOAT", {"default": 1, "min": -2, "max": 2,"step":0.01 ,"display": "slider"}),
|
||||
},
|
||||
|
||||
}
|
||||
|
||||
|
||||
|
||||
RETURN_TYPES = ("IMAGE","STRING",)
|
||||
RETURN_TYPES = ("STRING",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "Mixlab/workflow"
|
||||
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
# OUTPUT_NODE = True
|
||||
|
||||
# 运行的函数
|
||||
def run(self,workflow,prompt,image,input_node,output_node):
|
||||
print('#运行的函数',prompt,image,input_node,output_node)
|
||||
workflow=json.loads(workflow)
|
||||
input_node=input_node.split(".")
|
||||
workflow[input_node[0]][input_node[1]][input_node[2]]=prompt
|
||||
|
||||
workflow_new={}
|
||||
# 遍历,seed设为随机
|
||||
for key, value in workflow.items():
|
||||
if 'inputs' in value:
|
||||
if 'seed' in value['inputs']:
|
||||
value['inputs']['seed']= random.randint(1, 18446744073709551614)
|
||||
workflow_new[key]=value
|
||||
|
||||
queue_prompt(workflow_new)
|
||||
print('#运行的函数',workflow_new[input_node[0]])
|
||||
|
||||
def run(self,embedding,weight):
|
||||
weight = round(weight, 3)
|
||||
prompt='embedding:'+embedding
|
||||
if weight!=1:
|
||||
prompt='('+prompt+':'+str(weight)+')'
|
||||
prompt=" "+prompt+' '
|
||||
# return (new_prompt)
|
||||
return {"ui":{"images": []},"result": ([image],['text'],)}
|
||||
return (prompt,)
|
||||
|
||||
# RETURN_TYPES = (any_type,)
|
||||
|
||||
# conditioning :提示,正向or负向
|
||||
# clip:clip模型
|
||||
# gligen_textbox_model:gligen模型
|
||||
# grids:矩形框的集合
|
||||
# labels:每个矩形框对应的标签的集合
|
||||
# index:选取第几个矩形框作为gligen的box
|
||||
|
||||
class GLIGENTextBoxApply_Advanced:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {"conditioning": ("CONDITIONING", ),
|
||||
"clip": ("CLIP", ),
|
||||
"gligen_textbox_model": ("GLIGEN", ),
|
||||
"grids": ("_GRID",),
|
||||
"labels": ("STRING",
|
||||
{
|
||||
"multiline": True,
|
||||
"default": "",
|
||||
"forceInput": True
|
||||
}),
|
||||
"index": ("INT", {"default": -1, "min": -1, "max": 300, "step": 1}),
|
||||
"max_size": ("INT", {"default": 8, "min": 1, "max": 300, "step": 1}),
|
||||
"random_shuffle":(["on","off"],),
|
||||
},
|
||||
"optional":{
|
||||
"seed": (any_type, {"default": 0, "min": 0, "max": 0xffffffffffffffff,"step": 1}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("CONDITIONING","STRING",)
|
||||
RETURN_NAMES = ("CONDITIONING","label",)
|
||||
|
||||
FUNCTION = "run"
|
||||
# INPUT_IS_LIST = True
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
|
||||
def run(self, conditioning, clip, gligen_textbox_model, grids, labels, index,max_size,random_shuffle,seed=0):
|
||||
# print('grids',grids)
|
||||
# conditioning=conditioning[0]
|
||||
# clip=clip[0]
|
||||
# gligen_textbox_model=gligen_textbox_model[0]
|
||||
# index=index[0]
|
||||
# max_size=max_size[0]
|
||||
# random_shuffle=random_shuffle[0]
|
||||
|
||||
texts=labels
|
||||
|
||||
if index>-1:
|
||||
texts=[labels[index]]
|
||||
grids=[grids[index]]
|
||||
|
||||
if random_shuffle=='on':
|
||||
sss=[[texts[i],grids[i]] for i in range(len(texts))]
|
||||
random.shuffle(sss)
|
||||
texts=[s[0] for s in sss]
|
||||
grids=[s[1] for s in sss]
|
||||
|
||||
if len(texts) > max_size:
|
||||
texts = texts[:max_size]
|
||||
|
||||
c = []
|
||||
|
||||
for t in conditioning:
|
||||
n = [t[0], t[1].copy()]
|
||||
|
||||
|
||||
# 多个
|
||||
position_params=[]
|
||||
for i in range(len(texts)):
|
||||
text=texts[i]
|
||||
grid=grids[i]
|
||||
x,y,width,height=grid
|
||||
# print(text)
|
||||
cond, cond_pooled = clip.encode_from_tokens(clip.tokenize(text), return_pooled=True)
|
||||
position_params =position_params+ [(cond_pooled, height // 8, width // 8, y // 8, x // 8)]
|
||||
|
||||
# 前一个
|
||||
prev = []
|
||||
if "gligen" in n[1]:
|
||||
prev = n[1]['gligen'][2]
|
||||
|
||||
n[1]['gligen'] = ("position", gligen_textbox_model, prev + position_params)
|
||||
# print('gligen',n)
|
||||
c.append(n)
|
||||
|
||||
# 下面这个写法有bug
|
||||
# for i in range(len(texts)):
|
||||
# text=texts[i]
|
||||
# grid=grids[i]
|
||||
# x,y,width,height=grid
|
||||
|
||||
# cond, cond_pooled = clip.encode_from_tokens(clip.tokenize(text), return_pooled=True)
|
||||
# for t in conditioning:
|
||||
# n = [t[0], t[1].copy()]
|
||||
# position_params = [(cond_pooled, height // 8, width // 8, y // 8, x // 8)]
|
||||
# prev = []
|
||||
# if "gligen" in n[1]:
|
||||
# prev = n[1]['gligen'][2]
|
||||
|
||||
# n[1]['gligen'] = ("position", gligen_textbox_model, prev + position_params)
|
||||
# c.append(n)
|
||||
|
||||
|
||||
return (c,texts, )
|
||||
|
||||
|
||||
class JoinWithDelimiter:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"text_list": (any_type,),
|
||||
"delimiter":(["newline","comma","backslash","space"],),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Text"
|
||||
|
||||
INPUT_IS_LIST = True # 当true的时候,输入时list,当false的时候,如果输入是list,则会自动包一层for循环调用
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
def run(self,text_list,delimiter):
|
||||
delimiter=delimiter[0]
|
||||
if delimiter =='newline':
|
||||
delimiter='\n'
|
||||
elif delimiter=='comma':
|
||||
delimiter=','
|
||||
elif delimiter=='backslash':
|
||||
delimiter='\\'
|
||||
elif delimiter=='space':
|
||||
delimiter=' '
|
||||
t=''
|
||||
if isinstance(text_list, list):
|
||||
t=join_with_(text_list,delimiter)
|
||||
return (t,)
|
||||
@@ -0,0 +1,708 @@
|
||||
import os,sys
|
||||
import folder_paths
|
||||
|
||||
from PIL import Image
|
||||
import importlib.util
|
||||
|
||||
import comfy.utils
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from huggingface_hub import hf_hub_download
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torchvision.transforms.functional import normalize
|
||||
# BRIA-RMBG-1.4 / briarmbg.py
|
||||
class REBNCONV(nn.Module):
|
||||
def __init__(self,in_ch=3,out_ch=3,dirate=1,stride=1):
|
||||
super(REBNCONV,self).__init__()
|
||||
|
||||
self.conv_s1 = nn.Conv2d(in_ch,out_ch,3,padding=1*dirate,dilation=1*dirate,stride=stride)
|
||||
self.bn_s1 = nn.BatchNorm2d(out_ch)
|
||||
self.relu_s1 = nn.ReLU(inplace=True)
|
||||
|
||||
def forward(self,x):
|
||||
|
||||
hx = x
|
||||
xout = self.relu_s1(self.bn_s1(self.conv_s1(hx)))
|
||||
|
||||
return xout
|
||||
|
||||
## upsample tensor 'src' to have the same spatial size with tensor 'tar'
|
||||
def _upsample_like(src,tar):
|
||||
|
||||
src = F.interpolate(src,size=tar.shape[2:],mode='bilinear')
|
||||
|
||||
return src
|
||||
|
||||
|
||||
### RSU-7 ###
|
||||
class RSU7(nn.Module):
|
||||
|
||||
def __init__(self, in_ch=3, mid_ch=12, out_ch=3, img_size=512):
|
||||
super(RSU7,self).__init__()
|
||||
|
||||
self.in_ch = in_ch
|
||||
self.mid_ch = mid_ch
|
||||
self.out_ch = out_ch
|
||||
|
||||
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1) ## 1 -> 1/2
|
||||
|
||||
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
|
||||
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool4 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool5 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv6 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
|
||||
self.rebnconv7 = REBNCONV(mid_ch,mid_ch,dirate=2)
|
||||
|
||||
self.rebnconv6d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv5d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
|
||||
|
||||
def forward(self,x):
|
||||
b, c, h, w = x.shape
|
||||
|
||||
hx = x
|
||||
hxin = self.rebnconvin(hx)
|
||||
|
||||
hx1 = self.rebnconv1(hxin)
|
||||
hx = self.pool1(hx1)
|
||||
|
||||
hx2 = self.rebnconv2(hx)
|
||||
hx = self.pool2(hx2)
|
||||
|
||||
hx3 = self.rebnconv3(hx)
|
||||
hx = self.pool3(hx3)
|
||||
|
||||
hx4 = self.rebnconv4(hx)
|
||||
hx = self.pool4(hx4)
|
||||
|
||||
hx5 = self.rebnconv5(hx)
|
||||
hx = self.pool5(hx5)
|
||||
|
||||
hx6 = self.rebnconv6(hx)
|
||||
|
||||
hx7 = self.rebnconv7(hx6)
|
||||
|
||||
hx6d = self.rebnconv6d(torch.cat((hx7,hx6),1))
|
||||
hx6dup = _upsample_like(hx6d,hx5)
|
||||
|
||||
hx5d = self.rebnconv5d(torch.cat((hx6dup,hx5),1))
|
||||
hx5dup = _upsample_like(hx5d,hx4)
|
||||
|
||||
hx4d = self.rebnconv4d(torch.cat((hx5dup,hx4),1))
|
||||
hx4dup = _upsample_like(hx4d,hx3)
|
||||
|
||||
hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))
|
||||
hx3dup = _upsample_like(hx3d,hx2)
|
||||
|
||||
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
|
||||
hx2dup = _upsample_like(hx2d,hx1)
|
||||
|
||||
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
|
||||
|
||||
return hx1d + hxin
|
||||
|
||||
|
||||
### RSU-6 ###
|
||||
class RSU6(nn.Module):
|
||||
|
||||
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
|
||||
super(RSU6,self).__init__()
|
||||
|
||||
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
|
||||
|
||||
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
|
||||
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool4 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
|
||||
self.rebnconv6 = REBNCONV(mid_ch,mid_ch,dirate=2)
|
||||
|
||||
self.rebnconv5d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
|
||||
|
||||
def forward(self,x):
|
||||
|
||||
hx = x
|
||||
|
||||
hxin = self.rebnconvin(hx)
|
||||
|
||||
hx1 = self.rebnconv1(hxin)
|
||||
hx = self.pool1(hx1)
|
||||
|
||||
hx2 = self.rebnconv2(hx)
|
||||
hx = self.pool2(hx2)
|
||||
|
||||
hx3 = self.rebnconv3(hx)
|
||||
hx = self.pool3(hx3)
|
||||
|
||||
hx4 = self.rebnconv4(hx)
|
||||
hx = self.pool4(hx4)
|
||||
|
||||
hx5 = self.rebnconv5(hx)
|
||||
|
||||
hx6 = self.rebnconv6(hx5)
|
||||
|
||||
|
||||
hx5d = self.rebnconv5d(torch.cat((hx6,hx5),1))
|
||||
hx5dup = _upsample_like(hx5d,hx4)
|
||||
|
||||
hx4d = self.rebnconv4d(torch.cat((hx5dup,hx4),1))
|
||||
hx4dup = _upsample_like(hx4d,hx3)
|
||||
|
||||
hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))
|
||||
hx3dup = _upsample_like(hx3d,hx2)
|
||||
|
||||
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
|
||||
hx2dup = _upsample_like(hx2d,hx1)
|
||||
|
||||
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
|
||||
|
||||
return hx1d + hxin
|
||||
|
||||
### RSU-5 ###
|
||||
class RSU5(nn.Module):
|
||||
|
||||
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
|
||||
super(RSU5,self).__init__()
|
||||
|
||||
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
|
||||
|
||||
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
|
||||
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
|
||||
self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=2)
|
||||
|
||||
self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
|
||||
|
||||
def forward(self,x):
|
||||
|
||||
hx = x
|
||||
|
||||
hxin = self.rebnconvin(hx)
|
||||
|
||||
hx1 = self.rebnconv1(hxin)
|
||||
hx = self.pool1(hx1)
|
||||
|
||||
hx2 = self.rebnconv2(hx)
|
||||
hx = self.pool2(hx2)
|
||||
|
||||
hx3 = self.rebnconv3(hx)
|
||||
hx = self.pool3(hx3)
|
||||
|
||||
hx4 = self.rebnconv4(hx)
|
||||
|
||||
hx5 = self.rebnconv5(hx4)
|
||||
|
||||
hx4d = self.rebnconv4d(torch.cat((hx5,hx4),1))
|
||||
hx4dup = _upsample_like(hx4d,hx3)
|
||||
|
||||
hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))
|
||||
hx3dup = _upsample_like(hx3d,hx2)
|
||||
|
||||
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
|
||||
hx2dup = _upsample_like(hx2d,hx1)
|
||||
|
||||
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
|
||||
|
||||
return hx1d + hxin
|
||||
|
||||
### RSU-4 ###
|
||||
class RSU4(nn.Module):
|
||||
|
||||
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
|
||||
super(RSU4,self).__init__()
|
||||
|
||||
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
|
||||
|
||||
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
|
||||
self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)
|
||||
|
||||
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=2)
|
||||
|
||||
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)
|
||||
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
|
||||
|
||||
def forward(self,x):
|
||||
|
||||
hx = x
|
||||
|
||||
hxin = self.rebnconvin(hx)
|
||||
|
||||
hx1 = self.rebnconv1(hxin)
|
||||
hx = self.pool1(hx1)
|
||||
|
||||
hx2 = self.rebnconv2(hx)
|
||||
hx = self.pool2(hx2)
|
||||
|
||||
hx3 = self.rebnconv3(hx)
|
||||
|
||||
hx4 = self.rebnconv4(hx3)
|
||||
|
||||
hx3d = self.rebnconv3d(torch.cat((hx4,hx3),1))
|
||||
hx3dup = _upsample_like(hx3d,hx2)
|
||||
|
||||
hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))
|
||||
hx2dup = _upsample_like(hx2d,hx1)
|
||||
|
||||
hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))
|
||||
|
||||
return hx1d + hxin
|
||||
|
||||
### RSU-4F ###
|
||||
class RSU4F(nn.Module):
|
||||
|
||||
def __init__(self, in_ch=3, mid_ch=12, out_ch=3):
|
||||
super(RSU4F,self).__init__()
|
||||
|
||||
self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)
|
||||
|
||||
self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)
|
||||
self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=2)
|
||||
self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=4)
|
||||
|
||||
self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=8)
|
||||
|
||||
self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=4)
|
||||
self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=2)
|
||||
self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)
|
||||
|
||||
def forward(self,x):
|
||||
|
||||
hx = x
|
||||
|
||||
hxin = self.rebnconvin(hx)
|
||||
|
||||
hx1 = self.rebnconv1(hxin)
|
||||
hx2 = self.rebnconv2(hx1)
|
||||
hx3 = self.rebnconv3(hx2)
|
||||
|
||||
hx4 = self.rebnconv4(hx3)
|
||||
|
||||
hx3d = self.rebnconv3d(torch.cat((hx4,hx3),1))
|
||||
hx2d = self.rebnconv2d(torch.cat((hx3d,hx2),1))
|
||||
hx1d = self.rebnconv1d(torch.cat((hx2d,hx1),1))
|
||||
|
||||
return hx1d + hxin
|
||||
|
||||
|
||||
class myrebnconv(nn.Module):
|
||||
def __init__(self, in_ch=3,
|
||||
out_ch=1,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1,
|
||||
dilation=1,
|
||||
groups=1):
|
||||
super(myrebnconv,self).__init__()
|
||||
|
||||
self.conv = nn.Conv2d(in_ch,
|
||||
out_ch,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
dilation=dilation,
|
||||
groups=groups)
|
||||
self.bn = nn.BatchNorm2d(out_ch)
|
||||
self.rl = nn.ReLU(inplace=True)
|
||||
|
||||
def forward(self,x):
|
||||
return self.rl(self.bn(self.conv(x)))
|
||||
|
||||
|
||||
class BriaRMBG(nn.Module):
|
||||
|
||||
def __init__(self,in_ch=3,out_ch=1):
|
||||
super(BriaRMBG,self).__init__()
|
||||
|
||||
self.conv_in = nn.Conv2d(in_ch,64,3,stride=2,padding=1)
|
||||
self.pool_in = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.stage1 = RSU7(64,32,64)
|
||||
self.pool12 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.stage2 = RSU6(64,32,128)
|
||||
self.pool23 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.stage3 = RSU5(128,64,256)
|
||||
self.pool34 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.stage4 = RSU4(256,128,512)
|
||||
self.pool45 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.stage5 = RSU4F(512,256,512)
|
||||
self.pool56 = nn.MaxPool2d(2,stride=2,ceil_mode=True)
|
||||
|
||||
self.stage6 = RSU4F(512,256,512)
|
||||
|
||||
# decoder
|
||||
self.stage5d = RSU4F(1024,256,512)
|
||||
self.stage4d = RSU4(1024,128,256)
|
||||
self.stage3d = RSU5(512,64,128)
|
||||
self.stage2d = RSU6(256,32,64)
|
||||
self.stage1d = RSU7(128,16,64)
|
||||
|
||||
self.side1 = nn.Conv2d(64,out_ch,3,padding=1)
|
||||
self.side2 = nn.Conv2d(64,out_ch,3,padding=1)
|
||||
self.side3 = nn.Conv2d(128,out_ch,3,padding=1)
|
||||
self.side4 = nn.Conv2d(256,out_ch,3,padding=1)
|
||||
self.side5 = nn.Conv2d(512,out_ch,3,padding=1)
|
||||
self.side6 = nn.Conv2d(512,out_ch,3,padding=1)
|
||||
|
||||
# self.outconv = nn.Conv2d(6*out_ch,out_ch,1)
|
||||
|
||||
def forward(self,x):
|
||||
|
||||
hx = x
|
||||
|
||||
hxin = self.conv_in(hx)
|
||||
#hx = self.pool_in(hxin)
|
||||
|
||||
#stage 1
|
||||
hx1 = self.stage1(hxin)
|
||||
hx = self.pool12(hx1)
|
||||
|
||||
#stage 2
|
||||
hx2 = self.stage2(hx)
|
||||
hx = self.pool23(hx2)
|
||||
|
||||
#stage 3
|
||||
hx3 = self.stage3(hx)
|
||||
hx = self.pool34(hx3)
|
||||
|
||||
#stage 4
|
||||
hx4 = self.stage4(hx)
|
||||
hx = self.pool45(hx4)
|
||||
|
||||
#stage 5
|
||||
hx5 = self.stage5(hx)
|
||||
hx = self.pool56(hx5)
|
||||
|
||||
#stage 6
|
||||
hx6 = self.stage6(hx)
|
||||
hx6up = _upsample_like(hx6,hx5)
|
||||
|
||||
#-------------------- decoder --------------------
|
||||
hx5d = self.stage5d(torch.cat((hx6up,hx5),1))
|
||||
hx5dup = _upsample_like(hx5d,hx4)
|
||||
|
||||
hx4d = self.stage4d(torch.cat((hx5dup,hx4),1))
|
||||
hx4dup = _upsample_like(hx4d,hx3)
|
||||
|
||||
hx3d = self.stage3d(torch.cat((hx4dup,hx3),1))
|
||||
hx3dup = _upsample_like(hx3d,hx2)
|
||||
|
||||
hx2d = self.stage2d(torch.cat((hx3dup,hx2),1))
|
||||
hx2dup = _upsample_like(hx2d,hx1)
|
||||
|
||||
hx1d = self.stage1d(torch.cat((hx2dup,hx1),1))
|
||||
|
||||
|
||||
#side output
|
||||
d1 = self.side1(hx1d)
|
||||
d1 = _upsample_like(d1,x)
|
||||
|
||||
d2 = self.side2(hx2d)
|
||||
d2 = _upsample_like(d2,x)
|
||||
|
||||
d3 = self.side3(hx3d)
|
||||
d3 = _upsample_like(d3,x)
|
||||
|
||||
d4 = self.side4(hx4d)
|
||||
d4 = _upsample_like(d4,x)
|
||||
|
||||
d5 = self.side5(hx5d)
|
||||
d5 = _upsample_like(d5,x)
|
||||
|
||||
d6 = self.side6(hx6)
|
||||
d6 = _upsample_like(d6,x)
|
||||
|
||||
return [F.sigmoid(d1), F.sigmoid(d2), F.sigmoid(d3), F.sigmoid(d4), F.sigmoid(d5), F.sigmoid(d6)],[hx1d,hx2d,hx3d,hx4d,hx5d,hx6]
|
||||
|
||||
|
||||
|
||||
|
||||
def get_U2NET_model_path():
|
||||
try:
|
||||
return folder_paths.get_folder_paths('rembg')[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, "rembg")
|
||||
|
||||
|
||||
U2NET_HOME=get_U2NET_model_path()
|
||||
os.environ["U2NET_HOME"] = U2NET_HOME
|
||||
|
||||
global _available
|
||||
_available=False
|
||||
|
||||
|
||||
def get_rembg_models(path):
|
||||
"""从目录中获取文件并提取文件名
|
||||
Args:
|
||||
path: 目录路径
|
||||
Returns:
|
||||
文件名列表
|
||||
"""
|
||||
filenames = []
|
||||
for root, _, files in os.walk(path):
|
||||
for filename in files:
|
||||
# 过滤隐藏文件
|
||||
if not filename.startswith('.'):
|
||||
name, ext = os.path.splitext(os.path.basename(filename))
|
||||
filenames.append(name)
|
||||
return filenames
|
||||
|
||||
|
||||
def is_installed(package):
|
||||
try:
|
||||
spec = importlib.util.find_spec(package)
|
||||
except ModuleNotFoundError:
|
||||
return False
|
||||
return spec is not None
|
||||
|
||||
|
||||
try:
|
||||
if is_installed('rembg')==False:
|
||||
import subprocess
|
||||
|
||||
# 安装
|
||||
print('#pip install rembg[gpu]')
|
||||
|
||||
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'rembg[gpu]'], capture_output=True, text=True)
|
||||
|
||||
#检查命令执行结果
|
||||
if result.returncode == 0:
|
||||
print("#install success")
|
||||
from rembg import new_session, remove
|
||||
_available=True
|
||||
else:
|
||||
print("#install error")
|
||||
|
||||
else:
|
||||
from rembg import new_session, remove
|
||||
_available=True
|
||||
|
||||
except:
|
||||
_available=False
|
||||
|
||||
|
||||
def run_briarmbg(images=[]):
|
||||
mroot=U2NET_HOME
|
||||
m=os.path.join(mroot,'briarmbg.pth')
|
||||
if os.path.exists(m)==False:
|
||||
# 下载
|
||||
m1=hf_hub_download("briaai/RMBG-1.4",
|
||||
local_dir=mroot,
|
||||
filename='model.pth',
|
||||
local_dir_use_symlinks=False,
|
||||
endpoint='https://hf-mirror.com')
|
||||
os.rename(m1, m)
|
||||
|
||||
net=BriaRMBG()
|
||||
if torch.cuda.is_available():
|
||||
net.load_state_dict(torch.load(m))
|
||||
net=net.cuda()
|
||||
else:
|
||||
net.load_state_dict(torch.load(m,map_location="cpu"))
|
||||
net.eval()
|
||||
|
||||
masks=[]
|
||||
rgba_images=[]
|
||||
rgb_images=[]
|
||||
for orig_image in images:
|
||||
|
||||
w,h = orig_im_size = orig_image.size
|
||||
|
||||
image = orig_image.convert('RGB')
|
||||
model_input_size = (1024, 1024)
|
||||
image = image.resize(model_input_size, Image.BILINEAR)
|
||||
|
||||
im_np = np.array(image)
|
||||
im_tensor = torch.tensor(im_np, dtype=torch.float32).permute(2,0,1)
|
||||
im_tensor = torch.unsqueeze(im_tensor,0)
|
||||
im_tensor = torch.divide(im_tensor,255.0)
|
||||
im_tensor = normalize(im_tensor,[0.5,0.5,0.5],[1.0,1.0,1.0])
|
||||
if torch.cuda.is_available():
|
||||
im_tensor=im_tensor.cuda()
|
||||
|
||||
result=net(im_tensor)
|
||||
result = torch.squeeze(F.interpolate(result[0][0], size=(h,w), mode='bilinear') ,0)
|
||||
ma = torch.max(result)
|
||||
mi = torch.min(result)
|
||||
result = (result-mi)/(ma-mi)
|
||||
im_array = (result*255).cpu().data.numpy().astype(np.uint8)
|
||||
mask = Image.fromarray(np.squeeze(im_array))
|
||||
# mask.save('test.png')
|
||||
# mask=tensor2pil(result)
|
||||
mask=mask.convert('L')
|
||||
|
||||
masks.append(mask)
|
||||
|
||||
# rgba图
|
||||
image_rgba =orig_image.convert("RGBA")
|
||||
image_rgba.putalpha(mask)
|
||||
rgba_images.append(image_rgba)
|
||||
|
||||
#rgb
|
||||
rgb_image = Image.new("RGB", image_rgba.size, (0, 0, 0))
|
||||
rgb_image.paste(image_rgba, mask=image_rgba.split()[3])
|
||||
rgb_images.append(rgb_image)
|
||||
return (masks,rgba_images,rgb_images)
|
||||
|
||||
|
||||
def run_rembg(model_name= "unet",images=[],callback=None):
|
||||
# model_name = "unet" # "isnet-general-use"
|
||||
# print('#run_rembg',model_name)
|
||||
rembg_session = new_session(model_name)
|
||||
masks=[]
|
||||
rgba_images=[]
|
||||
rgb_images=[]
|
||||
# 进度条
|
||||
pbar=callback
|
||||
for img in images:
|
||||
# use the post_process_mask argument to post process the mask to get better results.
|
||||
mask = remove(img, session=rembg_session,only_mask=True,post_process_mask=True)
|
||||
# mask=mask.convert('L')
|
||||
# masks.append(mask)
|
||||
if model_name=="u2net_cloth_seg":
|
||||
width, original_height = mask.size
|
||||
num_slices = original_height // img.height
|
||||
for i in range(num_slices):
|
||||
top = i * img.height
|
||||
bottom = (i + 1) * img.height
|
||||
slice_image = mask.crop((0, top, width, bottom))
|
||||
slice_mask=slice_image.convert('L')
|
||||
masks.append(slice_mask)
|
||||
|
||||
# rgba图
|
||||
image_rgba = img.convert("RGBA")
|
||||
image_rgba.putalpha(slice_mask)
|
||||
rgba_images.append(image_rgba)
|
||||
|
||||
#rgb
|
||||
rgb_image = Image.new("RGB", image_rgba.size, (0, 0, 0))
|
||||
rgb_image.paste(image_rgba, mask=image_rgba.split()[3])
|
||||
rgb_images.append(rgb_image)
|
||||
|
||||
else:
|
||||
mask=mask.convert('L')
|
||||
# mask.save(output_path)
|
||||
masks.append(mask)
|
||||
|
||||
# rgba图
|
||||
image_rgba = img.convert("RGBA")
|
||||
image_rgba.putalpha(mask)
|
||||
rgba_images.append(image_rgba)
|
||||
|
||||
#rgb
|
||||
rgb_image = Image.new("RGB", image_rgba.size, (0, 0, 0))
|
||||
rgb_image.paste(image_rgba, mask=image_rgba.split()[3])
|
||||
rgb_images.append(rgb_image)
|
||||
|
||||
if pbar:
|
||||
pbar.update(1)
|
||||
return (masks,rgba_images,rgb_images)
|
||||
|
||||
|
||||
# Tensor to PIL
|
||||
def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
# Convert PIL to Tensor
|
||||
def pil2tensor(image):
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
|
||||
class RembgNode_:
|
||||
|
||||
global _available
|
||||
available=_available
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"image": ("IMAGE",),
|
||||
"model_name": (get_rembg_models(U2NET_HOME),),
|
||||
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MASK","IMAGE","RGBA",)
|
||||
RETURN_NAMES = ("masks","images","RGBAs")
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Mask"
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,True,True,)
|
||||
|
||||
def run(self,image,model_name):
|
||||
# 兼容list输入和batch输入
|
||||
|
||||
model_name=model_name[0]
|
||||
|
||||
images=[]
|
||||
|
||||
for ims in image:
|
||||
for im in ims:
|
||||
im=tensor2pil(im)
|
||||
images.append(im)
|
||||
|
||||
if model_name=='briarmbg':
|
||||
masks,rgba_images,rgb_images=run_briarmbg(images)
|
||||
else:
|
||||
masks,rgba_images,rgb_images=run_rembg(model_name,images, comfy.utils.ProgressBar(len(images) ))
|
||||
|
||||
masks=[pil2tensor(m) for m in masks]
|
||||
|
||||
rgba_images=[pil2tensor(m) for m in rgba_images]
|
||||
|
||||
rgb_images=[pil2tensor(m) for m in rgb_images]
|
||||
|
||||
return (masks,rgb_images,rgba_images,)
|
||||
@@ -79,33 +79,37 @@ class ScreenShareNode:
|
||||
def INPUT_TYPES(s):
|
||||
return { "required":{
|
||||
"image_base64": ("CHEESE",),
|
||||
"refresh_rate": ("INT", {"default": 500, "min": 0,"step": 50, "max": 0xffffffffffffffff}),
|
||||
},
|
||||
"optional":{
|
||||
"prompt": ("PROMPT",),
|
||||
"slide": ("SLIDE",),
|
||||
"seed": ("SEED",),
|
||||
|
||||
# "seed": ("INT", {"default": 1, "min": 0, "max": 0xffffffffffffffff}),
|
||||
} }
|
||||
|
||||
RETURN_TYPES = ('IMAGE','STRING')
|
||||
|
||||
RETURN_TYPES = ('IMAGE','STRING','FLOAT',"INT")
|
||||
RETURN_NAMES = ("current frame (image)","prompt","denoise (float)","seed (int)")
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "Mixlab/image"
|
||||
CATEGORY = "♾️Mixlab/Screen"
|
||||
|
||||
# INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (False,False,False)
|
||||
OUTPUT_IS_LIST = (False,False,False,False)
|
||||
|
||||
# 运行的函数
|
||||
def run(self,image_base64,prompt):
|
||||
def run(self,image_base64,refresh_rate ,prompt,slide,seed):
|
||||
im,mask=base64_save(image_base64)
|
||||
# print('##########prompt',prompt)
|
||||
return (im,prompt)
|
||||
|
||||
return {"ui":{"refresh_rate": [refresh_rate]},"result": (im,prompt,slide,seed,)}
|
||||
|
||||
|
||||
class FloatingVideo:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return { "required":{
|
||||
"images": ("IMAGE",)
|
||||
"image": ("IMAGE",)
|
||||
}, }
|
||||
|
||||
# RETURN_TYPES = ('IMAGE','MASK')
|
||||
@@ -114,22 +118,22 @@ class FloatingVideo:
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "Mixlab/image"
|
||||
CATEGORY = "♾️Mixlab/Screen"
|
||||
|
||||
# INPUT_IS_LIST = True
|
||||
# OUTPUT_IS_LIST = (False,False,)
|
||||
|
||||
# 运行的函数
|
||||
def run(self,images):
|
||||
def run(self,image):
|
||||
|
||||
results = list()
|
||||
|
||||
for image in images:
|
||||
image=tensor2pil(image)
|
||||
for im in image:
|
||||
im=tensor2pil(im)
|
||||
# image_base64 = base64.b64encode(image.tobytes())
|
||||
|
||||
buffered = BytesIO()
|
||||
image.save(buffered, format="JPEG")
|
||||
im.save(buffered, format="JPEG")
|
||||
image_base64 = base64.b64encode(buffered.getvalue()).decode("utf-8")
|
||||
|
||||
results.append(image_base64)
|
||||
@@ -137,3 +141,15 @@ class FloatingVideo:
|
||||
|
||||
return { "ui": { "images_": results } }
|
||||
|
||||
|
||||
|
||||
# class SildeNode:
|
||||
# CATEGORY = "quicknodes"
|
||||
# @classmethod
|
||||
# def INPUT_TYPES(s):
|
||||
# return { "required":{} }
|
||||
# RETURN_TYPES = ()
|
||||
# RETURN_NAMES = ()
|
||||
# FUNCTION = "func"
|
||||
# def func(self):
|
||||
# return ()
|
||||
@@ -0,0 +1,226 @@
|
||||
# -*- coding:utf-8 -*-
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
|
||||
from huggingface_hub import snapshot_download
|
||||
import torch,re
|
||||
from sensevoice.onnx.sense_voice_ort_session import SenseVoiceInferenceSession
|
||||
from sensevoice.utils.frontend import WavFrontend
|
||||
from sensevoice.utils.fsmn_vad import FSMNVad
|
||||
import comfy.utils
|
||||
import folder_paths
|
||||
|
||||
languages = {"auto": 0, "zh": 3, "en": 4, "yue": 7, "ja": 11, "ko": 12, "nospeech": 13}
|
||||
|
||||
# 设置环境变量
|
||||
os.environ['HF_ENDPOINT'] = 'https://hf-mirror.com'
|
||||
|
||||
#
|
||||
def get_model_path():
|
||||
try:
|
||||
return folder_paths.get_folder_paths('sense_voice')[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, "sense_voice")
|
||||
|
||||
class AnyType(str):
|
||||
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
|
||||
|
||||
def __ne__(self, __value: object) -> bool:
|
||||
return False
|
||||
|
||||
any_type = AnyType("*")
|
||||
|
||||
# 字幕
|
||||
def format_to_srt(channel_id, start_time_ms, end_time_ms, asr_result):
|
||||
start_time = start_time_ms / 1000
|
||||
end_time = end_time_ms / 1000
|
||||
|
||||
def format_time(seconds):
|
||||
hours = int(seconds // 3600)
|
||||
minutes = int((seconds % 3600) // 60)
|
||||
seconds = seconds % 60
|
||||
milliseconds = int((seconds - int(seconds)) * 1000)
|
||||
return f"{hours:02}:{minutes:02}:{int(seconds):02},{milliseconds:03}"
|
||||
|
||||
start_time_str = format_time(start_time)
|
||||
end_time_str = format_time(end_time)
|
||||
|
||||
pattern = r"<\|(.+?)\|><\|(.+?)\|><\|(.+?)\|><\|(.+?)\|>(.+)"
|
||||
match = re.match(pattern,asr_result)
|
||||
print('#format_to_srt',match,asr_result)
|
||||
if match==None:
|
||||
return None, None, None, None,None,start_time,end_time,None
|
||||
lang, emotion, audio_type, itn, text = match.groups()
|
||||
# 😊 表示高兴,😡 表示愤怒,😔 表示悲伤。对于音频事件,🎼 表示音乐,😀 表示笑声,👏 表示掌声
|
||||
|
||||
srt_content = f"1\n{start_time_str} --> {end_time_str}\n{text}\n"
|
||||
|
||||
logging.info(f"[Channel {channel_id}] [{start_time}s - {end_time}s] [{lang}] [{emotion}] [{audio_type}] [{itn}] {text}")
|
||||
|
||||
return lang, emotion, audio_type, itn,srt_content,start_time,end_time,text
|
||||
|
||||
|
||||
class SenseVoiceProcessor:
|
||||
def __init__(self, download_model_path, device, num_threads, use_int8):
|
||||
|
||||
if not os.path.exists(download_model_path):
|
||||
logging.info(
|
||||
"Downloading model from huggingface hub from https://huggingface.co/lovemefan/SenseVoice-onnx"
|
||||
)
|
||||
logging.info(
|
||||
"You can speed up with `export HF_ENDPOINT=https://hf-mirror.com`"
|
||||
)
|
||||
snapshot_download(
|
||||
repo_id="lovemefan/SenseVoice-onnx", local_dir=download_model_path
|
||||
)
|
||||
|
||||
self.download_model_path = download_model_path
|
||||
self.device = device
|
||||
self.num_threads = num_threads
|
||||
self.use_int8 = use_int8
|
||||
self.front = WavFrontend(os.path.join(download_model_path, "am.mvn"))
|
||||
self.model = SenseVoiceInferenceSession(
|
||||
os.path.join(download_model_path, "embedding.npy"),
|
||||
os.path.join(
|
||||
download_model_path,
|
||||
"sense-voice-encoder-int8.onnx"
|
||||
if use_int8
|
||||
else "sense-voice-encoder.onnx",
|
||||
),
|
||||
os.path.join(download_model_path, "chn_jpn_yue_eng_ko_spectok.bpe.model"),
|
||||
device,
|
||||
num_threads,
|
||||
)
|
||||
self.vad = FSMNVad(download_model_path)
|
||||
|
||||
def process_audio(self, waveform, _sample_rate, language, use_itn):
|
||||
|
||||
start = time.time()
|
||||
pbar = comfy.utils.ProgressBar(waveform.shape[1]) # 进度条
|
||||
|
||||
results = []
|
||||
|
||||
for channel_id, channel_data in enumerate(waveform.T):
|
||||
segments = self.vad.segments_offline(channel_data)
|
||||
|
||||
for part in segments:
|
||||
audio_feats = self.front.get_features(channel_data[part[0] * 16 : part[1] * 16])
|
||||
asr_result = self.model(
|
||||
audio_feats[None, ...],
|
||||
language=languages[language],
|
||||
use_itn=use_itn,
|
||||
)
|
||||
|
||||
lang, emotion, audio_type, itn,srt_content,start_time,end_time,text=format_to_srt(
|
||||
channel_id,
|
||||
part[0] ,
|
||||
part[1],
|
||||
asr_result)
|
||||
|
||||
if lang!=None:
|
||||
results.append({
|
||||
"language":lang,
|
||||
"emotion":emotion,
|
||||
"audio_type":audio_type,
|
||||
"itn":itn,
|
||||
"srt_content":srt_content,
|
||||
"start_time":start_time,
|
||||
"end_time":end_time,
|
||||
"text":text
|
||||
})
|
||||
|
||||
self.vad.vad.all_reset_detection()
|
||||
pbar.update(1) # 更新进度条
|
||||
|
||||
decoding_time = time.time() - start
|
||||
logging.info(f"Decoder audio takes {decoding_time} seconds")
|
||||
logging.info(f"The RTF is {decoding_time/(waveform.shape[1] * len(waveform) / _sample_rate)}.")
|
||||
return results
|
||||
|
||||
|
||||
class SenseVoiceNode:
|
||||
|
||||
def __init__(self):
|
||||
self.processor = None
|
||||
self.download_model_path=get_model_path()
|
||||
self.device="cpu"
|
||||
self.num_threads = 4
|
||||
self.use_int8 = True
|
||||
self.language='auto'
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
|
||||
return {"required": {
|
||||
"audio": ("AUDIO", ),
|
||||
"device": ( ['auto','cpu'], {"default": 'auto'}),
|
||||
"language": (list(languages.keys()), {"default": 'auto'}),# 不能直接写 languages.keys(),json.dumps会报错
|
||||
"num_threads":("INT",{
|
||||
"default":4,
|
||||
"min": 1, #Minimum value
|
||||
"max": 32, #Maximum value
|
||||
"step": 1, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
},),
|
||||
"use_int8":("BOOLEAN", {"default": True},),
|
||||
"use_itn":("BOOLEAN", {"default": True},),
|
||||
},
|
||||
}
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
|
||||
RETURN_TYPES = (any_type,"STRING","STRING","FLOAT",)
|
||||
RETURN_NAMES = ("result","srt","text","total_seconds",)
|
||||
|
||||
def run(self,audio,device,language,num_threads,use_int8,use_itn ):
|
||||
|
||||
if device!=self.device:
|
||||
self.device=device
|
||||
self.processor=None
|
||||
if language!=self.language:
|
||||
self.language=language
|
||||
self.processor=None
|
||||
if num_threads!=self.num_threads:
|
||||
self.num_threads=num_threads
|
||||
self.processor=None
|
||||
if use_int8!=self.use_int8:
|
||||
self.use_int8=use_int8
|
||||
self.processor=None
|
||||
|
||||
if device=='auto' and torch.cuda.is_available():
|
||||
self.device='cuda'
|
||||
|
||||
# num_threads=4
|
||||
# use_int8=True
|
||||
|
||||
if self.processor==None:
|
||||
self.processor = SenseVoiceProcessor(self.download_model_path,
|
||||
self.device,
|
||||
self.num_threads,
|
||||
self.use_int8)
|
||||
|
||||
if 'waveform' in audio and 'sample_rate' in audio:
|
||||
waveform = audio['waveform']
|
||||
sample_rate = audio['sample_rate']
|
||||
# print("Original shape:", waveform.shape) # 打印原始形状
|
||||
if waveform.ndim == 3 and waveform.shape[0] == 1: # 检查是否为三维且 batch_size 为 1
|
||||
waveform = waveform.squeeze(0) # 移除 batch_size 维度
|
||||
else:
|
||||
raise ValueError("Unexpected waveform dimensions")
|
||||
|
||||
print("waveform.shape:", waveform.shape)
|
||||
total_length_seconds = waveform.shape[1] / sample_rate
|
||||
|
||||
waveform_numpy = waveform.numpy().transpose(1, 0) # 转换为 (num_samples, num_channels)
|
||||
|
||||
results=self.processor.process_audio(waveform_numpy, sample_rate, language, use_itn)
|
||||
|
||||
srt_content="\n".join([s['srt_content'] for s in results])
|
||||
text="\n".join([s['text'] for s in results])
|
||||
|
||||
return (results,srt_content,text,total_length_seconds,)
|
||||
|
||||
@@ -0,0 +1,503 @@
|
||||
import comfy
|
||||
import torch
|
||||
|
||||
from dataclasses import dataclass
|
||||
import torch.nn as nn
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
import comfy.ops
|
||||
from typing import Union
|
||||
import comfy.sample
|
||||
import latent_preview
|
||||
import comfy.utils
|
||||
|
||||
T = torch.Tensor
|
||||
|
||||
|
||||
from .VisualStylePrompting.attention_functions import VisualStyleProcessor
|
||||
|
||||
class ApplyVisualStylePrompting:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"reference_image": ("IMAGE",),
|
||||
"reference_image_text": ("STRING", {"multiline": True}),
|
||||
"model": ("MODEL",),
|
||||
"clip": ("CLIP", ),
|
||||
"vae": ("VAE", ),
|
||||
"positive": ("CONDITIONING",),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"enabled": ("BOOLEAN", {"default": True}),
|
||||
"denoise": ("FLOAT", {"default": 1., "min": 0., "max": 1., "step": 1e-2}),
|
||||
"batch_size": ("INT", {"default": 1, "min": 1, "max": 4096,"step":2})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL", "CONDITIONING","CONDITIONING", "LATENT")
|
||||
RETURN_NAMES = ("model", "positive", "negative", "latents")
|
||||
|
||||
CATEGORY = "♾️Mixlab/Style"
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
def run(
|
||||
self,
|
||||
reference_image,
|
||||
reference_image_text,
|
||||
model: comfy.model_patcher.ModelPatcher,
|
||||
clip,
|
||||
vae,
|
||||
positive,
|
||||
negative,
|
||||
enabled,
|
||||
denoise,
|
||||
batch_size=1
|
||||
):
|
||||
|
||||
tokens = clip.tokenize(reference_image_text)
|
||||
cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True)
|
||||
reference_image_prompt=[[cond, {"pooled_output": pooled}]]
|
||||
|
||||
reference_image = reference_image.repeat(((batch_size+1)//2, 1,1,1))
|
||||
|
||||
self.model = model
|
||||
reference_latent = vae.encode(reference_image[:,:,:,:3])
|
||||
|
||||
for n, m in model.model.diffusion_model.named_modules():
|
||||
if m.__class__.__name__ == "CrossAttention":
|
||||
processor = VisualStyleProcessor(m, enabled=enabled)
|
||||
setattr(m, 'forward', processor.visual_style_forward)
|
||||
|
||||
conditioning_prompt = reference_image_prompt + positive
|
||||
negative_prompt = negative * 2
|
||||
|
||||
latents = torch.zeros_like(reference_latent)
|
||||
latents = torch.cat([latents] * 2)
|
||||
|
||||
if denoise < 1.0:
|
||||
latents[::1] = reference_latent[:1]
|
||||
else:
|
||||
latents[::2] = reference_latent
|
||||
|
||||
denoise_mask = torch.ones_like(latents)[:, :1, ...] * denoise
|
||||
|
||||
denoise_mask[0] = 0.
|
||||
|
||||
return (model, conditioning_prompt, negative_prompt, {"samples": latents, "noise_mask": denoise_mask})
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def exists(val):
|
||||
return val is not None
|
||||
|
||||
def default(val, d):
|
||||
if exists(val):
|
||||
return val
|
||||
return d
|
||||
|
||||
|
||||
class StyleAlignedArgs:
|
||||
def __init__(self, share_attn: str) -> None:
|
||||
self.adain_keys = "k" in share_attn
|
||||
self.adain_values = "v" in share_attn
|
||||
self.adain_queries = "q" in share_attn
|
||||
|
||||
share_attention: bool = True
|
||||
adain_queries: bool = True
|
||||
adain_keys: bool = True
|
||||
adain_values: bool = True
|
||||
|
||||
|
||||
def expand_first(
|
||||
feat: T,
|
||||
scale=1.0,
|
||||
) -> T:
|
||||
"""
|
||||
Expand the first element so it has the same shape as the rest of the batch.
|
||||
"""
|
||||
b = feat.shape[0]
|
||||
feat_style = torch.stack((feat[0], feat[b // 2])).unsqueeze(1)
|
||||
if scale == 1:
|
||||
feat_style = feat_style.expand(2, b // 2, *feat.shape[1:])
|
||||
else:
|
||||
feat_style = feat_style.repeat(1, b // 2, 1, 1, 1)
|
||||
feat_style = torch.cat([feat_style[:, :1], scale * feat_style[:, 1:]], dim=1)
|
||||
return feat_style.reshape(*feat.shape)
|
||||
|
||||
|
||||
def concat_first(feat: T, dim=2, scale=1.0) -> T:
|
||||
"""
|
||||
concat the the feature and the style feature expanded above
|
||||
"""
|
||||
feat_style = expand_first(feat, scale=scale)
|
||||
return torch.cat((feat, feat_style), dim=dim)
|
||||
|
||||
|
||||
def calc_mean_std(feat, eps: float = 1e-5) -> "tuple[T, T]":
|
||||
feat_std = (feat.var(dim=-2, keepdims=True) + eps).sqrt()
|
||||
feat_mean = feat.mean(dim=-2, keepdims=True)
|
||||
return feat_mean, feat_std
|
||||
|
||||
def adain(feat: T) -> T:
|
||||
feat_mean, feat_std = calc_mean_std(feat)
|
||||
feat_style_mean = expand_first(feat_mean)
|
||||
feat_style_std = expand_first(feat_std)
|
||||
feat = (feat - feat_mean) / feat_std
|
||||
feat = feat * feat_style_std + feat_style_mean
|
||||
return feat
|
||||
|
||||
class SharedAttentionProcessor:
|
||||
def __init__(self, args: StyleAlignedArgs, scale: float):
|
||||
self.args = args
|
||||
self.scale = scale
|
||||
|
||||
def __call__(self, q, k, v, extra_options):
|
||||
if self.args.adain_queries:
|
||||
q = adain(q)
|
||||
if self.args.adain_keys:
|
||||
k = adain(k)
|
||||
if self.args.adain_values:
|
||||
v = adain(v)
|
||||
if self.args.share_attention:
|
||||
k = concat_first(k, -2, scale=self.scale)
|
||||
v = concat_first(v, -2)
|
||||
|
||||
return q, k, v
|
||||
|
||||
|
||||
def get_norm_layers(
|
||||
layer: nn.Module,
|
||||
norm_layers_: "dict[str, list[Union[nn.GroupNorm, nn.LayerNorm]]]",
|
||||
share_layer_norm: bool,
|
||||
share_group_norm: bool,
|
||||
):
|
||||
if isinstance(layer, nn.LayerNorm) and share_layer_norm:
|
||||
norm_layers_["layer"].append(layer)
|
||||
if isinstance(layer, nn.GroupNorm) and share_group_norm:
|
||||
norm_layers_["group"].append(layer)
|
||||
else:
|
||||
for child_layer in layer.children():
|
||||
get_norm_layers(
|
||||
child_layer, norm_layers_, share_layer_norm, share_group_norm
|
||||
)
|
||||
|
||||
|
||||
def register_norm_forward(
|
||||
norm_layer: Union[nn.GroupNorm, nn.LayerNorm],
|
||||
) -> Union[nn.GroupNorm, nn.LayerNorm]:
|
||||
if not hasattr(norm_layer, "orig_forward"):
|
||||
setattr(norm_layer, "orig_forward", norm_layer.forward)
|
||||
orig_forward = norm_layer.orig_forward
|
||||
|
||||
def forward_(hidden_states: T) -> T:
|
||||
n = hidden_states.shape[-2]
|
||||
hidden_states = concat_first(hidden_states, dim=-2)
|
||||
hidden_states = orig_forward(hidden_states) # type: ignore
|
||||
return hidden_states[..., :n, :]
|
||||
|
||||
norm_layer.forward = forward_ # type: ignore
|
||||
return norm_layer
|
||||
|
||||
|
||||
def register_shared_norm(
|
||||
model: ModelPatcher,
|
||||
share_group_norm: bool = True,
|
||||
share_layer_norm: bool = True,
|
||||
):
|
||||
norm_layers = {"group": [], "layer": []}
|
||||
get_norm_layers(model.model, norm_layers, share_layer_norm, share_group_norm)
|
||||
print(
|
||||
f"Patching {len(norm_layers['group'])} group norms, {len(norm_layers['layer'])} layer norms."
|
||||
)
|
||||
return [register_norm_forward(layer) for layer in norm_layers["group"]] + [
|
||||
register_norm_forward(layer) for layer in norm_layers["layer"]
|
||||
]
|
||||
|
||||
|
||||
SHARE_NORM_OPTIONS = ["both", "group", "layer", "disabled"]
|
||||
SHARE_ATTN_OPTIONS = ["q+k", "q+k+v", "disabled"]
|
||||
|
||||
class StyleAlignedSampleReferenceLatents:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{
|
||||
"reference_image": ("IMAGE",),
|
||||
"positive": ("CONDITIONING",),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"model": ("MODEL",),
|
||||
"vae": ("VAE", ),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
|
||||
|
||||
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, ),
|
||||
"denoise": ("FLOAT", {"default": 1, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STEP_LATENTS","LATENT")
|
||||
RETURN_NAMES = ("ref_latents", "noised_output")
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
# CATEGORY = "style_aligned"
|
||||
CATEGORY = "♾️Mixlab/Style"
|
||||
|
||||
def run(self, reference_image, positive, negative, model, vae, seed, steps, cfg,scheduler,denoise):
|
||||
|
||||
# TODO noise_mask?
|
||||
def vae_encode_crop_pixels(pixels):
|
||||
x = (pixels.shape[1] // 8) * 8
|
||||
y = (pixels.shape[2] // 8) * 8
|
||||
if pixels.shape[1] != x or pixels.shape[2] != y:
|
||||
x_offset = (pixels.shape[1] % 8) // 2
|
||||
y_offset = (pixels.shape[2] % 8) // 2
|
||||
pixels = pixels[:, x_offset:x + x_offset, y_offset:y + y_offset, :]
|
||||
return pixels
|
||||
|
||||
pixels=vae_encode_crop_pixels(reference_image)
|
||||
t = vae.encode(pixels[:,:,:,:3])
|
||||
latent_image = {"samples":t}
|
||||
|
||||
noise_seed=seed
|
||||
|
||||
sampler_name="ddim"
|
||||
|
||||
sampler = comfy.samplers.sampler_object(sampler_name)
|
||||
|
||||
total_steps = steps
|
||||
if denoise < 1.0:
|
||||
total_steps = int(steps/denoise)
|
||||
|
||||
comfy.model_management.load_models_gpu([model])
|
||||
sigmas = comfy.samplers.calculate_sigmas_scheduler(model.model, scheduler, total_steps).cpu()
|
||||
sigmas = sigmas[-(steps + 1):]
|
||||
|
||||
sigmas = sigmas.flip(0)
|
||||
if sigmas[0] == 0:
|
||||
sigmas[0] = 0.0001
|
||||
|
||||
latent = latent_image
|
||||
latent_image = latent["samples"]
|
||||
noise = torch.zeros(latent_image.size(), dtype=latent_image.dtype, layout=latent_image.layout, device="cpu")
|
||||
|
||||
|
||||
noise_mask = None
|
||||
if "noise_mask" in latent:
|
||||
noise_mask = latent["noise_mask"]
|
||||
|
||||
ref_latents = []
|
||||
def callback(step: int, x0: T, x: T, steps: int):
|
||||
ref_latents.insert(0, x[0])
|
||||
|
||||
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
|
||||
samples = comfy.sample.sample_custom(model, noise, cfg, sampler, sigmas, positive, negative, latent_image, noise_mask=noise_mask, callback=callback, disable_pbar=disable_pbar, seed=noise_seed)
|
||||
|
||||
out = latent.copy()
|
||||
out["samples"] = samples
|
||||
out_noised = out
|
||||
|
||||
ref_latents = torch.stack(ref_latents)
|
||||
|
||||
return (ref_latents, out_noised)
|
||||
|
||||
class StyleAlignedReferenceSampler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
|
||||
"ref_latents": ("STEP_LATENTS",),
|
||||
"reference_image_text": ("STRING", {"multiline": True}),
|
||||
"model": ("MODEL",),
|
||||
"clip": ("CLIP", ),
|
||||
|
||||
"positive": ("CONDITIONING",),
|
||||
"negative": ("CONDITIONING",),
|
||||
|
||||
"share_norm": (SHARE_NORM_OPTIONS,),
|
||||
"share_attn": (SHARE_ATTN_OPTIONS,),
|
||||
"scale": ("FLOAT", {"default": 1, "min": 0, "max": 2.0, "step": 0.01}),
|
||||
"batch_size": ("INT", {"default": 2, "min": 1, "max": 8, "step": 1}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
|
||||
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step":0.1, "round": 0.01}),
|
||||
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, ),
|
||||
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT", "LATENT")
|
||||
RETURN_NAMES = ("output", "denoised_output")
|
||||
FUNCTION = "patch"
|
||||
# CATEGORY = "style_aligned"
|
||||
CATEGORY = "♾️Mixlab/Style"
|
||||
def patch(
|
||||
self,
|
||||
ref_latents,
|
||||
reference_image_text,
|
||||
model,
|
||||
clip,
|
||||
positive,
|
||||
negative,
|
||||
share_norm,
|
||||
share_attn,
|
||||
scale,
|
||||
batch_size,
|
||||
seed,steps,cfg,scheduler,denoise
|
||||
|
||||
) -> "tuple[dict, dict]":
|
||||
|
||||
m = model.clone()
|
||||
|
||||
# ref_latents = vae.encode(reference_image[:,:,:,:3])
|
||||
|
||||
tokens = clip.tokenize(reference_image_text)
|
||||
cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True)
|
||||
ref_positive=[[cond, {"pooled_output": pooled}]]
|
||||
|
||||
noise_seed=seed
|
||||
|
||||
|
||||
total_steps = steps
|
||||
if denoise < 1.0:
|
||||
total_steps = int(steps/denoise)
|
||||
|
||||
# comfy.model_management.load_models_gpu([model])
|
||||
sigmas = comfy.samplers.calculate_sigmas_scheduler(model.model, scheduler, total_steps).cpu()
|
||||
sigmas = sigmas[-(steps + 1):]
|
||||
|
||||
sampler_name="ddim"
|
||||
|
||||
sampler = comfy.samplers.sampler_object(sampler_name)
|
||||
|
||||
args = StyleAlignedArgs(share_attn)
|
||||
|
||||
# Concat batch with style latent
|
||||
style_latent_tensor = ref_latents[0].unsqueeze(0)
|
||||
height, width = style_latent_tensor.shape[-2:]
|
||||
latent_t = torch.zeros(
|
||||
[batch_size, 4, height, width], device=ref_latents.device
|
||||
)
|
||||
latent = {"samples": latent_t}
|
||||
noise = comfy.sample.prepare_noise(latent_t, noise_seed)
|
||||
|
||||
latent_t = torch.cat((style_latent_tensor, latent_t), dim=0)
|
||||
ref_noise = torch.zeros_like(noise[0]).unsqueeze(0)
|
||||
noise = torch.cat((ref_noise, noise), dim=0)
|
||||
|
||||
x0_output = {}
|
||||
preview_callback = latent_preview.prepare_callback(m, sigmas.shape[-1] - 1, x0_output)
|
||||
|
||||
# Replace first latent with the corresponding reference latent after each step
|
||||
def callback(step: int, x0: T, x: T, steps: int):
|
||||
preview_callback(step, x0, x, steps)
|
||||
if (step + 1 < steps):
|
||||
# 当ref_latents的step不够时
|
||||
if step+1>len(ref_latents)-1:
|
||||
step=len(ref_latents)-2
|
||||
|
||||
x[0] = ref_latents[step+1]
|
||||
x0[0] = ref_latents[step+1]
|
||||
|
||||
# Register shared norms
|
||||
share_group_norm = share_norm in ["group", "both"]
|
||||
share_layer_norm = share_norm in ["layer", "both"]
|
||||
register_shared_norm(m, share_group_norm, share_layer_norm)
|
||||
|
||||
# Patch cross attn
|
||||
m.set_model_attn1_patch(SharedAttentionProcessor(args, scale))
|
||||
|
||||
# Add reference conditioning to batch
|
||||
batched_condition = []
|
||||
for i,condition in enumerate(positive):
|
||||
additional = condition[1].copy()
|
||||
batch_with_reference = torch.cat([ref_positive[i][0], condition[0].repeat([batch_size] + [1] * len(condition[0].shape[1:]))], dim=0)
|
||||
if 'pooled_output' in additional and 'pooled_output' in ref_positive[i][1]:
|
||||
# combine pooled output
|
||||
pooled_output = torch.cat([ref_positive[i][1]['pooled_output'], additional['pooled_output'].repeat([batch_size]
|
||||
+ [1] * len(additional['pooled_output'].shape[1:]))], dim=0)
|
||||
additional['pooled_output'] = pooled_output
|
||||
if 'control' in additional:
|
||||
if 'control' in ref_positive[i][1]:
|
||||
# combine control conditioning
|
||||
control_hint = torch.cat([ref_positive[i][1]['control'].cond_hint_original, additional['control'].cond_hint_original.repeat([batch_size]
|
||||
+ [1] * len(additional['control'].cond_hint_original.shape[1:]))], dim=0)
|
||||
cloned_controlnet = additional['control'].copy()
|
||||
cloned_controlnet.set_cond_hint(control_hint, strength=additional['control'].strength, timestep_percent_range=additional['control'].timestep_percent_range)
|
||||
additional['control'] = cloned_controlnet
|
||||
else:
|
||||
# add zeros for first in batch
|
||||
control_hint = torch.cat([torch.zeros_like(additional['control'].cond_hint_original), additional['control'].cond_hint_original.repeat([batch_size]
|
||||
+ [1] * len(additional['control'].cond_hint_original.shape[1:]))], dim=0)
|
||||
cloned_controlnet = additional['control'].copy()
|
||||
cloned_controlnet.set_cond_hint(control_hint, strength=additional['control'].strength, timestep_percent_range=additional['control'].timestep_percent_range)
|
||||
additional['control'] = cloned_controlnet
|
||||
batched_condition.append([batch_with_reference, additional])
|
||||
|
||||
disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED
|
||||
samples = comfy.sample.sample_custom(
|
||||
m,
|
||||
noise,
|
||||
cfg,
|
||||
sampler,
|
||||
sigmas,
|
||||
batched_condition,
|
||||
negative,
|
||||
latent_t,
|
||||
callback=callback,
|
||||
disable_pbar=disable_pbar,
|
||||
seed=noise_seed,
|
||||
)
|
||||
|
||||
# remove reference image
|
||||
samples = samples[1:]
|
||||
|
||||
out = latent.copy()
|
||||
out["samples"] = samples
|
||||
if "x0" in x0_output:
|
||||
out_denoised = latent.copy()
|
||||
x0 = x0_output["x0"][1:]
|
||||
out_denoised["samples"] = m.model.process_latent_out(x0.cpu())
|
||||
else:
|
||||
out_denoised = out
|
||||
return (out, out_denoised)
|
||||
|
||||
|
||||
class StyleAlignedBatchAlign:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"share_norm": (SHARE_NORM_OPTIONS,),
|
||||
"share_attn": (SHARE_ATTN_OPTIONS,),
|
||||
"scale": ("FLOAT", {"default": 1, "min": 0, "max": 1.0, "step": 0.1}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "patch"
|
||||
# CATEGORY = "style_aligned"
|
||||
CATEGORY = "♾️Mixlab/Style"
|
||||
def patch(
|
||||
self,
|
||||
model: ModelPatcher,
|
||||
share_norm: str,
|
||||
share_attn: str,
|
||||
scale: float,
|
||||
):
|
||||
m = model.clone()
|
||||
share_group_norm = share_norm in ["group", "both"]
|
||||
share_layer_norm = share_norm in ["layer", "both"]
|
||||
register_shared_norm(model, share_group_norm, share_layer_norm)
|
||||
args = StyleAlignedArgs(share_attn)
|
||||
m.set_model_attn1_patch(SharedAttentionProcessor(args, scale))
|
||||
return (m,)
|
||||
|
||||
|
||||
@@ -0,0 +1,439 @@
|
||||
from transformers import pipeline, set_seed,AutoTokenizer, AutoModelForSeq2SeqLM
|
||||
import random
|
||||
import re
|
||||
|
||||
import os,sys
|
||||
import folder_paths
|
||||
|
||||
# from PIL import Image
|
||||
import importlib.util
|
||||
|
||||
import comfy.utils
|
||||
# import numpy as np
|
||||
import torch
|
||||
import random
|
||||
from lark import Lark, Transformer, v_args
|
||||
|
||||
|
||||
global _available
|
||||
_available=True
|
||||
|
||||
def get_text_generator_path():
|
||||
try:
|
||||
return folder_paths.get_folder_paths('prompt_generator')[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, "prompt_generator")
|
||||
|
||||
prompt_generator=get_text_generator_path()
|
||||
|
||||
text_generator_model_path=os.path.join(prompt_generator, "text2image-prompt-generator")
|
||||
if not os.path.exists(text_generator_model_path):
|
||||
print(f"## text_generator_model not found: {text_generator_model_path}, pls download from https://huggingface.co/succinctly/text2image-prompt-generator/tree/main")
|
||||
text_generator_model_path='succinctly/text2image-prompt-generator'
|
||||
|
||||
zh_en_model_path=os.path.join(prompt_generator, "opus-mt-zh-en")
|
||||
if not os.path.exists(zh_en_model_path):
|
||||
print(f"## zh_en_model not found: {zh_en_model_path}, pls download from https://huggingface.co/Helsinki-NLP/opus-mt-zh-en/tree/main")
|
||||
zh_en_model_path='Helsinki-NLP/opus-mt-zh-en'
|
||||
|
||||
|
||||
def is_installed(package):
|
||||
try:
|
||||
spec = importlib.util.find_spec(package)
|
||||
except ModuleNotFoundError:
|
||||
return False
|
||||
return spec is not None
|
||||
|
||||
|
||||
try:
|
||||
if is_installed('sentencepiece')==False:
|
||||
import subprocess
|
||||
|
||||
# 安装
|
||||
print('#pip install sentencepiece')
|
||||
|
||||
result = subprocess.run([sys.executable, '-s', '-m', 'pip', 'install', 'sentencepiece'], capture_output=True, text=True)
|
||||
|
||||
#检查命令执行结果
|
||||
if result.returncode == 0 and is_installed('sentencepiece'):
|
||||
print("#install success")
|
||||
_available=True
|
||||
else:
|
||||
print("#install error")
|
||||
_available=False
|
||||
|
||||
else:
|
||||
_available=True
|
||||
|
||||
except:
|
||||
_available=False
|
||||
|
||||
|
||||
|
||||
def translate(text):
|
||||
global text_pipe,zh_en_model,zh_en_tokenizer
|
||||
|
||||
if zh_en_model==None:
|
||||
zh_en_model = AutoModelForSeq2SeqLM.from_pretrained(zh_en_model_path).eval()
|
||||
zh_en_tokenizer = AutoTokenizer.from_pretrained(zh_en_model_path,padding=True, truncation=True)
|
||||
|
||||
zh_en_model.to("cuda" if torch.cuda.is_available() else "cpu")
|
||||
with torch.no_grad():
|
||||
encoded = zh_en_tokenizer([text], return_tensors="pt")
|
||||
encoded.to(zh_en_model.device)
|
||||
sequences = zh_en_model.generate(**encoded)
|
||||
return zh_en_tokenizer.batch_decode(sequences, skip_special_tokens=True)[0]
|
||||
|
||||
# input = "青春不能回头,所以青春没有终点。 ——《火影忍者》"
|
||||
# print(input, translate(input))
|
||||
|
||||
|
||||
def text_generate(text_pipe,input,seed=None):
|
||||
|
||||
if seed==None:
|
||||
seed = random.randint(100, 1000000)
|
||||
|
||||
set_seed(seed)
|
||||
|
||||
for count in range(6):
|
||||
sequences = text_pipe(input, max_length=random.randint(60, 90), num_return_sequences=8)
|
||||
list = []
|
||||
for sequence in sequences:
|
||||
line = sequence['generated_text'].strip()
|
||||
if line != input and len(line) > (len(input) + 4) and line.endswith((":", "-", "—")) is False:
|
||||
list.append(line)
|
||||
|
||||
result = "\n".join(list)
|
||||
result = re.sub('[^ ]+\.[^ ]+','', result)
|
||||
result = result.replace("<", "").replace(">", "")
|
||||
if result != "":
|
||||
return result
|
||||
if count == 5:
|
||||
return result
|
||||
|
||||
# input = "Youth can't turn back, so there's no end to youth."
|
||||
# print(input, text_generate(input))
|
||||
|
||||
|
||||
import re
|
||||
|
||||
def correct_prompt_syntax(prompt=""):
|
||||
|
||||
# print("input prompt",prompt)
|
||||
corrected_elements = []
|
||||
# 处理成统一的英文标点
|
||||
prompt = prompt.replace('(', '(').replace(')', ')').replace(',', ',').replace(';', ',').replace('。', '.').replace(':',':')
|
||||
# 删除多余的空格
|
||||
prompt = re.sub(r'\s+', ' ', prompt).strip()
|
||||
prompt = prompt.replace("< ","<").replace(" >",">").replace("( ","(").replace(" )",")").replace("[ ","[").replace(' ]',']')
|
||||
|
||||
# 分词
|
||||
prompt_elements = prompt.split(',')
|
||||
|
||||
def balance_brackets(element, open_bracket, close_bracket):
|
||||
open_brackets_count = element.count(open_bracket)
|
||||
close_brackets_count = element.count(close_bracket)
|
||||
return element + close_bracket * (open_brackets_count - close_brackets_count)
|
||||
|
||||
for element in prompt_elements:
|
||||
element = element.strip()
|
||||
|
||||
# 处理空元素
|
||||
if not element:
|
||||
continue
|
||||
|
||||
# 检查并处理圆括号、方括号、尖括号
|
||||
if element[0] in '([':
|
||||
corrected_element = balance_brackets(element, '(', ')') if element[0] == '(' else balance_brackets(element, '[', ']')
|
||||
elif element[0] == '<':
|
||||
corrected_element = balance_brackets(element, '<', '>')
|
||||
else:
|
||||
# 删除开头的右括号或右方括号
|
||||
corrected_element = element.lstrip(')]')
|
||||
|
||||
corrected_elements.append(corrected_element)
|
||||
|
||||
# 重组修正后的prompt
|
||||
return ','.join(corrected_elements)
|
||||
|
||||
|
||||
# # 示例使用
|
||||
# test_prompt = "((middle-century castles)), [forsaken: 0.8], (mystery dragons: 1.3, mist forests, sunsets, quiet; (((dummy)), [fisting city: 0.5] background, radiant, soft and flavoured,] promising mountains, ((starry: 1.6), [[crowds], [middle-century castle: urban landscapes of the future: 0.5], [yellow: bright sun: 0.7], overlooking"
|
||||
# corrected_prompt = correct_prompt_syntax(test_prompt)
|
||||
# print(corrected_prompt)
|
||||
|
||||
def detect_language(input_str):
|
||||
# 统计中文和英文字符的数量
|
||||
count_cn = count_en = 0
|
||||
for char in input_str:
|
||||
if '\u4e00' <= char <= '\u9fff':
|
||||
count_cn += 1
|
||||
elif char.isalpha():
|
||||
count_en += 1
|
||||
|
||||
# 根据统计的字符数量判断主要语言
|
||||
if count_cn > count_en:
|
||||
return "cn"
|
||||
elif count_en > count_cn:
|
||||
return "en"
|
||||
else:
|
||||
return "unknow"
|
||||
|
||||
|
||||
|
||||
|
||||
#定义Prompt文法
|
||||
grammar = """
|
||||
start: sentence
|
||||
sentence: phrase ("," phrase)*
|
||||
phrase: emphasis | weight | word | lora | embedding | schedule
|
||||
emphasis: "(" sentence ")" -> emphasis
|
||||
| "[" sentence "]" -> weak_emphasis
|
||||
weight: "(" word ":" NUMBER ")"
|
||||
schedule: "[" word ":" word ":" NUMBER "]"
|
||||
lora: "<" WORD ":" WORD (":" NUMBER)? (":" NUMBER)? ">"
|
||||
embedding: "embedding" ":" WORD (":" NUMBER)? (":" NUMBER)?
|
||||
word: WORD
|
||||
|
||||
NUMBER: /\s*-?\d+(\.\d+)?\s*/
|
||||
WORD: /[^,:\(\)\[\]<>]+/
|
||||
"""
|
||||
|
||||
|
||||
|
||||
@v_args(inline=True) # Decorator to flatten the tree directly into the function arguments
|
||||
class ChinesePromptTranslate(Transformer):
|
||||
|
||||
def sentence(self, *args):
|
||||
return ", ".join(args)
|
||||
|
||||
def phrase(self, *args):
|
||||
return "".join(args)
|
||||
|
||||
def emphasis(self, *args):
|
||||
# Reconstruct the emphasis with translated content
|
||||
return "(" + "".join(args) + ")"
|
||||
|
||||
def weak_emphasis(self, *args):
|
||||
print('weak_emphasis:',args)
|
||||
return "[" + "".join(args) + "]"
|
||||
|
||||
def embedding(self,*args):
|
||||
print('prompt embedding',args[0])
|
||||
if len(args) == 1:
|
||||
# print('prompt embedding',str(args[0]))
|
||||
# 只传递了一个参数,意味着只有embedding名称没有数字
|
||||
embedding_name = str(args[0])
|
||||
return f"embedding:{embedding_name}"
|
||||
elif len(args) > 1:
|
||||
embedding_name,*numbers = args
|
||||
|
||||
if len(numbers)==2:
|
||||
return f"embedding:{embedding_name}:{numbers[0]}:{numbers[1]}"
|
||||
elif len(numbers)==1:
|
||||
return f"embedding:{embedding_name}:{numbers[0]}"
|
||||
else:
|
||||
return f"embedding:{embedding_name}"
|
||||
|
||||
def lora(self,*args):
|
||||
print('lora prompt',*args)
|
||||
if len(args) == 1:
|
||||
return f"<lora:{loar_name}>"
|
||||
elif len(args) > 1:
|
||||
# print('lora', args)
|
||||
_,loar_name,*numbers = args
|
||||
loar_name = str(loar_name).strip()
|
||||
if len(numbers)==2:
|
||||
return f"<lora:{loar_name}:{numbers[0]}:{numbers[1]}>"
|
||||
elif len(numbers)==1:
|
||||
return f"<lora:{loar_name}:{numbers[0]}>"
|
||||
else:
|
||||
return f"<lora:{loar_name}>"
|
||||
|
||||
def weight(self, word,number):
|
||||
translated_word = translate(str(word)).rstrip('.')
|
||||
return f"({translated_word}:{str(number).strip()})"
|
||||
|
||||
def schedule(self,*args):
|
||||
print('prompt schedule',args)
|
||||
data = [str(arg).strip() for arg in args]
|
||||
|
||||
return f"[{':'.join(data)}]"
|
||||
|
||||
def word(self, word):
|
||||
# Translate each word using the dictionary
|
||||
if detect_language(str(word)) == "cn":
|
||||
return translate(str(word)).rstrip('.')
|
||||
else:
|
||||
return str(word).rstrip('.')
|
||||
|
||||
class ChinesePrompt:
|
||||
|
||||
global _available
|
||||
available=_available
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"text": ("STRING",{"multiline": True,"default": "", "dynamicPrompts": False}),
|
||||
"generation": (["on","off"],{"default": "off"}),
|
||||
},
|
||||
|
||||
"optional":{
|
||||
"seed":("INT", {"default": 100, "min": 100, "max": 0xffffffffffffffff}),
|
||||
|
||||
},
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("prompt",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
global text_pipe,zh_en_model,zh_en_tokenizer
|
||||
|
||||
text_pipe= None
|
||||
zh_en_model=None
|
||||
zh_en_tokenizer=None
|
||||
|
||||
def run(self,text,seed,generation):
|
||||
|
||||
|
||||
seed=seed[0]
|
||||
generation=generation[0]
|
||||
|
||||
# 进度条
|
||||
pbar = comfy.utils.ProgressBar(len(text)+1)
|
||||
texts = [correct_prompt_syntax(t) for t in text]
|
||||
|
||||
global text_pipe,zh_en_model,zh_en_tokenizer
|
||||
if zh_en_model==None:
|
||||
zh_en_model = AutoModelForSeq2SeqLM.from_pretrained(zh_en_model_path).eval()
|
||||
zh_en_tokenizer = AutoTokenizer.from_pretrained(zh_en_model_path,padding=True, truncation=True)
|
||||
|
||||
zh_en_model.to("cuda" if torch.cuda.is_available() else "cpu")
|
||||
# zh_en_tokenizer.to("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
text_pipe=pipeline('text-generation', model=text_generator_model_path,device="cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
# text_pipe.model.to("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
prompt_result=[]
|
||||
|
||||
# print('zh_en_model device',zh_en_model.device,text_pipe.model.device,torch.cuda.current_device() )
|
||||
en_texts=[]
|
||||
|
||||
for t in texts:
|
||||
if t:
|
||||
parser = Lark(grammar, start="start", parser="lalr", transformer=ChinesePromptTranslate())
|
||||
try:
|
||||
result = parser.parse(t).children
|
||||
en_texts.append(result[0])
|
||||
except:
|
||||
print(f"Error parsing '{t}'")
|
||||
t = translate(str(t))
|
||||
en_texts.append(t)
|
||||
|
||||
|
||||
zh_en_model.to('cpu')
|
||||
print("test en_text",en_texts)
|
||||
# en_text.to("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
pbar.update(1)
|
||||
for t in en_texts:
|
||||
if generation=='on':
|
||||
prompt =text_generate(text_pipe,t,seed)
|
||||
# 多条,还是单条
|
||||
lines = prompt.split("\n")
|
||||
longest_line = max(lines, key=len)
|
||||
# print(longest_line)
|
||||
prompt_result.append(longest_line)
|
||||
else:
|
||||
prompt_result.append(t)
|
||||
pbar.update(1)
|
||||
|
||||
text_pipe.model.to('cpu')
|
||||
|
||||
print('prompt_result',prompt_result,)
|
||||
# prompt_result = [','.join(correct_prompt_syntax(p)) for p in prompt_result]
|
||||
if len(prompt_result)==0:
|
||||
prompt_result=[""]
|
||||
return {
|
||||
"ui":{
|
||||
"prompt": prompt_result
|
||||
},
|
||||
"result": (prompt_result,)}
|
||||
|
||||
|
||||
|
||||
|
||||
class PromptGenerate:
|
||||
|
||||
global _available
|
||||
available=_available
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"text": ("STRING",{"multiline": True,"default": "", "dynamicPrompts": False}),
|
||||
},
|
||||
|
||||
"optional":{
|
||||
"multiple": (["off","on"],),
|
||||
"seed":("INT", {"default": 100, "min": 100, "max": 0xffffffffffffffff}),
|
||||
},
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("prompt",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Prompt"
|
||||
OUTPUT_NODE = True
|
||||
INPUT_IS_LIST = True
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
global text_pipe
|
||||
|
||||
text_pipe= None
|
||||
#
|
||||
|
||||
def run(self,text,multiple,seed):
|
||||
global text_pipe
|
||||
|
||||
seed=seed[0]
|
||||
|
||||
multiple=multiple[0]
|
||||
|
||||
# 进度条
|
||||
pbar = comfy.utils.ProgressBar(len(text))
|
||||
|
||||
text_pipe=pipeline('text-generation', model=text_generator_model_path,device="cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
prompt_result=[]
|
||||
|
||||
for t in text:
|
||||
prompt =text_generate(text_pipe,t,seed)
|
||||
prompt = prompt.split("\n")
|
||||
if multiple=='off':
|
||||
prompt = [max(prompt, key=len)]
|
||||
|
||||
for p in prompt:
|
||||
prompt_result.append(p)
|
||||
pbar.update(1)
|
||||
|
||||
text_pipe.model.to('cpu')
|
||||
|
||||
return {
|
||||
"ui":{
|
||||
"prompt": prompt_result
|
||||
},
|
||||
"result": (prompt_result,)}
|
||||
@@ -0,0 +1,180 @@
|
||||
import sys
|
||||
from os import path
|
||||
sys.path.insert(0, path.dirname(__file__))
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from folder_paths import get_folder_paths, get_full_path, get_save_image_path, get_output_directory,models_dir
|
||||
from comfy.model_management import get_torch_device
|
||||
from .tsr.system import TSR
|
||||
|
||||
import comfy.utils
|
||||
|
||||
|
||||
def get_triposr_model_path():
|
||||
try:
|
||||
return path.join(get_folder_paths('triposr')[0],'model.ckpt')
|
||||
except:
|
||||
return path.join(path.join(models_dir, "triposr"),'model.ckpt')
|
||||
|
||||
triposr_model_path=get_triposr_model_path()
|
||||
|
||||
|
||||
# Tensor to PIL
|
||||
def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
# Convert PIL to Tensor
|
||||
def pil2tensor(image):
|
||||
return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
|
||||
def fill_background(image):
|
||||
im = np.array(image).astype(np.float32) / 255.0
|
||||
im = im[:, :, :3] * im[:, :, 3:4] + (1 - im[:, :, 3:4]) * 0.5
|
||||
im = Image.fromarray((im * 255.0).astype(np.uint8))
|
||||
return im
|
||||
|
||||
|
||||
class LoadTripoSRModel:
|
||||
def __init__(self):
|
||||
self.initialized_model = None
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
# "model": (get_filename_list("checkpoints"),),
|
||||
"chunk_size": ("INT", {"default": 8192, "min": 0, "max": 10000})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("TRIPOSR_MODEL",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/3D/TripoSR"
|
||||
|
||||
def run(self, chunk_size):
|
||||
device = get_torch_device()
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
device = "cpu"
|
||||
|
||||
if not self.initialized_model:
|
||||
# triposr_model_path
|
||||
print("#Loading TripoSR model",triposr_model_path)
|
||||
self.initialized_model = TSR.from_pretrained_custom(
|
||||
weight_path=triposr_model_path,
|
||||
config_path=path.join(path.dirname(__file__), "tsr/config.yaml")
|
||||
)
|
||||
self.initialized_model.renderer.set_chunk_size(chunk_size)
|
||||
self.initialized_model.to(device)
|
||||
|
||||
return (self.initialized_model,)
|
||||
|
||||
|
||||
class TripoSRSampler:
|
||||
def __init__(self):
|
||||
self.initialized_model = None
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("TRIPOSR_MODEL",),
|
||||
"image": ("IMAGE",),
|
||||
"resolution": ("INT", {"default": 256, "min": 128, "max": 12288}),
|
||||
"threshold": ("FLOAT", {"default": 25.0, "min": 0.0, "step": 0.01}),
|
||||
"device":(["auto","cpu"],),
|
||||
},
|
||||
"optional": {
|
||||
"mask": ("MASK",)
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MESH",)
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/3D/TripoSR"
|
||||
|
||||
def run(self, model, image, resolution, threshold,device='auto', mask=None):
|
||||
|
||||
reference_image=image
|
||||
reference_mask=mask
|
||||
|
||||
device = get_torch_device()
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
device = "cpu"
|
||||
|
||||
if device=='cpu':
|
||||
device = "cpu"
|
||||
|
||||
print('#TripoSRSampler device',device)
|
||||
|
||||
to_images=[]
|
||||
|
||||
for i in range(len(reference_image)):
|
||||
|
||||
image = reference_image[i]
|
||||
|
||||
if reference_mask is not None:
|
||||
mask = reference_mask[i].unsqueeze(2)
|
||||
image = torch.cat((image, mask), dim=2).detach().cpu().numpy()
|
||||
image = Image.fromarray(np.clip(255. * image, 0, 255).astype(np.uint8))
|
||||
image = fill_background(image)
|
||||
else:
|
||||
image = tensor2pil(image)
|
||||
|
||||
image = image.convert('RGB')
|
||||
|
||||
to_images.append(image)
|
||||
|
||||
# 进度条
|
||||
pbar = comfy.utils.ProgressBar(len(to_images))
|
||||
def callback(c):
|
||||
pbar.update(1)
|
||||
|
||||
scene_codes = model(to_images, device)
|
||||
meshes = model.extract_mesh(scene_codes, resolution=resolution, threshold=threshold,callback=callback)
|
||||
|
||||
del model
|
||||
return (meshes,)
|
||||
|
||||
|
||||
class SaveTripoSRMesh:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"mesh": ("MESH",),
|
||||
# "format":(["glb","obj"],),
|
||||
"filename_prefix":("STRING", {"multiline": False,"default": "TripoSR_"})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
OUTPUT_NODE = True
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "♾️Mixlab/3D/TripoSR"
|
||||
|
||||
def run(self, mesh,filename_prefix):
|
||||
format='glb'
|
||||
saved = list()
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = get_save_image_path(filename_prefix,
|
||||
get_output_directory())
|
||||
|
||||
for (index, single_mesh) in enumerate(mesh):
|
||||
filename_with_batch_num = filename.replace("%batch_num%", str(index))
|
||||
file = f"{filename_with_batch_num}_{counter:05}_.{format}"
|
||||
single_mesh.apply_transform(np.array([[1, 0, 0, 0], [0, 0, 1, 0], [0, -1, 0, 0], [0, 0, 0, 1]]))
|
||||
single_mesh.export(path.join(full_output_folder, file))
|
||||
saved.append({
|
||||
"filename": file,
|
||||
"type": "output",
|
||||
"subfolder": subfolder
|
||||
})
|
||||
|
||||
return {"ui": {"mesh": saved}}
|
||||
|
||||
|
||||
|
||||
@@ -1,179 +0,0 @@
|
||||
# https://github.com/openai/consistencydecoder/blob/main/consistencydecoder/__init__.py
|
||||
|
||||
import folder_paths
|
||||
from comfy import model_management
|
||||
|
||||
import math
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
|
||||
class ConsistencyDecoderWrapper:
|
||||
def __init__(self, decoder):
|
||||
self.decoder = decoder
|
||||
def decode(self, x):
|
||||
return self.decoder(x)
|
||||
|
||||
def _extract_into_tensor(arr, timesteps, broadcast_shape):
|
||||
|
||||
# from: https://github.com/openai/guided-diffusion/blob/22e0df8183507e13a7813f8d38d51b072ca1e67c/guided_diffusion/gaussian_diffusion.py#L895 """
|
||||
res = arr[timesteps].float()
|
||||
dims_to_append = len(broadcast_shape) - len(res.shape)
|
||||
return res[(...,) + (None,) * dims_to_append]
|
||||
|
||||
|
||||
def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999):
|
||||
# from: https://github.com/openai/guided-diffusion/blob/22e0df8183507e13a7813f8d38d51b072ca1e67c/guided_diffusion/gaussian_diffusion.py#L45
|
||||
betas = []
|
||||
for i in range(num_diffusion_timesteps):
|
||||
t1 = i / num_diffusion_timesteps
|
||||
t2 = (i + 1) / num_diffusion_timesteps
|
||||
betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta))
|
||||
return torch.tensor(betas)
|
||||
|
||||
class ConsistencyDecoder:
|
||||
def __init__(self, device="cuda:0", download_target=""):
|
||||
self.n_distilled_steps = 64
|
||||
# download_target = _download("https://openaipublic.azureedge.net/diff-vae/c9cebd3132dd9c42936d803e33424145a748843c8f716c0814838bdc8a2fe7cb/decoder.pt", download_root)
|
||||
self.ckpt = torch.jit.load(download_target).to(device)
|
||||
self.device = device
|
||||
sigma_data = 0.5
|
||||
betas = betas_for_alpha_bar(
|
||||
1024, lambda t: math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
|
||||
).to(device)
|
||||
alphas = 1.0 - betas
|
||||
alphas_cumprod = torch.cumprod(alphas, dim=0)
|
||||
self.sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)
|
||||
self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)
|
||||
sqrt_recip_alphas_cumprod = torch.sqrt(1.0 / alphas_cumprod)
|
||||
sigmas = torch.sqrt(1.0 / alphas_cumprod - 1)
|
||||
self.c_skip = (
|
||||
sqrt_recip_alphas_cumprod
|
||||
* sigma_data**2
|
||||
/ (sigmas**2 + sigma_data**2)
|
||||
)
|
||||
self.c_out = sigmas * sigma_data / (sigmas**2 + sigma_data**2) ** 0.5
|
||||
self.c_in = sqrt_recip_alphas_cumprod / (sigmas**2 + sigma_data**2) ** 0.5
|
||||
|
||||
@staticmethod
|
||||
def round_timesteps(
|
||||
timesteps, total_timesteps, n_distilled_steps, truncate_start=True
|
||||
):
|
||||
with torch.no_grad():
|
||||
space = torch.div(total_timesteps, n_distilled_steps, rounding_mode="floor")
|
||||
rounded_timesteps = (
|
||||
torch.div(timesteps, space, rounding_mode="floor") + 1
|
||||
) * space
|
||||
if truncate_start:
|
||||
rounded_timesteps[rounded_timesteps == total_timesteps] -= space
|
||||
else:
|
||||
rounded_timesteps[rounded_timesteps == total_timesteps] -= space
|
||||
rounded_timesteps[rounded_timesteps == 0] += space
|
||||
return rounded_timesteps
|
||||
|
||||
@staticmethod
|
||||
def ldm_transform_latent(z, extra_scale_factor=1):
|
||||
channel_means = [0.38862467, 0.02253063, 0.07381133, -0.0171294]
|
||||
channel_stds = [0.9654121, 1.0440036, 0.76147926, 0.77022034]
|
||||
|
||||
if len(z.shape) != 4:
|
||||
raise ValueError()
|
||||
|
||||
z = z * 0.18215
|
||||
channels = [z[:, i] for i in range(z.shape[1])]
|
||||
|
||||
channels = [
|
||||
extra_scale_factor * (c - channel_means[i]) / channel_stds[i]
|
||||
for i, c in enumerate(channels)
|
||||
]
|
||||
return torch.stack(channels, dim=1)
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(
|
||||
self,
|
||||
features: torch.Tensor,
|
||||
schedule=[1.0, 0.5],
|
||||
):
|
||||
features = self.ldm_transform_latent(features)
|
||||
|
||||
ts = self.round_timesteps(
|
||||
torch.arange(0, 1024),
|
||||
1024,
|
||||
self.n_distilled_steps,
|
||||
truncate_start=False,
|
||||
)
|
||||
shape = (
|
||||
features.size(0),
|
||||
3,
|
||||
8 * features.size(2),
|
||||
8 * features.size(3),
|
||||
)
|
||||
|
||||
x_start = torch.zeros(shape, device=features.device, dtype=features.dtype)
|
||||
schedule_timesteps = [int((1024 - 1) * s) for s in schedule]
|
||||
for i in schedule_timesteps:
|
||||
t = ts[i].item()
|
||||
t_ = torch.tensor([t] * features.shape[0]).to(self.device)
|
||||
noise = torch.randn_like(x_start)
|
||||
|
||||
x_start = (
|
||||
_extract_into_tensor(self.sqrt_alphas_cumprod, t_, x_start.shape)
|
||||
* x_start
|
||||
+ _extract_into_tensor(
|
||||
self.sqrt_one_minus_alphas_cumprod, t_, x_start.shape
|
||||
)
|
||||
* noise
|
||||
)
|
||||
c_in = _extract_into_tensor(self.c_in, t_, x_start.shape)
|
||||
model_output = self.ckpt(c_in * x_start, t_, features=features)
|
||||
B, C = x_start.shape[:2]
|
||||
model_output, _ = torch.split(model_output, C, dim=1)
|
||||
pred_xstart = (
|
||||
_extract_into_tensor(self.c_out, t_, x_start.shape) * model_output
|
||||
+ _extract_into_tensor(self.c_skip, t_, x_start.shape) * x_start
|
||||
).clamp(-1, 1)
|
||||
x_start = pred_xstart
|
||||
return x_start
|
||||
|
||||
|
||||
|
||||
class VAELoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "vae_name": (folder_paths.get_filename_list("vae"), )}}
|
||||
RETURN_TYPES = ("VAE",)
|
||||
FUNCTION = "load_vae"
|
||||
|
||||
CATEGORY = "Mixlab/ConsistencyDecoder"
|
||||
|
||||
#TODO: scale factor?
|
||||
def load_vae(self, vae_name):
|
||||
vae_path = folder_paths.get_full_path("vae", vae_name)
|
||||
device = 'cuda:0'
|
||||
# print('device',device)
|
||||
consistencyDecoder = ConsistencyDecoder(device=device,
|
||||
download_target=vae_path) # Model size: 2.49 GB
|
||||
vae = ConsistencyDecoderWrapper(consistencyDecoder)
|
||||
return (vae,)
|
||||
|
||||
|
||||
class VAEDecode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "samples": ("LATENT", ), "vae": ("VAE", )}}
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "decode"
|
||||
|
||||
CATEGORY = "Mixlab/ConsistencyDecoder"
|
||||
|
||||
def decode(self, vae, samples):
|
||||
image = vae.decode(samples["samples"].to("cuda:0"))
|
||||
image = image[0].cpu().numpy()
|
||||
image = (image + 1.0) * 127.5
|
||||
image = image.clip(0, 255).astype(np.uint8)
|
||||
image = Image.fromarray(image.transpose(1, 2, 0))
|
||||
image = image.convert("RGB")
|
||||
image = np.array(image).astype(np.float32) / 255.0
|
||||
image = torch.from_numpy(image)[None,]
|
||||
return (image, )
|
||||
@@ -0,0 +1,45 @@
|
||||
from comfy.ldm.modules.attention import default, optimized_attention, optimized_attention_masked
|
||||
from .style_functions import adain, concat_first
|
||||
|
||||
class VisualStyleProcessor(object):
|
||||
def __init__(self,
|
||||
module_self,
|
||||
keys_scale: float = 1.0,
|
||||
enabled: bool = True,
|
||||
adain_queries: bool = True,
|
||||
adain_keys: bool = True,
|
||||
adain_values: bool = False
|
||||
):
|
||||
self.module_self = module_self
|
||||
self.keys_scale = keys_scale
|
||||
self.enabled = enabled
|
||||
self.adain_queries = adain_queries
|
||||
self.adain_keys = adain_keys
|
||||
self.adain_values = adain_values
|
||||
|
||||
def visual_style_forward(self, x, context, value, mask=None):
|
||||
q = self.module_self.to_q(x)
|
||||
context = default(context, x)
|
||||
k = self.module_self.to_k(context)
|
||||
if value is not None:
|
||||
v = self.module_self.to_v(value)
|
||||
del value
|
||||
else:
|
||||
v = self.module_self.to_v(context)
|
||||
|
||||
if self.enabled:
|
||||
if self.adain_queries:
|
||||
q = adain(q)
|
||||
if self.adain_keys:
|
||||
k = adain(k)
|
||||
if self.adain_values:
|
||||
v = adain(v)
|
||||
|
||||
k = concat_first(k, -2, self.keys_scale)
|
||||
v = concat_first(v, -2)
|
||||
|
||||
if mask is None:
|
||||
out = optimized_attention(q, k, v, self.module_self.heads)
|
||||
else:
|
||||
out = optimized_attention_masked(q, k, v, self.module_self.heads, mask)
|
||||
return self.module_self.to_out(out)
|
||||
@@ -0,0 +1,60 @@
|
||||
import torch
|
||||
|
||||
from einops import rearrange
|
||||
from dataclasses import dataclass
|
||||
|
||||
T = torch.Tensor
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StyleAlignedArgs:
|
||||
share_group_norm: bool = True
|
||||
share_layer_norm: bool = True,
|
||||
share_attention: bool = True
|
||||
adain_queries: bool = True
|
||||
adain_keys: bool = True
|
||||
adain_values: bool = False
|
||||
full_attention_share: bool = False
|
||||
keys_scale: float = 1.
|
||||
only_self_level: float = 0.
|
||||
|
||||
def expand_first(feat: T, scale=1., ) -> T:
|
||||
b = feat.shape[0]
|
||||
feat_style = torch.stack((feat[0], feat[b // 2])).unsqueeze(1)
|
||||
if scale == 1:
|
||||
feat_style = feat_style.expand(2, b // 2, *feat.shape[1:])
|
||||
else:
|
||||
feat_style = feat_style.repeat(1, b // 2, 1, 1, 1)
|
||||
feat_style = torch.cat([feat_style[:, :1], scale * feat_style[:, 1:]], dim=1)
|
||||
return feat_style.reshape(*feat.shape)
|
||||
|
||||
|
||||
def concat_first(feat: T, dim=2, scale=1.) -> T:
|
||||
feat_style = expand_first(feat, scale=scale)
|
||||
return torch.cat((feat, feat_style), dim=dim)
|
||||
|
||||
|
||||
def calc_mean_std(feat, eps: float = 1e-5) -> tuple[T, T]:
|
||||
feat_std = (feat.var(dim=-2, keepdims=True) + eps).sqrt()
|
||||
feat_mean = feat.mean(dim=-2, keepdims=True)
|
||||
return feat_mean, feat_std
|
||||
|
||||
|
||||
def adain(feat: T) -> T:
|
||||
feat_mean, feat_std = calc_mean_std(feat)
|
||||
feat_style_mean = expand_first(feat_mean)
|
||||
feat_style_std = expand_first(feat_std)
|
||||
feat = (feat - feat_mean) / feat_std
|
||||
feat = feat * feat_style_std + feat_style_mean
|
||||
return feat
|
||||
|
||||
def swapping_attention(key, value, chunk_size=2):
|
||||
chunk_length = key.size()[0] // chunk_size # [text-condition, null-condition]
|
||||
reference_image_index = [0] * chunk_length # [0 0 0 0 0]
|
||||
key = rearrange(key, "(b f) d c -> b f d c", f=chunk_length)
|
||||
key = key[:, reference_image_index] # ref to all
|
||||
key = rearrange(key, "b f d c -> (b f) d c")
|
||||
value = rearrange(value, "(b f) d c -> b f d c", f=chunk_length)
|
||||
value = value[:, reference_image_index] # ref to all
|
||||
value = rearrange(value, "b f d c -> (b f) d c")
|
||||
|
||||
return key, value
|
||||
@@ -0,0 +1,173 @@
|
||||
import os,re
|
||||
import sys,time
|
||||
from pathlib import Path
|
||||
import torchaudio
|
||||
import hashlib
|
||||
import torch
|
||||
import folder_paths
|
||||
import comfy.utils
|
||||
|
||||
from faster_whisper import WhisperModel
|
||||
|
||||
class AnyType(str):
|
||||
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
|
||||
|
||||
def __ne__(self, __value: object) -> bool:
|
||||
return False
|
||||
|
||||
any_type = AnyType("*")
|
||||
|
||||
def get_model_dir(m):
|
||||
try:
|
||||
return folder_paths.get_folder_paths(m)[0]
|
||||
except:
|
||||
return os.path.join(folder_paths.models_dir, m)
|
||||
|
||||
|
||||
|
||||
whisper_model_path=get_model_dir('whisper')
|
||||
|
||||
model_sizes=[
|
||||
d for d in os.listdir(whisper_model_path) if os.path.isdir(
|
||||
os.path.join(whisper_model_path, d)
|
||||
) and os.path.isfile(os.path.join(os.path.join(whisper_model_path, d), "config.json"))
|
||||
]
|
||||
|
||||
|
||||
class LoadWhisperModel:
|
||||
def __init__(self):
|
||||
self.model = None
|
||||
self.device="cuda" if torch.cuda.is_available() else "cpu"
|
||||
self.model_size=model_sizes[0]
|
||||
self.compute_type='float16'
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"model_size": (model_sizes,),
|
||||
"device": (["auto","cpu"],),
|
||||
"compute_type": (["float16","int8_float16","int8"],),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WHISPER",)
|
||||
RETURN_NAMES = ("whisper_model",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/Whisper"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
|
||||
def run(self,model_size,device,compute_type):
|
||||
|
||||
if device=="auto" and self.device!='cuda':
|
||||
self.device="cuda" if torch.cuda.is_available() else "cpu"
|
||||
self.model=None
|
||||
|
||||
if device=='cpu' and self.device!='cpu':
|
||||
self.device="cpu"
|
||||
self.model=None
|
||||
|
||||
if model_size!= self.model_size:
|
||||
self.model_size=model_size
|
||||
self.model=None
|
||||
|
||||
if compute_type!=self.compute_type:
|
||||
self.compute_type=compute_type
|
||||
self.model=None
|
||||
|
||||
if self.model==None:
|
||||
self.model = WhisperModel(
|
||||
os.path.join(whisper_model_path, self.model_size),
|
||||
device=self.device,
|
||||
compute_type=self.compute_type
|
||||
)
|
||||
|
||||
return (self.model,)
|
||||
|
||||
|
||||
class WhisperTranscribe:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"whisper_model": ("WHISPER",),
|
||||
"audio": ("AUDIO",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (any_type,"STRING","STRING","FLOAT",)
|
||||
RETURN_NAMES = ("result","srt","text","total_seconds",)
|
||||
|
||||
FUNCTION = "run"
|
||||
|
||||
CATEGORY = "♾️Mixlab/Audio/Whisper"
|
||||
|
||||
INPUT_IS_LIST = False
|
||||
# OUTPUT_IS_LIST = (False,False,False,)
|
||||
|
||||
def run(self,whisper_model,audio):
|
||||
|
||||
if 'audio_path' in audio and (not 'waveform' in audio):
|
||||
waveform, sample_rate = torchaudio.load(audio['audio_path'])
|
||||
waveform=waveform.mean(0)
|
||||
total_length_seconds = waveform.shape[0] / sample_rate
|
||||
waveform=waveform.numpy()
|
||||
|
||||
elif 'waveform' in audio and 'sample_rate' in audio:
|
||||
print("Original shape:", audio["waveform"].shape, isinstance(audio["waveform"], torch.Tensor)) # 打印原始形状
|
||||
waveform = audio["waveform"].squeeze(0) # Remove the added batch dimension
|
||||
sample_rate = audio["sample_rate"]
|
||||
|
||||
# if audio_sf != sampling_rate:
|
||||
# waveform = torchaudio.functional.resample(
|
||||
# waveform, orig_freq=audio_sf, new_freq=sampling_rate
|
||||
# )
|
||||
|
||||
waveform=waveform.mean(0)
|
||||
|
||||
total_length_seconds = waveform.shape[0] / sample_rate
|
||||
|
||||
waveform=waveform.numpy() #whisper_model.transcribe 旧版不支持直接传tensor,先用numpy
|
||||
|
||||
segments, info = whisper_model.transcribe(waveform, beam_size=5)
|
||||
|
||||
print("Detected language '%s' with probability %f" % (info.language, info.language_probability))
|
||||
|
||||
# Function to format time for SRT
|
||||
def format_time(seconds):
|
||||
millis = int((seconds - int(seconds)) * 1000)
|
||||
hours, remainder = divmod(int(seconds), 3600)
|
||||
minutes, seconds = divmod(remainder, 60)
|
||||
return f"{hours:02}:{minutes:02}:{seconds:02},{millis:03}"
|
||||
|
||||
# Prepare SRT content as a string
|
||||
results = []
|
||||
for i, segment in enumerate(segments):
|
||||
start_time = format_time(segment.start)
|
||||
end_time = format_time(segment.end)
|
||||
srt_content = f"{i + 1}\n"
|
||||
srt_content += f"{start_time} --> {end_time}\n"
|
||||
|
||||
text=segment.text.strip()
|
||||
|
||||
srt_content += f"{text}\n\n"
|
||||
|
||||
start_time=segment.start
|
||||
end_time=segment.end
|
||||
|
||||
|
||||
results.append({
|
||||
"srt_content":srt_content,
|
||||
"start_time":start_time,
|
||||
"end_time":end_time,
|
||||
"text":text,
|
||||
"language":[info.language]
|
||||
})
|
||||
|
||||
srt_content="\n".join([s['srt_content'] for s in results])
|
||||
text="\n".join([s['text'] for s in results])
|
||||
|
||||
return (results,srt_content,text,total_length_seconds,)
|
||||
|
||||
@@ -0,0 +1,174 @@
|
||||
import torch
|
||||
from PIL import Image, ImageOps, ImageSequence, ImageFile
|
||||
from PIL.PngImagePlugin import PngInfo
|
||||
|
||||
import numpy as np
|
||||
import os
|
||||
import folder_paths
|
||||
import node_helpers
|
||||
import hashlib
|
||||
from uuid import uuid4
|
||||
|
||||
# Tensor to PIL
|
||||
def tensor2pil(image):
|
||||
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
|
||||
|
||||
# tensor 取hash值
|
||||
def tensor_to_hash(tensor):
|
||||
# 将 Tensor 转换为 NumPy 数组
|
||||
np_array = tensor.cpu().numpy()
|
||||
|
||||
# 将 NumPy 数组转换为字节数据
|
||||
byte_data = np_array.tobytes()
|
||||
|
||||
# 计算哈希值
|
||||
hash_value = hashlib.md5(byte_data).hexdigest()
|
||||
|
||||
return hash_value
|
||||
|
||||
|
||||
def create_temp_file(image, uuid):
|
||||
output_dir = folder_paths.get_temp_directory()
|
||||
|
||||
(
|
||||
full_output_folder,
|
||||
filename,
|
||||
counter,
|
||||
subfolder,
|
||||
_,
|
||||
) = folder_paths.get_save_image_path(f'material_{uuid}', output_dir)
|
||||
|
||||
|
||||
image=tensor2pil(image)
|
||||
|
||||
image_file = f"{filename}_{counter:05}.png"
|
||||
|
||||
image_path=os.path.join(full_output_folder, image_file)
|
||||
|
||||
image.save(image_path,compress_level=4)
|
||||
|
||||
return (image_path,[{
|
||||
"filename": image_file,
|
||||
"subfolder": subfolder,
|
||||
"type": "temp"
|
||||
}])
|
||||
|
||||
|
||||
# image - tensor - 文件路径
|
||||
# loadImage的方法( 文件路径 - image-mask )
|
||||
class EditMask:
|
||||
|
||||
def __init__(self):
|
||||
self.image_id = None
|
||||
self.uuid = str(uuid4())
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"image": ("IMAGE",), # 表示一个张量
|
||||
|
||||
},
|
||||
|
||||
"optional":{
|
||||
"image_update": ("IMAGE_FILE",)
|
||||
},
|
||||
|
||||
}
|
||||
|
||||
CATEGORY = "♾️Mixlab/Mask"
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK")
|
||||
RETURN_NAMES = ("image", "mask")
|
||||
|
||||
FUNCTION = "edit"
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def edit(self, image,image_update=None):
|
||||
|
||||
# 根据image输入来判断是否是新的图片
|
||||
if self.image_id==None:
|
||||
self.image_id=tensor_to_hash(image)
|
||||
image_update=None
|
||||
else:
|
||||
image_id=tensor_to_hash(image)
|
||||
if image_id!=self.image_id:
|
||||
image_update=None
|
||||
self.image_id=image_id
|
||||
|
||||
|
||||
image_path=None
|
||||
# print('#image_update',self.image_id,image_update)
|
||||
if image_update==None:
|
||||
print('--')
|
||||
else:
|
||||
if 'images' in image_update:
|
||||
images=image_update['images']
|
||||
filename=images[0]['filename']
|
||||
subfolder=images[0]['subfolder']
|
||||
type=images[0]['type']
|
||||
name, base_dir=folder_paths.annotated_filepath(filename)
|
||||
if type.endswith("output"):
|
||||
base_dir = folder_paths.get_output_directory()
|
||||
elif type.endswith("input"):
|
||||
base_dir = folder_paths.get_input_directory()
|
||||
elif type.endswith("temp"):
|
||||
base_dir = folder_paths.get_temp_directory()
|
||||
#base_dir = folder_paths.get_input_directory()
|
||||
# print(base_dir,subfolder, name)
|
||||
image_path = os.path.join(base_dir,subfolder, name)
|
||||
|
||||
if image_path==None:
|
||||
image_path,images=create_temp_file(image, self.uuid)
|
||||
|
||||
print('#image_path',os.path.exists(image_path),image_path)
|
||||
# image_path = folder_paths.get_annotated_filepath(image) #文件名
|
||||
|
||||
if not os.path.exists(image_path):
|
||||
image_path,images=create_temp_file(image, self.uuid)
|
||||
|
||||
|
||||
img = node_helpers.pillow(Image.open, image_path)
|
||||
|
||||
output_images = []
|
||||
output_masks = []
|
||||
w, h = None, None
|
||||
|
||||
excluded_formats = ['MPO']
|
||||
|
||||
for i in ImageSequence.Iterator(img):
|
||||
i = node_helpers.pillow(ImageOps.exif_transpose, i)
|
||||
|
||||
if i.mode == 'I':
|
||||
i = i.point(lambda i: i * (1 / 255))
|
||||
image = i.convert("RGB")
|
||||
|
||||
if len(output_images) == 0:
|
||||
w = image.size[0]
|
||||
h = image.size[1]
|
||||
|
||||
if image.size[0] != w or image.size[1] != h:
|
||||
continue
|
||||
|
||||
image = np.array(image).astype(np.float32) / 255.0
|
||||
image = torch.from_numpy(image)[None,]
|
||||
if 'A' in i.getbands():
|
||||
mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0
|
||||
mask = 1. - torch.from_numpy(mask)
|
||||
else:
|
||||
# 尺寸不对,需要按照image来
|
||||
mask = torch.zeros((h, w), dtype=torch.float32, device="cpu")
|
||||
|
||||
output_images.append(image)
|
||||
output_masks.append(mask.unsqueeze(0))
|
||||
|
||||
if len(output_images) > 1 and img.format not in excluded_formats:
|
||||
output_image = torch.cat(output_images, dim=0)
|
||||
output_mask = torch.cat(output_masks, dim=0)
|
||||
else:
|
||||
output_image = output_images[0]
|
||||
output_mask = output_masks[0]
|
||||
|
||||
return {"ui":{"images": images},"result": (output_image, output_mask)}
|
||||
|
||||
# return (output_image, output_mask)
|
||||
@@ -0,0 +1,69 @@
|
||||
import itertools
|
||||
import re
|
||||
|
||||
LANGUAGE_UNICODE_RANGE_MAP = {
|
||||
"ZH": [(0x4E00, 0x9FFF)],
|
||||
"JP": [(0x4E00, 0x9FFF), (0x3040, 0x309F), (0x30A0, 0x30FF), (0x31F0, 0x31FF)],
|
||||
"EN": [(0x0000, 0x007F)],
|
||||
}
|
||||
|
||||
SYMBOLS_MAPPING = {
|
||||
":": ",",
|
||||
";": ",",
|
||||
",": ",",
|
||||
"。": ".",
|
||||
"!": "!",
|
||||
"?": "?",
|
||||
"\n": ".",
|
||||
"·": ",",
|
||||
"、": ",",
|
||||
"...": "…",
|
||||
"“": "'",
|
||||
"”": "'",
|
||||
"‘": "'",
|
||||
"’": "'",
|
||||
"(": "'",
|
||||
")": "'",
|
||||
"(": "'",
|
||||
")": "'",
|
||||
"《": "'",
|
||||
"》": "'",
|
||||
"【": "'",
|
||||
"】": "'",
|
||||
"[": "'",
|
||||
"]": "'",
|
||||
"—": "-",
|
||||
"~": "-",
|
||||
"~": "-",
|
||||
"・": "-",
|
||||
"「": "'",
|
||||
"」": "'",
|
||||
";": ",",
|
||||
":": ",",
|
||||
}
|
||||
|
||||
REPLACE_SYMBOL_REGEX = re.compile(
|
||||
"|".join(re.escape(p) for p in SYMBOLS_MAPPING.keys())
|
||||
)
|
||||
ALL_KNOWN_UTF8_RANGE = list(
|
||||
itertools.chain.from_iterable(LANGUAGE_UNICODE_RANGE_MAP.values())
|
||||
)
|
||||
REMOVE_UNKNOWN_SYMBOL_REGEX = re.compile(
|
||||
"[^"
|
||||
+ "".join(
|
||||
f"{re.escape(chr(start))}-{re.escape(chr(end))}"
|
||||
for start, end in ALL_KNOWN_UTF8_RANGE
|
||||
)
|
||||
+ "]"
|
||||
)
|
||||
|
||||
|
||||
def clean_text(text):
|
||||
# Clean the text
|
||||
text = text.strip()
|
||||
|
||||
# Replace all chinese symbols with their english counterparts
|
||||
text = REPLACE_SYMBOL_REGEX.sub(lambda x: SYMBOLS_MAPPING[x.group()], text)
|
||||
text = REMOVE_UNKNOWN_SYMBOL_REGEX.sub("", text)
|
||||
|
||||
return text
|
||||
@@ -0,0 +1,87 @@
|
||||
# Base configuration for training a model
|
||||
paths:
|
||||
run_dir: results/${project}
|
||||
ckpt_dir: ${paths.run_dir}/checkpoints
|
||||
|
||||
hydra:
|
||||
run:
|
||||
dir: ${paths.run_dir}
|
||||
|
||||
# Lightning Trainer
|
||||
trainer:
|
||||
_target_: lightning.pytorch.trainer.Trainer
|
||||
|
||||
default_root_dir: ${paths.run_dir}
|
||||
accelerator: gpu
|
||||
num_nodes: 1
|
||||
devices: auto
|
||||
strategy:
|
||||
_target_: lightning.pytorch.strategies.DDPStrategy
|
||||
process_group_backend: nccl # This should be override when training on windows
|
||||
|
||||
precision: bf16-mixed
|
||||
|
||||
# disable validation by epoch end
|
||||
check_val_every_n_epoch: null
|
||||
val_check_interval: 5000
|
||||
max_steps: 100_000
|
||||
|
||||
# Use torch.backends.cudnn.benchmark to speed up training
|
||||
benchmark: true
|
||||
|
||||
# Callbacks
|
||||
callbacks:
|
||||
model_checkpoint:
|
||||
_target_: lightning.pytorch.callbacks.ModelCheckpoint
|
||||
dirpath: ${paths.ckpt_dir}
|
||||
filename: "step_{step:09d}"
|
||||
save_last: false # additionally always save an exact copy of the last checkpoint to a file last.ckpt
|
||||
save_top_k: 5 # save 5 latest checkpoints
|
||||
monitor: step # use step to monitor checkpoints
|
||||
mode: max # save the latest checkpoint with the highest global_step
|
||||
every_n_epochs: null # don't save checkpoints by epoch end
|
||||
every_n_train_steps: 5000 # save checkpoints every 5000 steps
|
||||
auto_insert_metric_name: false
|
||||
|
||||
model_summary:
|
||||
_target_: lightning.pytorch.callbacks.ModelSummary
|
||||
max_depth: 2 # the maximum depth of layer nesting that the summary will include
|
||||
|
||||
learning_rate_monitor:
|
||||
_target_: lightning.pytorch.callbacks.LearningRateMonitor
|
||||
logging_interval: step
|
||||
log_momentum: false
|
||||
|
||||
grad_norm_monitor:
|
||||
_target_: fish_speech.callbacks.GradNormMonitor
|
||||
norm_type: 2
|
||||
logging_interval: step
|
||||
|
||||
# Logger
|
||||
logger:
|
||||
tensorboard:
|
||||
_target_: lightning.pytorch.loggers.tensorboard.TensorBoardLogger
|
||||
save_dir: "${paths.run_dir}/tensorboard/"
|
||||
name: null
|
||||
log_graph: false
|
||||
default_hp_metric: true
|
||||
prefix: ""
|
||||
|
||||
# wandb:
|
||||
# _target_: lightning.pytorch.loggers.wandb.WandbLogger
|
||||
# # name: "" # name of the run (normally generated by wandb)
|
||||
# save_dir: "${paths.run_dir}"
|
||||
# offline: False
|
||||
# id: null # pass correct id to resume experiment!
|
||||
# anonymous: null # enable anonymous logging
|
||||
# project: "fish-speech"
|
||||
# log_model: False # upload lightning ckpts
|
||||
# prefix: "" # a string to put at the beginning of metric keys
|
||||
# # entity: "" # set to name of your wandb team
|
||||
# group: ""
|
||||
# tags: ["vq", "hq", "finetune"]
|
||||
# job_type: ""
|
||||
|
||||
# Loop
|
||||
train: true
|
||||
test: false
|
||||
@@ -0,0 +1,33 @@
|
||||
_target_: fish_speech.models.vqgan.modules.firefly.FireflyArchitecture
|
||||
spec_transform:
|
||||
_target_: fish_speech.utils.spectrogram.LogMelSpectrogram
|
||||
sample_rate: 44100
|
||||
n_mels: 160
|
||||
n_fft: 2048
|
||||
hop_length: 512
|
||||
win_length: 2048
|
||||
backbone:
|
||||
_target_: fish_speech.models.vqgan.modules.firefly.ConvNeXtEncoder
|
||||
input_channels: 160
|
||||
depths: [3, 3, 9, 3]
|
||||
dims: [128, 256, 384, 512]
|
||||
drop_path_rate: 0.2
|
||||
kernel_size: 7
|
||||
head:
|
||||
_target_: fish_speech.models.vqgan.modules.firefly.HiFiGANGenerator
|
||||
hop_length: 512
|
||||
upsample_rates: [8, 8, 2, 2, 2] # aka. strides
|
||||
upsample_kernel_sizes: [16, 16, 4, 4, 4]
|
||||
resblock_kernel_sizes: [3, 7, 11]
|
||||
resblock_dilation_sizes: [[1, 3, 5], [1, 3, 5], [1, 3, 5]]
|
||||
num_mels: 512
|
||||
upsample_initial_channel: 512
|
||||
pre_conv_kernel_size: 13
|
||||
post_conv_kernel_size: 13
|
||||
quantizer:
|
||||
_target_: fish_speech.models.vqgan.modules.fsq.DownsampleFiniteScalarQuantize
|
||||
input_dim: 512
|
||||
n_groups: 8
|
||||
n_codebooks: 1
|
||||
levels: [8, 5, 5, 5]
|
||||
downsample_factor: [2, 2]
|
||||
@@ -0,0 +1,4 @@
|
||||
_target_: fish_speech.models.text2semantic.lora.LoraConfig
|
||||
r: 8
|
||||
lora_alpha: 16
|
||||
lora_dropout: 0.01
|
||||
@@ -0,0 +1,83 @@
|
||||
defaults:
|
||||
- base
|
||||
- _self_
|
||||
|
||||
project: text2semantic_finetune_dual_ar
|
||||
max_length: 4096
|
||||
pretrained_ckpt_path: checkpoints/fish-speech-1.4
|
||||
|
||||
# Lightning Trainer
|
||||
trainer:
|
||||
accumulate_grad_batches: 1
|
||||
gradient_clip_val: 1.0
|
||||
gradient_clip_algorithm: "norm"
|
||||
max_steps: 1000
|
||||
precision: bf16-true
|
||||
limit_val_batches: 10
|
||||
val_check_interval: 100
|
||||
|
||||
# Dataset Configuration
|
||||
tokenizer:
|
||||
_target_: transformers.AutoTokenizer.from_pretrained
|
||||
pretrained_model_name_or_path: ${pretrained_ckpt_path}
|
||||
|
||||
# Dataset Configuration
|
||||
train_dataset:
|
||||
_target_: fish_speech.datasets.semantic.AutoTextSemanticInstructionDataset
|
||||
proto_files:
|
||||
- data/protos
|
||||
tokenizer: ${tokenizer}
|
||||
causal: true
|
||||
max_length: ${max_length}
|
||||
use_speaker: false
|
||||
interactive_prob: 0.7
|
||||
|
||||
val_dataset:
|
||||
_target_: fish_speech.datasets.semantic.AutoTextSemanticInstructionDataset
|
||||
proto_files:
|
||||
- data/protos
|
||||
tokenizer: ${tokenizer}
|
||||
causal: true
|
||||
max_length: ${max_length}
|
||||
use_speaker: false
|
||||
interactive_prob: 0.7
|
||||
|
||||
data:
|
||||
_target_: fish_speech.datasets.semantic.SemanticDataModule
|
||||
train_dataset: ${train_dataset}
|
||||
val_dataset: ${val_dataset}
|
||||
num_workers: 4
|
||||
batch_size: 8
|
||||
tokenizer: ${tokenizer}
|
||||
max_length: ${max_length}
|
||||
|
||||
# Model Configuration
|
||||
model:
|
||||
_target_: fish_speech.models.text2semantic.lit_module.TextToSemantic
|
||||
model:
|
||||
_target_: fish_speech.models.text2semantic.llama.BaseTransformer.from_pretrained
|
||||
path: ${pretrained_ckpt_path}
|
||||
load_weights: true
|
||||
max_length: ${max_length}
|
||||
lora_config: null
|
||||
|
||||
optimizer:
|
||||
_target_: torch.optim.AdamW
|
||||
_partial_: true
|
||||
lr: 1e-4
|
||||
weight_decay: 0
|
||||
betas: [0.9, 0.95]
|
||||
eps: 1e-5
|
||||
|
||||
lr_scheduler:
|
||||
_target_: torch.optim.lr_scheduler.LambdaLR
|
||||
_partial_: true
|
||||
lr_lambda:
|
||||
_target_: fish_speech.scheduler.get_constant_schedule_with_warmup_lr_lambda
|
||||
_partial_: true
|
||||
num_warmup_steps: 10
|
||||
|
||||
# Callbacks
|
||||
callbacks:
|
||||
model_checkpoint:
|
||||
every_n_train_steps: ${trainer.val_check_interval}
|
||||
@@ -0,0 +1,2 @@
|
||||
SEMANTIC_TOKEN = "<|semantic|>"
|
||||
CODEBOOK_PAD_TOKEN_ID = 0
|
||||
@@ -0,0 +1,53 @@
|
||||
import bisect
|
||||
import random
|
||||
from typing import Iterable
|
||||
|
||||
from torch.utils.data import Dataset, IterableDataset
|
||||
|
||||
|
||||
class ConcatRepeatDataset(Dataset):
|
||||
datasets: list[Dataset]
|
||||
cumulative_sizes: list[int]
|
||||
repeats: list[int]
|
||||
|
||||
@staticmethod
|
||||
def cumsum(sequence, repeats):
|
||||
r, s = [], 0
|
||||
for dataset, repeat in zip(sequence, repeats):
|
||||
l = len(dataset) * repeat
|
||||
r.append(l + s)
|
||||
s += l
|
||||
return r
|
||||
|
||||
def __init__(self, datasets: Iterable[Dataset], repeats: list[int]):
|
||||
super().__init__()
|
||||
|
||||
self.datasets = list(datasets)
|
||||
self.repeats = repeats
|
||||
|
||||
assert len(self.datasets) > 0, "datasets should not be an empty iterable"
|
||||
assert len(self.datasets) == len(
|
||||
repeats
|
||||
), "datasets and repeats should have the same length"
|
||||
|
||||
for d in self.datasets:
|
||||
assert not isinstance(
|
||||
d, IterableDataset
|
||||
), "ConcatRepeatDataset does not support IterableDataset"
|
||||
|
||||
self.cumulative_sizes = self.cumsum(self.datasets, self.repeats)
|
||||
|
||||
def __len__(self):
|
||||
return self.cumulative_sizes[-1]
|
||||
|
||||
def __getitem__(self, idx):
|
||||
dataset_idx = bisect.bisect_right(self.cumulative_sizes, idx)
|
||||
|
||||
if dataset_idx == 0:
|
||||
sample_idx = idx
|
||||
else:
|
||||
sample_idx = idx - self.cumulative_sizes[dataset_idx - 1]
|
||||
|
||||
dataset = self.datasets[dataset_idx]
|
||||
|
||||
return dataset[sample_idx % len(dataset)]
|
||||
@@ -0,0 +1,24 @@
|
||||
syntax = "proto3";
|
||||
|
||||
package text_data;
|
||||
|
||||
message Semantics {
|
||||
repeated uint32 values = 1;
|
||||
}
|
||||
|
||||
message Sentence {
|
||||
repeated string texts = 1;
|
||||
repeated Semantics semantics = 3;
|
||||
}
|
||||
|
||||
message TextData {
|
||||
string source = 1;
|
||||
string name = 2;
|
||||
repeated Sentence sentences = 4;
|
||||
}
|
||||
|
||||
message SampledData {
|
||||
string source = 1;
|
||||
string name = 2;
|
||||
repeated Sentence samples = 3;
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Generated by the protocol buffer compiler. DO NOT EDIT!
|
||||
# source: text-data.proto
|
||||
# Protobuf Python Version: 4.25.1
|
||||
"""Generated protocol buffer code."""
|
||||
from google.protobuf import descriptor as _descriptor
|
||||
from google.protobuf import descriptor_pool as _descriptor_pool
|
||||
from google.protobuf import symbol_database as _symbol_database
|
||||
from google.protobuf.internal import builder as _builder
|
||||
|
||||
# @@protoc_insertion_point(imports)
|
||||
|
||||
_sym_db = _symbol_database.Default()
|
||||
|
||||
|
||||
DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(
|
||||
b'\n\x0ftext-data.proto\x12\ttext_data"\x1b\n\tSemantics\x12\x0e\n\x06values\x18\x01 \x03(\r"B\n\x08Sentence\x12\r\n\x05texts\x18\x01 \x03(\t\x12\'\n\tsemantics\x18\x03 \x03(\x0b\x32\x14.text_data.Semantics"P\n\x08TextData\x12\x0e\n\x06source\x18\x01 \x01(\t\x12\x0c\n\x04name\x18\x02 \x01(\t\x12&\n\tsentences\x18\x04 \x03(\x0b\x32\x13.text_data.Sentence"Q\n\x0bSampledData\x12\x0e\n\x06source\x18\x01 \x01(\t\x12\x0c\n\x04name\x18\x02 \x01(\t\x12$\n\x07samples\x18\x03 \x03(\x0b\x32\x13.text_data.Sentenceb\x06proto3'
|
||||
)
|
||||
|
||||
_globals = globals()
|
||||
_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals)
|
||||
_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, "text_data_pb2", _globals)
|
||||
if _descriptor._USE_C_DESCRIPTORS == False:
|
||||
DESCRIPTOR._options = None
|
||||
_globals["_SEMANTICS"]._serialized_start = 30
|
||||
_globals["_SEMANTICS"]._serialized_end = 57
|
||||
_globals["_SENTENCE"]._serialized_start = 59
|
||||
_globals["_SENTENCE"]._serialized_end = 125
|
||||
_globals["_TEXTDATA"]._serialized_start = 127
|
||||
_globals["_TEXTDATA"]._serialized_end = 207
|
||||
_globals["_SAMPLEDDATA"]._serialized_start = 209
|
||||
_globals["_SAMPLEDDATA"]._serialized_end = 290
|
||||
# @@protoc_insertion_point(module_scope)
|
||||
@@ -0,0 +1,36 @@
|
||||
import struct
|
||||
|
||||
from .text_data_pb2 import TextData
|
||||
|
||||
|
||||
def read_pb_stream(f):
|
||||
while True:
|
||||
buf = f.read(4)
|
||||
if len(buf) == 0:
|
||||
break
|
||||
size = struct.unpack("I", buf)[0]
|
||||
buf = f.read(size)
|
||||
text_data = TextData()
|
||||
text_data.ParseFromString(buf)
|
||||
yield text_data
|
||||
|
||||
|
||||
def write_pb_stream(f, text_data):
|
||||
buf = text_data.SerializeToString()
|
||||
f.write(struct.pack("I", len(buf)))
|
||||
f.write(buf)
|
||||
|
||||
|
||||
def pack_pb_stream(text_data):
|
||||
buf = text_data.SerializeToString()
|
||||
return struct.pack("I", len(buf)) + buf
|
||||
|
||||
|
||||
def split_pb_stream(f):
|
||||
while True:
|
||||
head = f.read(4)
|
||||
if len(head) == 0:
|
||||
break
|
||||
size = struct.unpack("I", head)[0]
|
||||
buf = f.read(size)
|
||||
yield head + buf
|
||||
@@ -0,0 +1,496 @@
|
||||
import random
|
||||
from dataclasses import dataclass
|
||||
from itertools import chain
|
||||
from pathlib import Path
|
||||
from random import Random
|
||||
from typing import Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from datasets.download.streaming_download_manager import xopen
|
||||
from huggingface_hub import HfApi
|
||||
from lightning import LightningDataModule
|
||||
from torch.distributed import get_rank, get_world_size, is_initialized
|
||||
from torch.utils.data import DataLoader, IterableDataset, get_worker_info
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fish_speech.conversation import CODEBOOK_PAD_TOKEN_ID
|
||||
from fish_speech.datasets.protos.text_data_pb2 import SampledData
|
||||
from fish_speech.datasets.protos.text_data_stream import read_pb_stream
|
||||
from fish_speech.text.clean import clean_text
|
||||
from fish_speech.utils import RankedLogger
|
||||
from fish_speech.utils.braceexpand import braceexpand
|
||||
|
||||
log = RankedLogger(__name__, rank_zero_only=True)
|
||||
|
||||
|
||||
def split_by_rank_worker(files):
|
||||
# We need to know the total number of devices
|
||||
# to split the data properly
|
||||
|
||||
total_devices = 1
|
||||
if is_initialized():
|
||||
total_devices = get_world_size()
|
||||
|
||||
worker_info = get_worker_info()
|
||||
if worker_info is not None:
|
||||
total_devices *= worker_info.num_workers
|
||||
|
||||
if len(files) < total_devices:
|
||||
# Repeat the files N times to match the number of devices
|
||||
files = files * (total_devices // len(files) + 1)
|
||||
|
||||
# DDP
|
||||
if is_initialized():
|
||||
files = files[get_rank() :: get_world_size()]
|
||||
|
||||
# Split by worker
|
||||
if worker_info is not None:
|
||||
files = files[worker_info.id :: worker_info.num_workers]
|
||||
|
||||
return files
|
||||
|
||||
|
||||
class AutoTextSemanticInstructionDataset(IterableDataset):
|
||||
"""
|
||||
Auto Augment Dataset by Speaker
|
||||
|
||||
1. Random concatenate multiple sentences from the same speaker to form a longer sentence
|
||||
2. Automatically normalize the text
|
||||
|
||||
For interactive mode, we use the following format (multiple sequences):
|
||||
<s> [INST] [SPK: speaker] text [/INST] ... [INST] text [/INST] </s>
|
||||
|
||||
For non-interactive mode, we use the following format (one long sequence):
|
||||
<s> [INST] text [/INST] ... </s>
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
proto_files: list[str],
|
||||
seed: int = 42,
|
||||
interactive_prob: float = 0.5,
|
||||
max_length: int = 1024,
|
||||
tokenizer: AutoTokenizer = None,
|
||||
use_speaker: bool | float = True,
|
||||
causal: bool = True,
|
||||
num_codebooks: Optional[int] = None,
|
||||
skip_text_prob: float = 0.0,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
proto_files: proto buf files if using local data
|
||||
seed: random seed
|
||||
interactive_prob: probability to use interactive mode
|
||||
max_length: max length of the text
|
||||
tokenizer: tokenizer
|
||||
use_speaker: include speaker information in the prompt
|
||||
causal: use causal sampling when using local data, disable will lead to random sampling
|
||||
num_codebooks: number of codebooks, if None, it will be automatically detected
|
||||
skip_text_prob: probability to skip the text (audio only), this only applies to interactive mode
|
||||
"""
|
||||
|
||||
super().__init__()
|
||||
|
||||
assert 0 <= interactive_prob <= 1, "interactive_prob must be in [0, 1]"
|
||||
|
||||
self.seed = seed
|
||||
self.max_length = max_length
|
||||
self.tokenizer = tokenizer
|
||||
self.interactive_prob = interactive_prob
|
||||
self.use_speaker = use_speaker
|
||||
self.proto_files = proto_files
|
||||
self.causal = causal
|
||||
self.num_codebooks = num_codebooks
|
||||
self.skip_text_prob = skip_text_prob
|
||||
|
||||
self.semantic_token_id = self.tokenizer.convert_tokens_to_ids("<|semantic|>")
|
||||
self.groups = None
|
||||
|
||||
def init_mock_data_server(self):
|
||||
if self.groups is not None:
|
||||
return
|
||||
|
||||
# Expand the proto files
|
||||
expanded_proto_files = []
|
||||
for filename in self.proto_files:
|
||||
for i in braceexpand(filename):
|
||||
i = Path(i)
|
||||
if i.is_file():
|
||||
expanded_proto_files.append(i)
|
||||
elif i.is_dir():
|
||||
expanded_proto_files.extend(i.rglob("*.proto"))
|
||||
expanded_proto_files.extend(i.rglob("*.protos"))
|
||||
else:
|
||||
raise ValueError(f"{i} is not a file or directory")
|
||||
|
||||
expanded_proto_files = sorted(expanded_proto_files)
|
||||
Random(self.seed).shuffle(expanded_proto_files)
|
||||
|
||||
self.groups = []
|
||||
shard_proto_files = split_by_rank_worker(expanded_proto_files)
|
||||
log.info(
|
||||
f"Reading {len(shard_proto_files)} / {len(expanded_proto_files)} files"
|
||||
)
|
||||
|
||||
count = 0
|
||||
for filename in shard_proto_files:
|
||||
with open(filename, "rb") as f:
|
||||
for text_data in read_pb_stream(f):
|
||||
self.groups.append(text_data)
|
||||
count += 1
|
||||
|
||||
log.info(f"Read total {count} groups of data")
|
||||
|
||||
# Shuffle the lines
|
||||
Random(self.seed).shuffle(self.groups)
|
||||
self.group_weights = [len(i.sentences) for i in self.groups]
|
||||
|
||||
def __iter__(self):
|
||||
while True:
|
||||
yield self.augment()
|
||||
|
||||
def tokenize_sentence(self, sentence: str):
|
||||
sentence = clean_text(sentence)
|
||||
tokens = self.tokenizer.encode(
|
||||
f"{sentence}",
|
||||
max_length=10**6,
|
||||
add_special_tokens=False,
|
||||
truncation=False,
|
||||
)
|
||||
return sentence, len(tokens)
|
||||
|
||||
def sample_data(self):
|
||||
if self.groups is None:
|
||||
self.init_mock_data_server()
|
||||
|
||||
# Shuffle unique lines, estimate that each sample is at least 20 tokens
|
||||
num_samples = self.max_length // 20
|
||||
|
||||
# choice group based on their number of samples
|
||||
group = random.choices(self.groups, weights=self.group_weights, k=1)[0]
|
||||
|
||||
if self.causal:
|
||||
# Sample in order
|
||||
if num_samples >= len(group.sentences):
|
||||
samples = group.sentences
|
||||
else:
|
||||
begin = random.randint(0, len(group.sentences) - num_samples)
|
||||
samples = group.sentences[begin : begin + num_samples]
|
||||
else:
|
||||
samples = random.choices(
|
||||
group.sentences, k=min(num_samples, len(group.sentences))
|
||||
)
|
||||
|
||||
return SampledData(
|
||||
source=group.source,
|
||||
name=group.name,
|
||||
samples=samples,
|
||||
)
|
||||
|
||||
def augment(self):
|
||||
final_text, final_semantic = [], []
|
||||
response = self.sample_data()
|
||||
if len(response.samples) == 0:
|
||||
# Invalid group
|
||||
return None
|
||||
|
||||
samples = list(response.samples)
|
||||
idx = 0
|
||||
use_interactive = random.random() < self.interactive_prob
|
||||
|
||||
if use_interactive is False:
|
||||
# Random sample based on speaker using a truncated normal distribution
|
||||
a = torch.tensor([0], dtype=torch.float32)
|
||||
torch.nn.init.trunc_normal_(
|
||||
a,
|
||||
mean=self.max_length // 2,
|
||||
std=self.max_length // 4,
|
||||
a=10,
|
||||
b=self.max_length,
|
||||
)
|
||||
remaining_tokens = a.long().item() - 4
|
||||
else:
|
||||
remaining_tokens = self.max_length
|
||||
|
||||
# Use speaker
|
||||
if isinstance(self.use_speaker, float):
|
||||
use_speaker = random.random() < self.use_speaker
|
||||
else:
|
||||
use_speaker = self.use_speaker
|
||||
|
||||
all_tokens, all_labels = [], []
|
||||
while remaining_tokens > 0 and len(samples) > 0:
|
||||
sentence = samples.pop(0)
|
||||
|
||||
text = random.choice(sentence.texts)
|
||||
text, length = self.tokenize_sentence(text)
|
||||
remaining_tokens -= length + len(sentence.semantics[0].values)
|
||||
|
||||
if use_interactive is False:
|
||||
final_text.append(text)
|
||||
final_semantic.append(sentence.semantics)
|
||||
else:
|
||||
# For interactive mode, we only apply speaker for the first sentence
|
||||
# [INST] [SPK: speaker] text [/INST] ... [INST] text [/INST]
|
||||
tokens, labels = self.pack_sentences(
|
||||
sentences=[text],
|
||||
semantics=[sentence.semantics],
|
||||
speaker=response.name if use_speaker else None,
|
||||
skip_text=random.random() < self.skip_text_prob,
|
||||
)
|
||||
|
||||
all_tokens.append(tokens)
|
||||
all_labels.append(labels)
|
||||
|
||||
idx += 1
|
||||
|
||||
if use_interactive is False:
|
||||
tokens, labels = self.pack_sentences(
|
||||
final_text,
|
||||
semantics=final_semantic,
|
||||
speaker=response.name if use_speaker else None,
|
||||
)
|
||||
all_tokens.append(tokens)
|
||||
all_labels.append(labels)
|
||||
|
||||
tokens = torch.cat(all_tokens, dim=1)
|
||||
labels = torch.cat(all_labels, dim=1)
|
||||
|
||||
# Verify that the length is correct
|
||||
assert tokens.size(1) == labels.size(1), f"{tokens.size(1)} != {labels.size(1)}"
|
||||
|
||||
data = {"tokens": tokens, "labels": labels}
|
||||
|
||||
return data
|
||||
|
||||
def pack_sentences(
|
||||
self,
|
||||
sentences: list[str],
|
||||
semantics: list,
|
||||
speaker: Optional[str] = None,
|
||||
skip_text: bool = False,
|
||||
):
|
||||
if speaker is None:
|
||||
speaker = "assistant"
|
||||
|
||||
cated_sentences = " ".join(sentences)
|
||||
if skip_text:
|
||||
cated_sentences = "<|skip_text|>"
|
||||
|
||||
final_text = "<|im_start|>user\n" + cated_sentences + "<|im_end|>"
|
||||
final_text = final_text + f"<|im_start|>{speaker}\n"
|
||||
|
||||
encoded = self.tokenizer.encode(
|
||||
final_text,
|
||||
add_special_tokens=False,
|
||||
truncation=False,
|
||||
max_length=10**6,
|
||||
)
|
||||
semantic_length = sum([len(i[0].values) for i in semantics])
|
||||
prompt_length = len(encoded)
|
||||
num_codebooks = (
|
||||
len(semantics[0]) if self.num_codebooks is None else self.num_codebooks
|
||||
)
|
||||
|
||||
# Pack the tokens and semantics (add <s> and </s> to semantic tokens)
|
||||
tokens = (
|
||||
encoded
|
||||
+ [self.semantic_token_id] * semantic_length
|
||||
+ self.tokenizer.convert_tokens_to_ids(["<|im_end|>"])
|
||||
)
|
||||
|
||||
# Codebook bos/padding: 0, eos: 1
|
||||
codes = [[CODEBOOK_PAD_TOKEN_ID] * prompt_length for _ in range(num_codebooks)]
|
||||
for segment in semantics:
|
||||
for book_idx, book in zip(range(num_codebooks), segment):
|
||||
for j in book.values:
|
||||
codes[book_idx].append(int(j) + 1)
|
||||
|
||||
for book in codes:
|
||||
book.extend([CODEBOOK_PAD_TOKEN_ID] * 1)
|
||||
|
||||
tokens = [tokens] + codes
|
||||
|
||||
tokens = torch.tensor(tokens, dtype=torch.long)
|
||||
labels = tokens.clone()
|
||||
|
||||
if skip_text:
|
||||
# If text is not provided, the sentence is used for condition only, all labels are -100
|
||||
torch.fill_(labels, -100)
|
||||
return tokens, labels
|
||||
|
||||
# Mask out the <s> tokens for semantic, predict semantic tokens only
|
||||
# Since we don't mask out the input tokens, the language modeling still works
|
||||
labels[1:, :prompt_length] = -100
|
||||
|
||||
tokens = tokens[:, :-1]
|
||||
labels = labels[:, 1:]
|
||||
|
||||
# Verify the padding is correct, and the last token is eos
|
||||
assert (tokens[1:, :prompt_length] == CODEBOOK_PAD_TOKEN_ID).all()
|
||||
assert (labels[1:, -1:] == CODEBOOK_PAD_TOKEN_ID).all()
|
||||
|
||||
return tokens, labels
|
||||
|
||||
|
||||
@dataclass
|
||||
class TextDataCollator:
|
||||
tokenizer: AutoTokenizer
|
||||
max_length: int = 1024
|
||||
|
||||
def __call__(self, examples):
|
||||
if "negative_tokens" in examples:
|
||||
positive_examples = []
|
||||
negative_examples = []
|
||||
|
||||
for i in examples:
|
||||
positive_examples.append(
|
||||
{
|
||||
"tokens": i["tokens"],
|
||||
"labels": i["labels"],
|
||||
}
|
||||
)
|
||||
negative_examples.append(
|
||||
{
|
||||
"tokens": i["negative_tokens"],
|
||||
"labels": i["negative_labels"],
|
||||
}
|
||||
)
|
||||
|
||||
examples = positive_examples + negative_examples
|
||||
|
||||
return self.batchify(examples)
|
||||
|
||||
def batchify(self, examples, tokens_key="tokens", labels_key="labels"):
|
||||
tokens, attention_masks, labels = [], [], []
|
||||
|
||||
# Calculate the max length
|
||||
max_tokens_length = 0
|
||||
for example in examples:
|
||||
max_tokens_length = max(max_tokens_length, example[tokens_key].size(1))
|
||||
max_tokens_length = min(max_tokens_length, self.max_length)
|
||||
|
||||
for example in examples:
|
||||
_tokens = example[tokens_key][:, :max_tokens_length]
|
||||
_labels = example[labels_key][:, :max_tokens_length]
|
||||
_attention_mask = torch.ones((max_tokens_length,), dtype=torch.bool)
|
||||
tokens_length = _tokens.size(1)
|
||||
_attention_mask[:tokens_length] = False
|
||||
|
||||
assert tokens_length == _labels.size(
|
||||
1
|
||||
), f"{tokens_length} != {_labels.size(1)}"
|
||||
|
||||
if tokens_length < max_tokens_length:
|
||||
_tokens = F.pad(
|
||||
_tokens,
|
||||
(0, max_tokens_length - tokens_length),
|
||||
value=self.tokenizer.eos_token_id,
|
||||
)
|
||||
_tokens[1:, tokens_length:] = CODEBOOK_PAD_TOKEN_ID
|
||||
_labels = F.pad(
|
||||
_labels, (0, max_tokens_length - _labels.size(1)), value=-100
|
||||
)
|
||||
|
||||
tokens.append(_tokens)
|
||||
attention_masks.append(_attention_mask)
|
||||
labels.append(_labels)
|
||||
|
||||
tokens = torch.stack(tokens, dim=0)
|
||||
attention_masks = torch.stack(attention_masks, dim=0)
|
||||
labels = torch.stack(labels, dim=0)
|
||||
|
||||
return {
|
||||
"inputs": tokens,
|
||||
"attention_masks": attention_masks,
|
||||
"labels": labels,
|
||||
}
|
||||
|
||||
|
||||
class InterleaveDataset(IterableDataset):
|
||||
def __init__(
|
||||
self,
|
||||
datasets: list[IterableDataset],
|
||||
probabilities: list[float],
|
||||
seed: int = 42,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.datasets = datasets
|
||||
self.probabilities = probabilities
|
||||
self.seed = seed
|
||||
|
||||
def __iter__(self):
|
||||
rng = np.random.default_rng(self.seed)
|
||||
dataset_iterators = [iter(dataset) for dataset in self.datasets]
|
||||
|
||||
while True:
|
||||
# Random choice one
|
||||
dataset_idx = rng.choice(len(self.datasets), p=self.probabilities)
|
||||
dataset_iterator = dataset_iterators[dataset_idx]
|
||||
|
||||
try:
|
||||
yield next(dataset_iterator)
|
||||
except StopIteration:
|
||||
# Exhausted, create a new iterator
|
||||
dataset_iterators[dataset_idx] = iter(self.datasets[dataset_idx])
|
||||
yield next(dataset_iterators[dataset_idx])
|
||||
|
||||
|
||||
class SemanticDataModule(LightningDataModule):
|
||||
def __init__(
|
||||
self,
|
||||
train_dataset: Union[AutoTextSemanticInstructionDataset, InterleaveDataset],
|
||||
val_dataset: Union[AutoTextSemanticInstructionDataset, InterleaveDataset],
|
||||
batch_size: int = 32,
|
||||
tokenizer: AutoTokenizer = None,
|
||||
max_length: int = 1024,
|
||||
num_workers: int = 4,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.train_dataset = train_dataset
|
||||
self.val_dataset = val_dataset
|
||||
self.batch_size = batch_size
|
||||
self.tokenizer = tokenizer
|
||||
self.max_length = max_length
|
||||
self.num_workers = num_workers
|
||||
|
||||
def train_dataloader(self):
|
||||
return DataLoader(
|
||||
self.train_dataset,
|
||||
batch_size=self.batch_size,
|
||||
collate_fn=TextDataCollator(self.tokenizer, self.max_length),
|
||||
num_workers=self.num_workers,
|
||||
persistent_workers=True,
|
||||
)
|
||||
|
||||
def val_dataloader(self):
|
||||
return DataLoader(
|
||||
self.val_dataset,
|
||||
batch_size=self.batch_size,
|
||||
collate_fn=TextDataCollator(self.tokenizer, self.max_length),
|
||||
num_workers=self.num_workers,
|
||||
persistent_workers=True,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from tqdm import tqdm
|
||||
|
||||
ds = AutoTextSemanticInstructionDataset(
|
||||
["data/protos"],
|
||||
tokenizer=AutoTokenizer.from_pretrained("fishaudio/fish-speech-1"),
|
||||
use_speaker=False,
|
||||
interactive_prob=1.0,
|
||||
skip_text_prob=0.5,
|
||||
)
|
||||
|
||||
for i in ds:
|
||||
print(ds.tokenizer.decode(i["tokens"][0], skip_special_tokens=False))
|
||||
# i["labels"][0][i["labels"][0] == -100] = 0
|
||||
# print(ds.tokenizer.decode(i["labels"][0], skip_special_tokens=False))
|
||||
break
|
||||
@@ -0,0 +1,147 @@
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import librosa
|
||||
import numpy as np
|
||||
import torch
|
||||
from lightning import LightningDataModule
|
||||
from torch.utils.data import DataLoader, Dataset
|
||||
|
||||
from fish_speech.utils import RankedLogger
|
||||
|
||||
logger = RankedLogger(__name__, rank_zero_only=False)
|
||||
|
||||
|
||||
class VQGANDataset(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
filelist: str,
|
||||
sample_rate: int = 32000,
|
||||
hop_length: int = 640,
|
||||
slice_frames: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
filelist = Path(filelist)
|
||||
root = filelist.parent
|
||||
|
||||
self.files = [
|
||||
root / line.strip()
|
||||
for line in filelist.read_text(encoding="utf-8").splitlines()
|
||||
if line.strip()
|
||||
]
|
||||
self.sample_rate = sample_rate
|
||||
self.hop_length = hop_length
|
||||
self.slice_frames = slice_frames
|
||||
|
||||
def __len__(self):
|
||||
return len(self.files)
|
||||
|
||||
def get_item(self, idx):
|
||||
file = self.files[idx]
|
||||
|
||||
audio, _ = librosa.load(file, sr=self.sample_rate, mono=True)
|
||||
|
||||
# Slice audio and features
|
||||
if (
|
||||
self.slice_frames is not None
|
||||
and audio.shape[0] > self.slice_frames * self.hop_length
|
||||
):
|
||||
start = np.random.randint(
|
||||
0, audio.shape[0] - self.slice_frames * self.hop_length
|
||||
)
|
||||
audio = audio[start : start + self.slice_frames * self.hop_length]
|
||||
|
||||
if len(audio) == 0:
|
||||
return None
|
||||
|
||||
max_value = np.abs(audio).max()
|
||||
if max_value > 1.0:
|
||||
audio = audio / max_value
|
||||
|
||||
return {
|
||||
"audio": torch.from_numpy(audio),
|
||||
}
|
||||
|
||||
def __getitem__(self, idx):
|
||||
try:
|
||||
return self.get_item(idx)
|
||||
except Exception as e:
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
logger.error(f"Error loading {self.files[idx]}: {e}")
|
||||
return None
|
||||
|
||||
|
||||
@dataclass
|
||||
class VQGANCollator:
|
||||
def __call__(self, batch):
|
||||
batch = [x for x in batch if x is not None]
|
||||
|
||||
audio_lengths = torch.tensor([len(x["audio"]) for x in batch])
|
||||
audio_maxlen = audio_lengths.max()
|
||||
|
||||
# Rounds up to nearest multiple of 2 (audio_lengths)
|
||||
audios = []
|
||||
for x in batch:
|
||||
audios.append(
|
||||
torch.nn.functional.pad(x["audio"], (0, audio_maxlen - len(x["audio"])))
|
||||
)
|
||||
|
||||
return {
|
||||
"audios": torch.stack(audios),
|
||||
"audio_lengths": audio_lengths,
|
||||
}
|
||||
|
||||
|
||||
class VQGANDataModule(LightningDataModule):
|
||||
def __init__(
|
||||
self,
|
||||
train_dataset: VQGANDataset,
|
||||
val_dataset: VQGANDataset,
|
||||
batch_size: int = 32,
|
||||
num_workers: int = 4,
|
||||
val_batch_size: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.train_dataset = train_dataset
|
||||
self.val_dataset = val_dataset
|
||||
self.batch_size = batch_size
|
||||
self.val_batch_size = val_batch_size or batch_size
|
||||
self.num_workers = num_workers
|
||||
|
||||
def train_dataloader(self):
|
||||
return DataLoader(
|
||||
self.train_dataset,
|
||||
batch_size=self.batch_size,
|
||||
collate_fn=VQGANCollator(),
|
||||
num_workers=self.num_workers,
|
||||
shuffle=True,
|
||||
persistent_workers=True,
|
||||
)
|
||||
|
||||
def val_dataloader(self):
|
||||
return DataLoader(
|
||||
self.val_dataset,
|
||||
batch_size=self.val_batch_size,
|
||||
collate_fn=VQGANCollator(),
|
||||
num_workers=self.num_workers,
|
||||
persistent_workers=True,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
dataset = VQGANDataset("data/LibriTTS_R/vq_train_filelist.txt")
|
||||
dataloader = DataLoader(
|
||||
dataset, batch_size=4, shuffle=False, collate_fn=VQGANCollator()
|
||||
)
|
||||
|
||||
for batch in dataloader:
|
||||
print(batch["audios"].shape)
|
||||
print(batch["features"].shape)
|
||||
print(batch["audio_lengths"])
|
||||
print(batch["feature_lengths"])
|
||||
break
|
||||
@@ -0,0 +1,104 @@
|
||||
|
||||
import torch
|
||||
from .models.text2semantic.llama import BaseTransformer, NaiveTransformer, DualARTransformer
|
||||
from .tools.llama.generate import decode_one_token_ar, decode_one_token_naive, generate_long
|
||||
import numpy as np
|
||||
import time
|
||||
from typing import Union
|
||||
from loguru import logger
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
def load_model(checkpoint_path, device, precision, compile=False):
|
||||
model: Union[NaiveTransformer, DualARTransformer] = BaseTransformer.from_pretrained(
|
||||
checkpoint_path, load_weights=True
|
||||
)
|
||||
|
||||
model = model.to(device=device, dtype=precision)
|
||||
logger.info(f"Restored model from checkpoint")
|
||||
|
||||
if isinstance(model, DualARTransformer):
|
||||
decode_one_token = decode_one_token_ar
|
||||
logger.info("Using DualARTransformer")
|
||||
else:
|
||||
decode_one_token = decode_one_token_naive
|
||||
logger.info("Using NaiveTransformer")
|
||||
|
||||
if compile:
|
||||
logger.info("Compiling function...")
|
||||
decode_one_token = torch.compile(
|
||||
decode_one_token, mode="reduce-overhead", fullgraph=True
|
||||
)
|
||||
|
||||
return model.eval(), decode_one_token
|
||||
|
||||
|
||||
def prompt2semantic(
|
||||
model: DualARTransformer,
|
||||
decode_one_token: callable,
|
||||
text: str,
|
||||
prompt_text: Optional[list[str]],
|
||||
prompt_tokens: Optional[list[np.ndarray]],
|
||||
max_new_tokens: int,
|
||||
top_p: float,
|
||||
repetition_penalty: float,
|
||||
temperature: float,
|
||||
device: str,
|
||||
compile: bool,
|
||||
seed: int,
|
||||
iterative_prompt: bool,
|
||||
chunk_length: int,
|
||||
):
|
||||
|
||||
if prompt_text is not None and len(prompt_text) != len(prompt_tokens):
|
||||
raise ValueError(
|
||||
f"Number of prompt text ({len(prompt_text)}) and prompt tokens ({len(prompt_tokens)}) should be the same"
|
||||
)
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
|
||||
if prompt_tokens is not None:
|
||||
prompt_tokens = [torch.from_numpy(pt).to(device) for pt in prompt_tokens]
|
||||
|
||||
torch.manual_seed(seed)
|
||||
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed(seed)
|
||||
|
||||
generator = generate_long(
|
||||
model=model,
|
||||
device=device,
|
||||
decode_one_token=decode_one_token,
|
||||
text=text,
|
||||
num_samples=1,
|
||||
max_new_tokens=max_new_tokens,
|
||||
top_p=top_p,
|
||||
repetition_penalty=repetition_penalty,
|
||||
temperature=temperature,
|
||||
compile=compile,
|
||||
iterative_prompt=iterative_prompt,
|
||||
chunk_length=chunk_length,
|
||||
prompt_text=prompt_text,
|
||||
prompt_tokens=prompt_tokens,
|
||||
)
|
||||
|
||||
idx = 0
|
||||
all_codes = []
|
||||
codes = []
|
||||
|
||||
for response in generator:
|
||||
if response.action == "sample":
|
||||
codes.append(response.codes)
|
||||
logger.info(f"Sampled text: {response.text}")
|
||||
elif response.action == "next":
|
||||
if codes:
|
||||
all_codes.append(torch.cat(codes, dim=1).cpu().numpy())
|
||||
logger.info(f"Saved codes to codes_{idx}.npy")
|
||||
logger.info(f"Next sample")
|
||||
codes = []
|
||||
idx += 1
|
||||
else:
|
||||
logger.error(f"Error: {response}")
|
||||
|
||||
return all_codes
|
||||
@@ -0,0 +1,202 @@
|
||||
from typing import Any, Optional
|
||||
|
||||
import lightning as L
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from lightning.pytorch.utilities.types import OptimizerLRScheduler
|
||||
|
||||
import fish_speech.utils as utils
|
||||
from fish_speech.conversation import CODEBOOK_PAD_TOKEN_ID
|
||||
from fish_speech.models.text2semantic.llama import NaiveTransformer
|
||||
|
||||
log = utils.RankedLogger(__name__, rank_zero_only=True)
|
||||
|
||||
|
||||
class TextToSemantic(L.LightningModule):
|
||||
def __init__(
|
||||
self,
|
||||
model: NaiveTransformer,
|
||||
optimizer: Any,
|
||||
lr_scheduler: Any,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.model = model
|
||||
self.optimizer_builder = optimizer
|
||||
self.lr_scheduler_builder = lr_scheduler
|
||||
|
||||
def forward(self, x):
|
||||
return self.model(x)
|
||||
|
||||
def on_save_checkpoint(self, checkpoint):
|
||||
# Save only LoRA parameters
|
||||
state_dict = checkpoint["state_dict"]
|
||||
use_lora = any("lora" in name for name in state_dict.keys())
|
||||
if not use_lora:
|
||||
return
|
||||
|
||||
for name in list(state_dict.keys()):
|
||||
if "lora" not in name:
|
||||
state_dict.pop(name)
|
||||
|
||||
def configure_optimizers(self) -> OptimizerLRScheduler:
|
||||
# Get weight decay parameters
|
||||
weight_decay_parameters, other_parameters = [], []
|
||||
for name, param in self.named_parameters():
|
||||
if ".bias" in name or "norm.weight" in name or ".embeddings." in name:
|
||||
other_parameters.append(param)
|
||||
else:
|
||||
weight_decay_parameters.append(param)
|
||||
|
||||
optimizer = self.optimizer_builder(
|
||||
[
|
||||
{"params": weight_decay_parameters},
|
||||
{"params": other_parameters, "weight_decay": 0.0},
|
||||
]
|
||||
)
|
||||
|
||||
# Print the parameters and their weight decay
|
||||
for i in optimizer.param_groups:
|
||||
log.info(
|
||||
f"Set weight decay: {i['weight_decay']} for {len(i['params'])} parameters"
|
||||
)
|
||||
|
||||
lr_scheduler = self.lr_scheduler_builder(optimizer)
|
||||
|
||||
return {
|
||||
"optimizer": optimizer,
|
||||
"lr_scheduler": {
|
||||
"scheduler": lr_scheduler,
|
||||
"interval": "step",
|
||||
},
|
||||
}
|
||||
|
||||
# Copied from https://github.com/eric-mitchell/direct-preference-optimization/blob/main/trainers.py#L90
|
||||
def get_batch_logps(
|
||||
self,
|
||||
logits: torch.FloatTensor,
|
||||
labels: torch.LongTensor,
|
||||
average_log_prob: bool = False,
|
||||
) -> torch.FloatTensor:
|
||||
"""Compute the log probabilities of the given labels under the given logits.
|
||||
|
||||
Args:
|
||||
logits: Logits of the model (unnormalized). Shape: (batch_size, sequence_length, codebook_size, vocab_size)
|
||||
labels: Labels for which to compute the log probabilities. Label tokens with a value of -100 are ignored. Shape: (batch_size, sequence_length, codebook_size)
|
||||
average_log_prob: If True, return the average log probability per (non-masked) token. Otherwise, return the sum of the log probabilities of the (non-masked) tokens.
|
||||
|
||||
Returns:
|
||||
A tensor of shape (batch_size,) containing the average/sum log probabilities of the given labels under the given logits.
|
||||
"""
|
||||
assert logits.shape[:-1] == labels.shape
|
||||
|
||||
labels = labels.clone()
|
||||
loss_mask = labels != -100
|
||||
|
||||
# dummy token; we'll ignore the losses on these tokens later
|
||||
labels[labels == -100] = 0
|
||||
|
||||
per_token_logps = torch.gather(
|
||||
logits.log_softmax(-1), dim=-1, index=labels.unsqueeze(-1)
|
||||
).squeeze(-1)
|
||||
|
||||
if average_log_prob:
|
||||
return (per_token_logps * loss_mask).sum(-1) / loss_mask.sum(-1)
|
||||
else:
|
||||
return (per_token_logps * loss_mask).sum(-1)
|
||||
|
||||
def _step(self, batch, batch_idx, stage: str):
|
||||
is_train = stage == "train"
|
||||
|
||||
if is_train:
|
||||
# Key part to make lora work
|
||||
# Otherwise the parameters are merged, which lead to incorrect gradients
|
||||
self.model.train()
|
||||
|
||||
# Do positive and negative samples in the same batch to speed up training
|
||||
labels = batch["labels"]
|
||||
outputs = self.model(
|
||||
inp=batch["inputs"],
|
||||
key_padding_mask=batch["attention_masks"],
|
||||
)
|
||||
token_logits = outputs.token_logits
|
||||
codebook_logits = outputs.codebook_logits
|
||||
|
||||
# Generate labels
|
||||
base_loss = F.cross_entropy(
|
||||
token_logits.view(-1, token_logits.size(-1)),
|
||||
labels[:, 0].reshape(-1),
|
||||
ignore_index=-100,
|
||||
)
|
||||
|
||||
codebook_labels = labels[:, 1 : 1 + self.model.config.num_codebooks].mT
|
||||
semantic_loss = F.cross_entropy(
|
||||
codebook_logits.view(-1, codebook_logits.size(-1)),
|
||||
codebook_labels.reshape(-1),
|
||||
ignore_index=-100,
|
||||
)
|
||||
|
||||
loss = base_loss + semantic_loss
|
||||
|
||||
self.log(
|
||||
f"{stage}/loss",
|
||||
loss,
|
||||
on_step=is_train,
|
||||
on_epoch=not is_train,
|
||||
prog_bar=True,
|
||||
logger=True,
|
||||
sync_dist=not is_train,
|
||||
)
|
||||
|
||||
self.log(
|
||||
f"{stage}/base_loss",
|
||||
base_loss,
|
||||
on_step=is_train,
|
||||
on_epoch=not is_train,
|
||||
prog_bar=False,
|
||||
logger=True,
|
||||
sync_dist=not is_train,
|
||||
)
|
||||
|
||||
self.log(
|
||||
f"{stage}/semantic_loss",
|
||||
semantic_loss,
|
||||
on_step=is_train,
|
||||
on_epoch=not is_train,
|
||||
prog_bar=False,
|
||||
logger=True,
|
||||
sync_dist=not is_train,
|
||||
)
|
||||
|
||||
# Top-5 accuracy
|
||||
accuracy = self.get_accuracy(codebook_logits, codebook_labels)
|
||||
self.log(
|
||||
f"{stage}/top_5_accuracy",
|
||||
accuracy,
|
||||
on_step=is_train,
|
||||
on_epoch=not is_train,
|
||||
prog_bar=True,
|
||||
logger=True,
|
||||
sync_dist=not is_train,
|
||||
)
|
||||
|
||||
return loss
|
||||
|
||||
def get_accuracy(self, logits, labels):
|
||||
mask = (labels != -100) & (labels != CODEBOOK_PAD_TOKEN_ID)
|
||||
if mask.sum() == 0:
|
||||
return torch.tensor(0.0, device=logits.device)
|
||||
|
||||
_, indices = logits.topk(5, dim=-1)
|
||||
correct = indices.eq(labels.unsqueeze(-1))
|
||||
correct[~mask] = 0
|
||||
correct = correct.sum()
|
||||
accuracy = correct / mask.sum()
|
||||
|
||||
return accuracy
|
||||
|
||||
def training_step(self, batch, batch_idx):
|
||||
return self._step(batch, batch_idx, "train")
|
||||
|
||||
def validation_step(self, batch, batch_idx):
|
||||
return self._step(batch, batch_idx, "val")
|
||||
@@ -0,0 +1,779 @@
|
||||
import json
|
||||
import math
|
||||
from collections import OrderedDict
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
from loguru import logger
|
||||
from torch import Tensor
|
||||
from torch.nn import functional as F
|
||||
from torch.nn.attention import SDPBackend, sdpa_kernel
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fish_speech.conversation import SEMANTIC_TOKEN
|
||||
from fish_speech.utils import RankedLogger
|
||||
|
||||
from .lora import LoraConfig, setup_lora
|
||||
|
||||
log = RankedLogger(__name__, rank_zero_only=True)
|
||||
|
||||
|
||||
def find_multiple(n: int, k: int) -> int:
|
||||
if n % k == 0:
|
||||
return n
|
||||
return n + k - (n % k)
|
||||
|
||||
|
||||
@dataclass
|
||||
class BaseModelArgs:
|
||||
model_type: str = "base"
|
||||
|
||||
vocab_size: int = 32000
|
||||
n_layer: int = 32
|
||||
n_head: int = 32
|
||||
dim: int = 4096
|
||||
intermediate_size: int = None
|
||||
n_local_heads: int = -1
|
||||
head_dim: int = 64
|
||||
rope_base: float = 10000
|
||||
norm_eps: float = 1e-5
|
||||
max_seq_len: int = 2048
|
||||
dropout: float = 0.0
|
||||
tie_word_embeddings: bool = True
|
||||
attention_qkv_bias: bool = False
|
||||
|
||||
# Codebook configs
|
||||
codebook_size: int = 160
|
||||
num_codebooks: int = 4
|
||||
|
||||
# Gradient checkpointing
|
||||
use_gradient_checkpointing: bool = True
|
||||
|
||||
# Initialize the model
|
||||
initializer_range: float = 0.02
|
||||
|
||||
def __post_init__(self):
|
||||
if self.n_local_heads == -1:
|
||||
self.n_local_heads = self.n_head
|
||||
if self.intermediate_size is None:
|
||||
hidden_dim = 4 * self.dim
|
||||
n_hidden = int(2 * hidden_dim / 3)
|
||||
self.intermediate_size = find_multiple(n_hidden, 256)
|
||||
self.head_dim = self.dim // self.n_head
|
||||
|
||||
@staticmethod
|
||||
def from_pretrained(path: str):
|
||||
path = Path(path)
|
||||
|
||||
if path.is_dir():
|
||||
path = path / "config.json"
|
||||
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
|
||||
match data["model_type"]:
|
||||
case "naive":
|
||||
cls = NaiveModelArgs
|
||||
case "dual_ar":
|
||||
cls = DualARModelArgs
|
||||
case _:
|
||||
raise ValueError(f"Unknown model type: {data['model_type']}")
|
||||
|
||||
return cls(**data)
|
||||
|
||||
def save(self, path: str):
|
||||
with open(path, "w") as f:
|
||||
json.dump(self.__dict__, f, indent=4, sort_keys=True, ensure_ascii=False)
|
||||
|
||||
|
||||
@dataclass
|
||||
class NaiveModelArgs(BaseModelArgs):
|
||||
model_type: str = "naive"
|
||||
|
||||
|
||||
@dataclass
|
||||
class DualARModelArgs(BaseModelArgs):
|
||||
model_type: str = "dual_ar"
|
||||
n_fast_layer: int = 4
|
||||
|
||||
|
||||
class KVCache(nn.Module):
|
||||
def __init__(
|
||||
self, max_batch_size, max_seq_len, n_heads, head_dim, dtype=torch.bfloat16
|
||||
):
|
||||
super().__init__()
|
||||
cache_shape = (max_batch_size, n_heads, max_seq_len, head_dim)
|
||||
self.register_buffer("k_cache", torch.zeros(cache_shape, dtype=dtype))
|
||||
self.register_buffer("v_cache", torch.zeros(cache_shape, dtype=dtype))
|
||||
|
||||
def update(self, input_pos, k_val, v_val):
|
||||
# input_pos: [S], k_val: [B, H, S, D]
|
||||
assert input_pos.shape[0] == k_val.shape[2]
|
||||
|
||||
k_out = self.k_cache
|
||||
v_out = self.v_cache
|
||||
k_out[:, :, input_pos] = k_val
|
||||
v_out[:, :, input_pos] = v_val
|
||||
|
||||
return k_out, v_out
|
||||
|
||||
|
||||
@dataclass
|
||||
class TransformerForwardResult:
|
||||
token_logits: Tensor
|
||||
codebook_logits: Tensor
|
||||
|
||||
|
||||
@dataclass
|
||||
class BaseTransformerForwardResult:
|
||||
logits: Tensor
|
||||
hidden_states: Tensor
|
||||
|
||||
|
||||
class BaseTransformer(nn.Module):
|
||||
def __init__(
|
||||
self, config: BaseModelArgs, tokenizer: AutoTokenizer, init_weights: bool = True
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.semantic_token_id = tokenizer.convert_tokens_to_ids(SEMANTIC_TOKEN)
|
||||
|
||||
# Slow transformer
|
||||
self.embeddings = nn.Embedding(
|
||||
config.vocab_size,
|
||||
config.dim,
|
||||
)
|
||||
self.codebook_embeddings = nn.Embedding(
|
||||
config.codebook_size * config.num_codebooks,
|
||||
config.dim,
|
||||
)
|
||||
self.layers = nn.ModuleList(
|
||||
TransformerBlock(config, use_sdpa=True) for _ in range(config.n_layer)
|
||||
)
|
||||
self.norm = RMSNorm(config.dim, eps=config.norm_eps)
|
||||
|
||||
if self.config.tie_word_embeddings is False:
|
||||
self.output = nn.Linear(
|
||||
config.dim,
|
||||
config.vocab_size,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
self.register_buffer(
|
||||
"freqs_cis",
|
||||
precompute_freqs_cis(
|
||||
config.max_seq_len,
|
||||
config.dim // config.n_head,
|
||||
config.rope_base,
|
||||
),
|
||||
persistent=False,
|
||||
)
|
||||
self.register_buffer(
|
||||
"causal_mask",
|
||||
torch.tril(
|
||||
torch.ones(
|
||||
config.max_seq_len,
|
||||
config.max_seq_len,
|
||||
dtype=torch.bool,
|
||||
)
|
||||
),
|
||||
persistent=False,
|
||||
)
|
||||
|
||||
# For kv cache
|
||||
self.max_batch_size = -1
|
||||
self.max_seq_len = -1
|
||||
|
||||
if init_weights:
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def setup_caches(
|
||||
self, max_batch_size: int, max_seq_len: int, dtype: torch.dtype = torch.bfloat16
|
||||
):
|
||||
if self.max_seq_len >= max_seq_len and self.max_batch_size >= max_batch_size:
|
||||
return
|
||||
|
||||
head_dim = self.config.dim // self.config.n_head
|
||||
max_seq_len = find_multiple(max_seq_len, 8)
|
||||
self.max_seq_len = max_seq_len
|
||||
self.max_batch_size = max_batch_size
|
||||
|
||||
for b in self.layers:
|
||||
b.attention.kv_cache = KVCache(
|
||||
max_batch_size,
|
||||
max_seq_len,
|
||||
self.config.n_local_heads,
|
||||
head_dim,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
def embed(self, x: Tensor) -> Tensor:
|
||||
vocab_embeds = [self.embeddings(x[:, 0])]
|
||||
for i in range(self.config.num_codebooks):
|
||||
emb = self.codebook_embeddings(x[:, i + 1] + i * self.config.codebook_size)
|
||||
emb[x[:, 0] != self.semantic_token_id] = 0
|
||||
vocab_embeds.append(emb)
|
||||
|
||||
x = torch.stack(vocab_embeds, dim=3)
|
||||
x = x.sum(dim=3)
|
||||
|
||||
return x
|
||||
|
||||
def forward(
|
||||
self,
|
||||
inp: Tensor,
|
||||
key_padding_mask: Optional[Tensor] = None,
|
||||
) -> BaseTransformerForwardResult:
|
||||
seq_len = inp.size(2)
|
||||
|
||||
# Here we want to merge the embeddings of the codebooks
|
||||
x = self.embed(inp)
|
||||
|
||||
freqs_cis = self.freqs_cis[:seq_len]
|
||||
|
||||
# Not that the causal mask here follows the definition of scaled_dot_product_attention
|
||||
# That is, FALSE means masked out
|
||||
# To maintain consistency, key_padding_mask use TRUE to mask out
|
||||
mask = None
|
||||
if key_padding_mask is not None:
|
||||
mask = self.causal_mask[None, None, :seq_len, :seq_len] # (B, N, Q, K)
|
||||
mask = mask & key_padding_mask[:, None, None, :].logical_not()
|
||||
|
||||
for layer in self.layers:
|
||||
if self.config.use_gradient_checkpointing and self.training:
|
||||
x = checkpoint(layer, x, freqs_cis, mask, use_reentrant=True)
|
||||
else:
|
||||
x = layer(x, freqs_cis, mask)
|
||||
|
||||
# We got slow_out here
|
||||
slow_out = self.norm(x)
|
||||
|
||||
if self.config.tie_word_embeddings:
|
||||
token_logits = F.linear(slow_out, self.embeddings.weight)
|
||||
else:
|
||||
token_logits = self.output(slow_out)
|
||||
|
||||
return BaseTransformerForwardResult(
|
||||
logits=token_logits,
|
||||
hidden_states=x,
|
||||
)
|
||||
|
||||
def forward_generate(
|
||||
self,
|
||||
x: Tensor,
|
||||
input_pos: Optional[Tensor] = None,
|
||||
return_all: bool = False,
|
||||
) -> BaseTransformerForwardResult:
|
||||
# This is used for generation, optimized for torch compile
|
||||
assert (
|
||||
self.max_seq_len != -1 and self.max_batch_size != -1
|
||||
), "Please call setup_caches before forward_generate"
|
||||
|
||||
x = self.embed(x)
|
||||
|
||||
mask = self.causal_mask[
|
||||
None, None, input_pos, : self.max_seq_len
|
||||
] # (B, N, Q, K)
|
||||
freqs_cis = self.freqs_cis[input_pos]
|
||||
|
||||
for layer in self.layers:
|
||||
x = layer(x, freqs_cis, mask, input_pos=input_pos)
|
||||
|
||||
# If prefill, we only calculate the logits of last token
|
||||
if x.size(1) > 1 and not return_all:
|
||||
x = x[:, -1:]
|
||||
|
||||
# We got slow_out here
|
||||
slow_out = self.norm(x)
|
||||
|
||||
if self.config.tie_word_embeddings:
|
||||
token_logits = F.linear(slow_out, self.embeddings.weight)
|
||||
else:
|
||||
token_logits = self.output(slow_out)
|
||||
|
||||
return BaseTransformerForwardResult(
|
||||
logits=token_logits,
|
||||
hidden_states=x,
|
||||
)
|
||||
|
||||
def _init_weights(self, module):
|
||||
std = self.config.initializer_range
|
||||
if isinstance(module, nn.Linear):
|
||||
module.weight.data.normal_(mean=0.0, std=std)
|
||||
if module.bias is not None:
|
||||
module.bias.data.zero_()
|
||||
elif isinstance(module, nn.Embedding):
|
||||
module.weight.data.normal_(mean=0.0, std=std)
|
||||
if module.padding_idx is not None:
|
||||
module.weight.data[module.padding_idx].zero_()
|
||||
|
||||
@staticmethod
|
||||
def from_pretrained(
|
||||
path: str,
|
||||
load_weights: bool = False,
|
||||
max_length: int | None = None,
|
||||
lora_config: LoraConfig | None = None,
|
||||
rope_base: int | None = None,
|
||||
) -> "BaseTransformer":
|
||||
config = BaseModelArgs.from_pretrained(str(path))
|
||||
if max_length is not None:
|
||||
config.max_seq_len = max_length
|
||||
log.info(f"Override max_seq_len to {max_length}")
|
||||
|
||||
if rope_base is not None:
|
||||
config.rope_base = rope_base
|
||||
log.info(f"Override rope_base to {rope_base}")
|
||||
|
||||
match config.model_type:
|
||||
case "naive":
|
||||
model_cls = NaiveTransformer
|
||||
case "dual_ar":
|
||||
model_cls = DualARTransformer
|
||||
case _:
|
||||
raise ValueError(f"Unknown model type: {config.model_type}")
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(str(path))
|
||||
log.info(f"Loading model from {path}, config: {config}")
|
||||
model = model_cls(config, tokenizer=tokenizer)
|
||||
|
||||
if lora_config is not None:
|
||||
setup_lora(model, lora_config)
|
||||
log.info(f"LoRA setup: {lora_config}")
|
||||
|
||||
if load_weights is False:
|
||||
log.info("Randomly initialized model")
|
||||
else:
|
||||
|
||||
if "int8" in str(Path(path)):
|
||||
logger.info("Using int8 weight-only quantization!")
|
||||
from tools.llama.quantize import WeightOnlyInt8QuantHandler
|
||||
|
||||
simple_quantizer = WeightOnlyInt8QuantHandler(model)
|
||||
model = simple_quantizer.convert_for_runtime()
|
||||
|
||||
if "int4" in str(Path(path)):
|
||||
logger.info("Using int4 quantization!")
|
||||
path_comps = path.name.split("-")
|
||||
assert path_comps[-2].startswith("g")
|
||||
groupsize = int(path_comps[-2][1:])
|
||||
from tools.llama.quantize import WeightOnlyInt4QuantHandler
|
||||
|
||||
simple_quantizer = WeightOnlyInt4QuantHandler(model, groupsize)
|
||||
model = simple_quantizer.convert_for_runtime()
|
||||
|
||||
weights = torch.load(
|
||||
Path(path) / "model.pth", map_location="cpu", mmap=True
|
||||
)
|
||||
|
||||
if "state_dict" in weights:
|
||||
logger.warning(
|
||||
"Using a TextToSemantic LightningModule checkpoint, "
|
||||
"please make sure it is a full model, not a LoRA model."
|
||||
)
|
||||
weights = weights["state_dict"]
|
||||
|
||||
if next(iter(weights.keys())).startswith("model."):
|
||||
logger.info(
|
||||
f"Remove prefix 'model.' created by TextToSemantic LightningModule from keys"
|
||||
)
|
||||
new_weights = OrderedDict()
|
||||
for k, v in weights.items():
|
||||
new_weights[k.replace("model.", "")] = v
|
||||
weights = new_weights
|
||||
|
||||
# Verify the name and shape of parameters since strict=False in load_state_dict.
|
||||
for k, v in model.named_parameters():
|
||||
if k not in weights:
|
||||
logger.warning(f"No weight for {k}")
|
||||
elif v.shape != weights[k].shape:
|
||||
logger.warning(
|
||||
f"Shape mismatch for {k}: {v.shape} vs {weights[k].shape}"
|
||||
)
|
||||
|
||||
err = model.load_state_dict(weights, strict=False, assign=True)
|
||||
log.info(f"Loaded weights with error: {err}")
|
||||
|
||||
return model
|
||||
|
||||
def save_pretrained(self, path: str, drop_lora: bool = False):
|
||||
path = Path(path)
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
self.config.save(path / "config.json")
|
||||
state_dict = self.state_dict()
|
||||
|
||||
if drop_lora:
|
||||
for key in list(state_dict.keys()):
|
||||
if "lora" not in key:
|
||||
continue
|
||||
|
||||
state_dict.pop(key)
|
||||
log.info(f"Drop LoRA parameter: {key}")
|
||||
|
||||
torch.save(state_dict, path / "model.pth")
|
||||
self.tokenizer.save_pretrained(path)
|
||||
|
||||
|
||||
class NaiveTransformer(BaseTransformer):
|
||||
def __init__(self, config: NaiveModelArgs, tokenizer: AutoTokenizer) -> None:
|
||||
super().__init__(config, init_weights=False, tokenizer=tokenizer)
|
||||
|
||||
self.codebook_norm = RMSNorm(config.dim, eps=config.norm_eps)
|
||||
self.codebook_output = nn.Linear(
|
||||
config.dim,
|
||||
config.codebook_size * config.num_codebooks,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def decode(self, result: BaseTransformerForwardResult) -> TransformerForwardResult:
|
||||
token_logits = result.logits
|
||||
x = result.hidden_states
|
||||
|
||||
# Codebook
|
||||
codebook_logits = self.codebook_output(self.codebook_norm(x))
|
||||
codebook_logits = rearrange(
|
||||
codebook_logits, "b n (c d) -> b n c d", c=self.config.num_codebooks
|
||||
)
|
||||
|
||||
return TransformerForwardResult(
|
||||
token_logits=token_logits,
|
||||
codebook_logits=codebook_logits,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
inp: Tensor,
|
||||
key_padding_mask: Optional[Tensor] = None,
|
||||
) -> TransformerForwardResult:
|
||||
result = super().forward(
|
||||
inp=inp,
|
||||
key_padding_mask=key_padding_mask,
|
||||
)
|
||||
return self.decode(result)
|
||||
|
||||
def forward_generate(
|
||||
self, x: Tensor, input_pos: Optional[Tensor] = None
|
||||
) -> TransformerForwardResult:
|
||||
result = super().forward_generate(x, input_pos)
|
||||
return self.decode(result)
|
||||
|
||||
|
||||
class DualARTransformer(BaseTransformer):
|
||||
def __init__(self, config: NaiveModelArgs, tokenizer: AutoTokenizer) -> None:
|
||||
super().__init__(config, init_weights=False, tokenizer=tokenizer)
|
||||
|
||||
# Fast transformer
|
||||
self.fast_embeddings = nn.Embedding(config.codebook_size, config.dim)
|
||||
|
||||
# The equivalent bs is so large that sdpa doesn't work
|
||||
self.fast_layers = nn.ModuleList(
|
||||
TransformerBlock(config, use_sdpa=False) for _ in range(config.n_fast_layer)
|
||||
)
|
||||
self.fast_norm = RMSNorm(config.dim, eps=config.norm_eps)
|
||||
self.fast_output = nn.Linear(
|
||||
config.dim,
|
||||
config.codebook_size,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def setup_caches(
|
||||
self, max_batch_size: int, max_seq_len: int, dtype: torch.dtype = torch.bfloat16
|
||||
):
|
||||
super().setup_caches(max_batch_size, max_seq_len, dtype)
|
||||
|
||||
head_dim = self.config.dim // self.config.n_head
|
||||
|
||||
# Fast transformer
|
||||
# The max seq len here is the number of codebooks
|
||||
for b in self.fast_layers:
|
||||
b.attention.kv_cache = KVCache(
|
||||
max_batch_size,
|
||||
self.config.num_codebooks,
|
||||
self.config.n_local_heads,
|
||||
head_dim,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
inp: Tensor,
|
||||
key_padding_mask: Optional[Tensor] = None,
|
||||
) -> TransformerForwardResult:
|
||||
parent_result = super().forward(inp, key_padding_mask)
|
||||
token_logits = parent_result.logits
|
||||
x = parent_result.hidden_states
|
||||
|
||||
# Fast transformer
|
||||
fast_seq_len = self.config.num_codebooks
|
||||
fast_mask = self.causal_mask[
|
||||
None, None, :fast_seq_len, :fast_seq_len
|
||||
] # (B, N, Q, K)
|
||||
fast_freqs_cis = self.freqs_cis[:fast_seq_len]
|
||||
|
||||
# Drop the last token and rotate left
|
||||
codebooks = inp[:, 1:-1, 1:]
|
||||
codebooks = F.pad(codebooks, (0, 1), value=0)
|
||||
codebook_embeddings = self.fast_embeddings(codebooks)
|
||||
x = torch.cat([x[:, None], codebook_embeddings], dim=1)
|
||||
b, s = x.size(0), x.size(2)
|
||||
x = rearrange(x, "b n s d -> (b s) n d") # flatten the batch and seq_len
|
||||
|
||||
# Remove padded part
|
||||
codebooks = rearrange(codebooks, "b n s -> (b s) n")
|
||||
codebook_mask = (codebooks == 0).all(dim=-1)
|
||||
|
||||
if torch.all(codebook_mask):
|
||||
# If all codebooks are padded, we keep first 8 to make sure the model runs
|
||||
codebook_mask[:8] = False
|
||||
|
||||
x_bs, x_len = x.size(0), x.size(1)
|
||||
x = x[~codebook_mask]
|
||||
|
||||
for layer in self.fast_layers:
|
||||
if self.config.use_gradient_checkpointing and self.training:
|
||||
x = checkpoint(layer, x, fast_freqs_cis, fast_mask, use_reentrant=True)
|
||||
else:
|
||||
x = layer(x, fast_freqs_cis, fast_mask)
|
||||
|
||||
# unflatten the batch and num_codebooks
|
||||
fast_out = self.fast_norm(x)
|
||||
codebook_logits = self.fast_output(fast_out)
|
||||
|
||||
# Re-pad the codebook_logits
|
||||
buffer = torch.zeros(
|
||||
x_bs,
|
||||
x_len,
|
||||
codebook_logits.size(-1),
|
||||
device=codebook_logits.device,
|
||||
dtype=codebook_logits.dtype,
|
||||
)
|
||||
buffer[~codebook_mask] = codebook_logits
|
||||
codebook_logits = buffer
|
||||
|
||||
assert codebook_logits.shape[1] == self.config.num_codebooks
|
||||
codebook_logits = rearrange(
|
||||
codebook_logits,
|
||||
"(b s) n d -> b s n d",
|
||||
b=b,
|
||||
s=s,
|
||||
n=self.config.num_codebooks,
|
||||
)
|
||||
|
||||
return TransformerForwardResult(
|
||||
token_logits=token_logits,
|
||||
codebook_logits=codebook_logits,
|
||||
)
|
||||
|
||||
def forward_generate_fast(
|
||||
self, x: Tensor, input_pos: Optional[Tensor] = None
|
||||
) -> Tensor:
|
||||
# Fast transformer
|
||||
x = x.view(1, 1, -1)
|
||||
|
||||
fast_mask = self.causal_mask[
|
||||
None, None, input_pos, : self.config.num_codebooks
|
||||
] # (B, N, Q, K)
|
||||
fast_freqs_cis = self.freqs_cis[input_pos]
|
||||
|
||||
for layer in self.fast_layers:
|
||||
x = layer(x, fast_freqs_cis, fast_mask, input_pos=input_pos)
|
||||
|
||||
# unflatten the batch and num_codebooks
|
||||
fast_out = self.fast_norm(x) # only take the last token
|
||||
codebook_logits = self.fast_output(fast_out)
|
||||
|
||||
return codebook_logits
|
||||
|
||||
|
||||
class TransformerBlock(nn.Module):
|
||||
def __init__(self, config: BaseModelArgs, use_sdpa: bool = True) -> None:
|
||||
super().__init__()
|
||||
self.attention = Attention(config, use_sdpa=use_sdpa)
|
||||
self.feed_forward = FeedForward(config)
|
||||
self.ffn_norm = RMSNorm(config.dim, config.norm_eps)
|
||||
self.attention_norm = RMSNorm(config.dim, config.norm_eps)
|
||||
|
||||
def forward(
|
||||
self, x: Tensor, freqs_cis: Tensor, mask: Tensor, input_pos: Tensor = None
|
||||
) -> Tensor:
|
||||
h = x + self.attention(self.attention_norm(x), freqs_cis, mask, input_pos)
|
||||
out = h + self.feed_forward(self.ffn_norm(h))
|
||||
return out
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(self, config: BaseModelArgs, use_sdpa: bool = True):
|
||||
super().__init__()
|
||||
assert config.dim % config.n_head == 0
|
||||
|
||||
total_head_dim = (config.n_head + 2 * config.n_local_heads) * config.head_dim
|
||||
# key, query, value projections for all heads, but in a batch
|
||||
self.wqkv = nn.Linear(
|
||||
config.dim, total_head_dim, bias=config.attention_qkv_bias
|
||||
)
|
||||
self.wo = nn.Linear(config.dim, config.dim, bias=False)
|
||||
self.kv_cache = None
|
||||
|
||||
self.dropout = config.dropout
|
||||
self.n_head = config.n_head
|
||||
self.head_dim = config.head_dim
|
||||
self.n_local_heads = config.n_local_heads
|
||||
self.dim = config.dim
|
||||
self.use_sdpa = use_sdpa
|
||||
self._register_load_state_dict_pre_hook(self.load_hook)
|
||||
|
||||
def load_hook(self, state_dict, prefix, *args):
|
||||
if prefix + "wq.weight" in state_dict:
|
||||
wq = state_dict.pop(prefix + "wq.weight")
|
||||
wk = state_dict.pop(prefix + "wk.weight")
|
||||
wv = state_dict.pop(prefix + "wv.weight")
|
||||
state_dict[prefix + "wqkv.weight"] = torch.cat([wq, wk, wv])
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: Tensor,
|
||||
freqs_cis: Tensor,
|
||||
mask: Tensor,
|
||||
input_pos: Optional[Tensor] = None,
|
||||
) -> Tensor:
|
||||
bsz, seqlen, _ = x.shape
|
||||
|
||||
kv_size = self.n_local_heads * self.head_dim
|
||||
q, k, v = self.wqkv(x).split([self.dim, kv_size, kv_size], dim=-1)
|
||||
|
||||
q = q.view(bsz, seqlen, self.n_head, self.head_dim)
|
||||
k = k.view(bsz, seqlen, self.n_local_heads, self.head_dim)
|
||||
v = v.view(bsz, seqlen, self.n_local_heads, self.head_dim)
|
||||
|
||||
q = apply_rotary_emb(q, freqs_cis)
|
||||
k = apply_rotary_emb(k, freqs_cis)
|
||||
|
||||
q, k, v = map(lambda x: x.transpose(1, 2), (q, k, v))
|
||||
|
||||
if self.kv_cache is not None:
|
||||
k, v = self.kv_cache.update(input_pos, k, v)
|
||||
|
||||
k = k.repeat_interleave(self.n_head // self.n_local_heads, dim=1)
|
||||
v = v.repeat_interleave(self.n_head // self.n_local_heads, dim=1)
|
||||
|
||||
if self.use_sdpa:
|
||||
if mask is None:
|
||||
with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
|
||||
y = F.scaled_dot_product_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
dropout_p=self.dropout if self.training else 0.0,
|
||||
is_causal=True,
|
||||
# No third party attn_mask here to use flash_attention
|
||||
)
|
||||
else:
|
||||
y = F.scaled_dot_product_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
attn_mask=mask,
|
||||
dropout_p=self.dropout if self.training else 0.0,
|
||||
)
|
||||
else:
|
||||
y = self.eq_scaled_dot_product_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
attn_mask=mask,
|
||||
dropout_p=self.dropout if self.training else 0.0,
|
||||
)
|
||||
|
||||
y = y.transpose(1, 2).contiguous().view(bsz, seqlen, self.dim)
|
||||
|
||||
return self.wo(y)
|
||||
|
||||
def eq_scaled_dot_product_attention(
|
||||
self,
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
attn_mask=None,
|
||||
dropout_p=0.0,
|
||||
) -> torch.Tensor:
|
||||
# This is a standard scaled dot product attention
|
||||
# It's low efficient, but it doesn't raise cuda error
|
||||
|
||||
L, S = query.size(-2), key.size(-2)
|
||||
scale_factor = 1 / math.sqrt(query.size(-1))
|
||||
attn_bias = torch.zeros(1, 1, L, S, dtype=query.dtype, device=query.device)
|
||||
|
||||
if attn_mask is not None:
|
||||
if attn_mask.dtype == torch.bool:
|
||||
attn_bias.masked_fill_(attn_mask.logical_not(), float("-inf"))
|
||||
else:
|
||||
attn_bias += attn_mask
|
||||
|
||||
attn_weight = query @ key.transpose(-2, -1) * scale_factor
|
||||
attn_weight += attn_bias
|
||||
attn_weight = torch.softmax(attn_weight, dim=-1)
|
||||
attn_weight = torch.dropout(attn_weight, dropout_p, train=True)
|
||||
|
||||
return attn_weight @ value
|
||||
|
||||
|
||||
class FeedForward(nn.Module):
|
||||
def __init__(self, config: BaseModelArgs) -> None:
|
||||
super().__init__()
|
||||
self.w1 = nn.Linear(config.dim, config.intermediate_size, bias=False)
|
||||
self.w3 = nn.Linear(config.dim, config.intermediate_size, bias=False)
|
||||
self.w2 = nn.Linear(config.intermediate_size, config.dim, bias=False)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return self.w2(F.silu(self.w1(x)) * self.w3(x))
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
def __init__(self, dim: int, eps: float = 1e-5):
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def _norm(self, x):
|
||||
return x * torch.rsqrt(torch.mean(x * x, dim=-1, keepdim=True) + self.eps)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
output = self._norm(x.float()).type_as(x)
|
||||
return output * self.weight
|
||||
|
||||
|
||||
def precompute_freqs_cis(seq_len: int, n_elem: int, base: int = 10000) -> Tensor:
|
||||
freqs = 1.0 / (
|
||||
base ** (torch.arange(0, n_elem, 2)[: (n_elem // 2)].float() / n_elem)
|
||||
)
|
||||
t = torch.arange(seq_len, device=freqs.device)
|
||||
freqs = torch.outer(t, freqs)
|
||||
freqs_cis = torch.polar(torch.ones_like(freqs), freqs)
|
||||
cache = torch.stack([freqs_cis.real, freqs_cis.imag], dim=-1)
|
||||
return cache.to(dtype=torch.bfloat16)
|
||||
|
||||
|
||||
def apply_rotary_emb(x: Tensor, freqs_cis: Tensor) -> Tensor:
|
||||
xshaped = x.float().reshape(*x.shape[:-1], -1, 2)
|
||||
freqs_cis = freqs_cis.view(1, xshaped.size(1), 1, xshaped.size(3), 2)
|
||||
x_out2 = torch.stack(
|
||||
[
|
||||
xshaped[..., 0] * freqs_cis[..., 0] - xshaped[..., 1] * freqs_cis[..., 1],
|
||||
xshaped[..., 1] * freqs_cis[..., 0] + xshaped[..., 0] * freqs_cis[..., 1],
|
||||
],
|
||||
-1,
|
||||
)
|
||||
|
||||
x_out2 = x_out2.flatten(3)
|
||||
return x_out2.type_as(x)
|
||||
@@ -0,0 +1,92 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
import loralib as lora
|
||||
|
||||
|
||||
@dataclass
|
||||
class LoraConfig:
|
||||
r: int
|
||||
lora_alpha: float
|
||||
lora_dropout: float = 0.0
|
||||
|
||||
|
||||
def setup_lora(model, lora_config):
|
||||
# Replace the embedding layer with a LoRA layer
|
||||
model.embeddings = lora.Embedding(
|
||||
num_embeddings=model.embeddings.num_embeddings,
|
||||
embedding_dim=model.embeddings.embedding_dim,
|
||||
padding_idx=model.embeddings.padding_idx,
|
||||
r=lora_config.r,
|
||||
lora_alpha=lora_config.lora_alpha,
|
||||
)
|
||||
|
||||
model.codebook_embeddings = lora.Embedding(
|
||||
num_embeddings=model.codebook_embeddings.num_embeddings,
|
||||
embedding_dim=model.codebook_embeddings.embedding_dim,
|
||||
padding_idx=model.codebook_embeddings.padding_idx,
|
||||
r=lora_config.r,
|
||||
lora_alpha=lora_config.lora_alpha,
|
||||
)
|
||||
|
||||
# Replace output layer with a LoRA layer
|
||||
linears = [(model, "output")]
|
||||
|
||||
# Replace all linear layers with LoRA layers
|
||||
for layer in model.layers:
|
||||
linears.extend([(layer.attention, "wqkv"), (layer.attention, "wo")])
|
||||
linears.extend(
|
||||
[
|
||||
(layer.feed_forward, "w1"),
|
||||
(layer.feed_forward, "w2"),
|
||||
(layer.feed_forward, "w3"),
|
||||
]
|
||||
)
|
||||
|
||||
if hasattr(model, "fast_layers"):
|
||||
model.fast_embeddings = lora.Embedding(
|
||||
num_embeddings=model.fast_embeddings.num_embeddings,
|
||||
embedding_dim=model.fast_embeddings.embedding_dim,
|
||||
padding_idx=model.fast_embeddings.padding_idx,
|
||||
r=lora_config.r,
|
||||
lora_alpha=lora_config.lora_alpha,
|
||||
)
|
||||
|
||||
# Dual-AR model
|
||||
linears.append((model, "fast_output"))
|
||||
|
||||
for layer in model.fast_layers:
|
||||
linears.extend([(layer.attention, "wqkv"), (layer.attention, "wo")])
|
||||
linears.extend(
|
||||
[
|
||||
(layer.feed_forward, "w1"),
|
||||
(layer.feed_forward, "w2"),
|
||||
(layer.feed_forward, "w3"),
|
||||
]
|
||||
)
|
||||
|
||||
for module, layer in linears:
|
||||
updated_linear = lora.Linear(
|
||||
in_features=getattr(module, layer).in_features,
|
||||
out_features=getattr(module, layer).out_features,
|
||||
bias=getattr(module, layer).bias,
|
||||
r=lora_config.r,
|
||||
lora_alpha=lora_config.lora_alpha,
|
||||
lora_dropout=lora_config.lora_dropout,
|
||||
)
|
||||
setattr(module, layer, updated_linear)
|
||||
|
||||
# Mark only the LoRA layers as trainable
|
||||
lora.mark_only_lora_as_trainable(model, bias="none")
|
||||
|
||||
|
||||
def get_merged_state_dict(model):
|
||||
# This line will merge the state dict of the model and the LoRA parameters
|
||||
model.eval()
|
||||
|
||||
# Then we need to remove the LoRA parameters from the state dict
|
||||
state_dict = model.state_dict()
|
||||
for name in list(state_dict.keys()):
|
||||
if "lora" in name:
|
||||
state_dict.pop(name)
|
||||
|
||||
return state_dict
|
||||
@@ -0,0 +1,596 @@
|
||||
import math
|
||||
from functools import partial
|
||||
from math import prod
|
||||
from typing import Callable
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from torch.nn.utils.parametrizations import weight_norm
|
||||
from torch.nn.utils.parametrize import remove_parametrizations
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
|
||||
def sequence_mask(length, max_length=None):
|
||||
if max_length is None:
|
||||
max_length = length.max()
|
||||
x = torch.arange(max_length, dtype=length.dtype, device=length.device)
|
||||
return x.unsqueeze(0) < length.unsqueeze(1)
|
||||
|
||||
|
||||
def init_weights(m, mean=0.0, std=0.01):
|
||||
classname = m.__class__.__name__
|
||||
if classname.find("Conv1D") != -1:
|
||||
m.weight.data.normal_(mean, std)
|
||||
|
||||
|
||||
def get_padding(kernel_size, dilation=1):
|
||||
return (kernel_size * dilation - dilation) // 2
|
||||
|
||||
|
||||
def unpad1d(x: torch.Tensor, paddings: tuple[int, int]):
|
||||
"""Remove padding from x, handling properly zero padding. Only for 1d!"""
|
||||
padding_left, padding_right = paddings
|
||||
assert padding_left >= 0 and padding_right >= 0, (padding_left, padding_right)
|
||||
assert (padding_left + padding_right) <= x.shape[-1]
|
||||
end = x.shape[-1] - padding_right
|
||||
return x[..., padding_left:end]
|
||||
|
||||
|
||||
def get_extra_padding_for_conv1d(
|
||||
x: torch.Tensor, kernel_size: int, stride: int, padding_total: int = 0
|
||||
) -> int:
|
||||
"""See `pad_for_conv1d`."""
|
||||
length = x.shape[-1]
|
||||
n_frames = (length - kernel_size + padding_total) / stride + 1
|
||||
ideal_length = (math.ceil(n_frames) - 1) * stride + (kernel_size - padding_total)
|
||||
return ideal_length - length
|
||||
|
||||
|
||||
def pad1d(
|
||||
x: torch.Tensor,
|
||||
paddings: tuple[int, int],
|
||||
mode: str = "zeros",
|
||||
value: float = 0.0,
|
||||
):
|
||||
"""Tiny wrapper around F.pad, just to allow for reflect padding on small input.
|
||||
If this is the case, we insert extra 0 padding to the right
|
||||
before the reflection happen.
|
||||
"""
|
||||
length = x.shape[-1]
|
||||
padding_left, padding_right = paddings
|
||||
assert padding_left >= 0 and padding_right >= 0, (padding_left, padding_right)
|
||||
if mode == "reflect":
|
||||
max_pad = max(padding_left, padding_right)
|
||||
extra_pad = 0
|
||||
if length <= max_pad:
|
||||
extra_pad = max_pad - length + 1
|
||||
x = F.pad(x, (0, extra_pad))
|
||||
padded = F.pad(x, paddings, mode, value)
|
||||
end = padded.shape[-1] - extra_pad
|
||||
return padded[..., :end]
|
||||
else:
|
||||
return F.pad(x, paddings, mode, value)
|
||||
|
||||
|
||||
class FishConvNet(nn.Module):
|
||||
def __init__(
|
||||
self, in_channels, out_channels, kernel_size, dilation=1, stride=1, groups=1
|
||||
):
|
||||
super(FishConvNet, self).__init__()
|
||||
self.conv = nn.Conv1d(
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride=stride,
|
||||
dilation=dilation,
|
||||
groups=groups,
|
||||
)
|
||||
self.stride = stride
|
||||
self.kernel_size = (kernel_size - 1) * dilation + 1
|
||||
self.dilation = dilation
|
||||
|
||||
def forward(self, x):
|
||||
pad = self.kernel_size - self.stride
|
||||
extra_padding = get_extra_padding_for_conv1d(
|
||||
x, self.kernel_size, self.stride, pad
|
||||
)
|
||||
x = pad1d(x, (pad, extra_padding), mode="constant", value=0)
|
||||
return self.conv(x).contiguous()
|
||||
|
||||
def weight_norm(self, name="weight", dim=0):
|
||||
self.conv = weight_norm(self.conv, name=name, dim=dim)
|
||||
return self
|
||||
|
||||
def remove_weight_norm(self):
|
||||
self.conv = remove_parametrizations(self.conv)
|
||||
return self
|
||||
|
||||
|
||||
class FishTransConvNet(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, kernel_size, dilation=1, stride=1):
|
||||
super(FishTransConvNet, self).__init__()
|
||||
self.conv = nn.ConvTranspose1d(
|
||||
in_channels, out_channels, kernel_size, stride=stride, dilation=dilation
|
||||
)
|
||||
self.stride = stride
|
||||
self.kernel_size = kernel_size
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv(x)
|
||||
pad = self.kernel_size - self.stride
|
||||
padding_right = math.ceil(pad)
|
||||
padding_left = pad - padding_right
|
||||
x = unpad1d(x, (padding_left, padding_right))
|
||||
return x.contiguous()
|
||||
|
||||
def weight_norm(self, name="weight", dim=0):
|
||||
self.conv = weight_norm(self.conv, name=name, dim=dim)
|
||||
return self
|
||||
|
||||
def remove_weight_norm(self):
|
||||
self.conv = remove_parametrizations(self.conv)
|
||||
return self
|
||||
|
||||
|
||||
class ResBlock1(torch.nn.Module):
|
||||
def __init__(self, channels, kernel_size=3, dilation=(1, 3, 5)):
|
||||
super().__init__()
|
||||
|
||||
self.convs1 = nn.ModuleList(
|
||||
[
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[0]
|
||||
).weight_norm(),
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[1]
|
||||
).weight_norm(),
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[2]
|
||||
).weight_norm(),
|
||||
]
|
||||
)
|
||||
self.convs1.apply(init_weights)
|
||||
|
||||
self.convs2 = nn.ModuleList(
|
||||
[
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[0]
|
||||
).weight_norm(),
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[1]
|
||||
).weight_norm(),
|
||||
FishConvNet(
|
||||
channels, channels, kernel_size, stride=1, dilation=dilation[2]
|
||||
).weight_norm(),
|
||||
]
|
||||
)
|
||||
self.convs2.apply(init_weights)
|
||||
|
||||
def forward(self, x):
|
||||
for c1, c2 in zip(self.convs1, self.convs2):
|
||||
xt = F.silu(x)
|
||||
xt = c1(xt)
|
||||
xt = F.silu(xt)
|
||||
xt = c2(xt)
|
||||
x = xt + x
|
||||
return x
|
||||
|
||||
def remove_parametrizations(self):
|
||||
for conv in self.convs1:
|
||||
remove_parametrizations(conv, tensor_name="weight")
|
||||
for conv in self.convs2:
|
||||
remove_parametrizations(conv, tensor_name="weight")
|
||||
|
||||
|
||||
class ParallelBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
channels: int,
|
||||
kernel_sizes: tuple[int] = (3, 7, 11),
|
||||
dilation_sizes: tuple[tuple[int]] = ((1, 3, 5), (1, 3, 5), (1, 3, 5)),
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
assert len(kernel_sizes) == len(dilation_sizes)
|
||||
|
||||
self.blocks = nn.ModuleList()
|
||||
for k, d in zip(kernel_sizes, dilation_sizes):
|
||||
self.blocks.append(ResBlock1(channels, k, d))
|
||||
|
||||
def forward(self, x):
|
||||
return torch.stack([block(x) for block in self.blocks], dim=0).mean(dim=0)
|
||||
|
||||
def remove_parametrizations(self):
|
||||
for block in self.blocks:
|
||||
block.remove_parametrizations()
|
||||
|
||||
|
||||
class HiFiGANGenerator(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
hop_length: int = 512,
|
||||
upsample_rates: tuple[int] = (8, 8, 2, 2, 2),
|
||||
upsample_kernel_sizes: tuple[int] = (16, 16, 8, 2, 2),
|
||||
resblock_kernel_sizes: tuple[int] = (3, 7, 11),
|
||||
resblock_dilation_sizes: tuple[tuple[int]] = ((1, 3, 5), (1, 3, 5), (1, 3, 5)),
|
||||
num_mels: int = 128,
|
||||
upsample_initial_channel: int = 512,
|
||||
pre_conv_kernel_size: int = 7,
|
||||
post_conv_kernel_size: int = 7,
|
||||
post_activation: Callable = partial(nn.SiLU, inplace=True),
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
assert (
|
||||
prod(upsample_rates) == hop_length
|
||||
), f"hop_length must be {prod(upsample_rates)}"
|
||||
|
||||
self.conv_pre = FishConvNet(
|
||||
num_mels,
|
||||
upsample_initial_channel,
|
||||
pre_conv_kernel_size,
|
||||
stride=1,
|
||||
).weight_norm()
|
||||
|
||||
self.num_upsamples = len(upsample_rates)
|
||||
self.num_kernels = len(resblock_kernel_sizes)
|
||||
|
||||
self.noise_convs = nn.ModuleList()
|
||||
self.ups = nn.ModuleList()
|
||||
|
||||
for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):
|
||||
self.ups.append(
|
||||
FishTransConvNet(
|
||||
upsample_initial_channel // (2**i),
|
||||
upsample_initial_channel // (2 ** (i + 1)),
|
||||
k,
|
||||
stride=u,
|
||||
).weight_norm()
|
||||
)
|
||||
|
||||
self.resblocks = nn.ModuleList()
|
||||
for i in range(len(self.ups)):
|
||||
ch = upsample_initial_channel // (2 ** (i + 1))
|
||||
self.resblocks.append(
|
||||
ParallelBlock(ch, resblock_kernel_sizes, resblock_dilation_sizes)
|
||||
)
|
||||
|
||||
self.activation_post = post_activation()
|
||||
self.conv_post = FishConvNet(
|
||||
ch, 1, post_conv_kernel_size, stride=1
|
||||
).weight_norm()
|
||||
self.ups.apply(init_weights)
|
||||
self.conv_post.apply(init_weights)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv_pre(x)
|
||||
|
||||
for i in range(self.num_upsamples):
|
||||
x = F.silu(x, inplace=True)
|
||||
x = self.ups[i](x)
|
||||
|
||||
if self.training and self.checkpointing:
|
||||
x = checkpoint(
|
||||
self.resblocks[i],
|
||||
x,
|
||||
use_reentrant=False,
|
||||
)
|
||||
else:
|
||||
x = self.resblocks[i](x)
|
||||
|
||||
x = self.activation_post(x)
|
||||
x = self.conv_post(x)
|
||||
x = torch.tanh(x)
|
||||
|
||||
return x
|
||||
|
||||
def remove_parametrizations(self):
|
||||
for up in self.ups:
|
||||
remove_parametrizations(up, tensor_name="weight")
|
||||
for block in self.resblocks:
|
||||
block.remove_parametrizations()
|
||||
remove_parametrizations(self.conv_pre, tensor_name="weight")
|
||||
remove_parametrizations(self.conv_post, tensor_name="weight")
|
||||
|
||||
|
||||
# DropPath copied from timm library
|
||||
def drop_path(
|
||||
x, drop_prob: float = 0.0, training: bool = False, scale_by_keep: bool = True
|
||||
):
|
||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
|
||||
|
||||
This is the same as the DropConnect impl I created for EfficientNet, etc networks, however,
|
||||
the original name is misleading as 'Drop Connect' is a different form of dropout in a separate paper...
|
||||
See discussion: https://github.com/tensorflow/tpu/issues/494#issuecomment-532968956 ... I've opted for
|
||||
changing the layer and argument names to 'drop path' rather than mix DropConnect as a layer name and use
|
||||
'survival rate' as the argument.
|
||||
|
||||
""" # noqa: E501
|
||||
|
||||
if drop_prob == 0.0 or not training:
|
||||
return x
|
||||
keep_prob = 1 - drop_prob
|
||||
shape = (x.shape[0],) + (1,) * (
|
||||
x.ndim - 1
|
||||
) # work with diff dim tensors, not just 2D ConvNets
|
||||
random_tensor = x.new_empty(shape).bernoulli_(keep_prob)
|
||||
if keep_prob > 0.0 and scale_by_keep:
|
||||
random_tensor.div_(keep_prob)
|
||||
return x * random_tensor
|
||||
|
||||
|
||||
class DropPath(nn.Module):
|
||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).""" # noqa: E501
|
||||
|
||||
def __init__(self, drop_prob: float = 0.0, scale_by_keep: bool = True):
|
||||
super(DropPath, self).__init__()
|
||||
self.drop_prob = drop_prob
|
||||
self.scale_by_keep = scale_by_keep
|
||||
|
||||
def forward(self, x):
|
||||
return drop_path(x, self.drop_prob, self.training, self.scale_by_keep)
|
||||
|
||||
def extra_repr(self):
|
||||
return f"drop_prob={round(self.drop_prob,3):0.3f}"
|
||||
|
||||
|
||||
class LayerNorm(nn.Module):
|
||||
r"""LayerNorm that supports two data formats: channels_last (default) or channels_first.
|
||||
The ordering of the dimensions in the inputs. channels_last corresponds to inputs with
|
||||
shape (batch_size, height, width, channels) while channels_first corresponds to inputs
|
||||
with shape (batch_size, channels, height, width).
|
||||
""" # noqa: E501
|
||||
|
||||
def __init__(self, normalized_shape, eps=1e-6, data_format="channels_last"):
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(normalized_shape))
|
||||
self.bias = nn.Parameter(torch.zeros(normalized_shape))
|
||||
self.eps = eps
|
||||
self.data_format = data_format
|
||||
if self.data_format not in ["channels_last", "channels_first"]:
|
||||
raise NotImplementedError
|
||||
self.normalized_shape = (normalized_shape,)
|
||||
|
||||
def forward(self, x):
|
||||
if self.data_format == "channels_last":
|
||||
return F.layer_norm(
|
||||
x, self.normalized_shape, self.weight, self.bias, self.eps
|
||||
)
|
||||
elif self.data_format == "channels_first":
|
||||
u = x.mean(1, keepdim=True)
|
||||
s = (x - u).pow(2).mean(1, keepdim=True)
|
||||
x = (x - u) / torch.sqrt(s + self.eps)
|
||||
x = self.weight[:, None] * x + self.bias[:, None]
|
||||
return x
|
||||
|
||||
|
||||
# ConvNeXt Block copied from https://github.com/fishaudio/fish-diffusion/blob/main/fish_diffusion/modules/convnext.py
|
||||
class ConvNeXtBlock(nn.Module):
|
||||
r"""ConvNeXt Block. There are two equivalent implementations:
|
||||
(1) DwConv -> LayerNorm (channels_first) -> 1x1 Conv -> GELU -> 1x1 Conv; all in (N, C, H, W)
|
||||
(2) DwConv -> Permute to (N, H, W, C); LayerNorm (channels_last) -> Linear -> GELU -> Linear; Permute back
|
||||
We use (2) as we find it slightly faster in PyTorch
|
||||
|
||||
Args:
|
||||
dim (int): Number of input channels.
|
||||
drop_path (float): Stochastic depth rate. Default: 0.0
|
||||
layer_scale_init_value (float): Init value for Layer Scale. Default: 1e-6.
|
||||
mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.0.
|
||||
kernel_size (int): Kernel size for depthwise conv. Default: 7.
|
||||
dilation (int): Dilation for depthwise conv. Default: 1.
|
||||
""" # noqa: E501
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
drop_path: float = 0.0,
|
||||
layer_scale_init_value: float = 1e-6,
|
||||
mlp_ratio: float = 4.0,
|
||||
kernel_size: int = 7,
|
||||
dilation: int = 1,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.dwconv = FishConvNet(
|
||||
dim,
|
||||
dim,
|
||||
kernel_size=kernel_size,
|
||||
# padding=int(dilation * (kernel_size - 1) / 2),
|
||||
groups=dim,
|
||||
) # depthwise conv
|
||||
self.norm = LayerNorm(dim, eps=1e-6)
|
||||
self.pwconv1 = nn.Linear(
|
||||
dim, int(mlp_ratio * dim)
|
||||
) # pointwise/1x1 convs, implemented with linear layers
|
||||
self.act = nn.GELU()
|
||||
self.pwconv2 = nn.Linear(int(mlp_ratio * dim), dim)
|
||||
self.gamma = (
|
||||
nn.Parameter(layer_scale_init_value * torch.ones((dim)), requires_grad=True)
|
||||
if layer_scale_init_value > 0
|
||||
else None
|
||||
)
|
||||
self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
|
||||
|
||||
def forward(self, x, apply_residual: bool = True):
|
||||
input = x
|
||||
|
||||
x = self.dwconv(x)
|
||||
x = x.permute(0, 2, 1) # (N, C, L) -> (N, L, C)
|
||||
x = self.norm(x)
|
||||
x = self.pwconv1(x)
|
||||
x = self.act(x)
|
||||
x = self.pwconv2(x)
|
||||
|
||||
if self.gamma is not None:
|
||||
x = self.gamma * x
|
||||
|
||||
x = x.permute(0, 2, 1) # (N, L, C) -> (N, C, L)
|
||||
x = self.drop_path(x)
|
||||
|
||||
if apply_residual:
|
||||
x = input + x
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class ConvNeXtEncoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_channels: int = 3,
|
||||
depths: list[int] = [3, 3, 9, 3],
|
||||
dims: list[int] = [96, 192, 384, 768],
|
||||
drop_path_rate: float = 0.0,
|
||||
layer_scale_init_value: float = 1e-6,
|
||||
kernel_size: int = 7,
|
||||
):
|
||||
super().__init__()
|
||||
assert len(depths) == len(dims)
|
||||
|
||||
self.downsample_layers = nn.ModuleList()
|
||||
stem = nn.Sequential(
|
||||
FishConvNet(
|
||||
input_channels,
|
||||
dims[0],
|
||||
kernel_size=7,
|
||||
# padding=3,
|
||||
# padding_mode="replicate",
|
||||
# padding_mode="zeros",
|
||||
),
|
||||
LayerNorm(dims[0], eps=1e-6, data_format="channels_first"),
|
||||
)
|
||||
self.downsample_layers.append(stem)
|
||||
|
||||
for i in range(len(depths) - 1):
|
||||
mid_layer = nn.Sequential(
|
||||
LayerNorm(dims[i], eps=1e-6, data_format="channels_first"),
|
||||
nn.Conv1d(dims[i], dims[i + 1], kernel_size=1),
|
||||
)
|
||||
self.downsample_layers.append(mid_layer)
|
||||
|
||||
self.stages = nn.ModuleList()
|
||||
dp_rates = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))]
|
||||
|
||||
cur = 0
|
||||
for i in range(len(depths)):
|
||||
stage = nn.Sequential(
|
||||
*[
|
||||
ConvNeXtBlock(
|
||||
dim=dims[i],
|
||||
drop_path=dp_rates[cur + j],
|
||||
layer_scale_init_value=layer_scale_init_value,
|
||||
kernel_size=kernel_size,
|
||||
)
|
||||
for j in range(depths[i])
|
||||
]
|
||||
)
|
||||
self.stages.append(stage)
|
||||
cur += depths[i]
|
||||
|
||||
self.norm = LayerNorm(dims[-1], eps=1e-6, data_format="channels_first")
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, (nn.Conv1d, nn.Linear)):
|
||||
nn.init.trunc_normal_(m.weight, std=0.02)
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
for i in range(len(self.downsample_layers)):
|
||||
x = self.downsample_layers[i](x)
|
||||
x = self.stages[i](x)
|
||||
|
||||
return self.norm(x)
|
||||
|
||||
|
||||
class FireflyArchitecture(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
backbone: nn.Module,
|
||||
head: nn.Module,
|
||||
quantizer: nn.Module,
|
||||
spec_transform: nn.Module,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.backbone = backbone
|
||||
self.head = head
|
||||
self.quantizer = quantizer
|
||||
self.spec_transform = spec_transform
|
||||
self.downsample_factor = math.prod(self.quantizer.downsample_factor)
|
||||
|
||||
def forward(self, x: torch.Tensor, template=None, mask=None) -> torch.Tensor:
|
||||
if self.spec_transform is not None:
|
||||
x = self.spec_transform(x)
|
||||
|
||||
x = self.backbone(x)
|
||||
if mask is not None:
|
||||
x = x * mask
|
||||
|
||||
if self.quantizer is not None:
|
||||
vq_result = self.quantizer(x)
|
||||
x = vq_result.z
|
||||
|
||||
if mask is not None:
|
||||
x = x * mask
|
||||
|
||||
x = self.head(x, template=template)
|
||||
|
||||
if x.ndim == 2:
|
||||
x = x[:, None, :]
|
||||
|
||||
if self.vq is not None:
|
||||
return x, vq_result
|
||||
|
||||
return x
|
||||
|
||||
def encode(self, audios, audio_lengths):
|
||||
audios = audios.float()
|
||||
|
||||
mels = self.spec_transform(audios)
|
||||
mel_lengths = audio_lengths // self.spec_transform.hop_length
|
||||
mel_masks = sequence_mask(mel_lengths, mels.shape[2])
|
||||
mel_masks_float_conv = mel_masks[:, None, :].float()
|
||||
mels = mels * mel_masks_float_conv
|
||||
|
||||
# Encode
|
||||
encoded_features = self.backbone(mels) * mel_masks_float_conv
|
||||
feature_lengths = mel_lengths // self.downsample_factor
|
||||
|
||||
return self.quantizer.encode(encoded_features), feature_lengths
|
||||
|
||||
def decode(self, indices, feature_lengths) -> torch.Tensor:
|
||||
mel_masks = sequence_mask(
|
||||
feature_lengths * self.downsample_factor,
|
||||
indices.shape[2] * self.downsample_factor,
|
||||
)
|
||||
mel_masks_float_conv = mel_masks[:, None, :].float()
|
||||
audio_lengths = (
|
||||
feature_lengths * self.downsample_factor * self.spec_transform.hop_length
|
||||
)
|
||||
|
||||
audio_masks = sequence_mask(
|
||||
audio_lengths,
|
||||
indices.shape[2] * self.downsample_factor * self.spec_transform.hop_length,
|
||||
)
|
||||
audio_masks_float_conv = audio_masks[:, None, :].float()
|
||||
|
||||
z = self.quantizer.decode(indices) * mel_masks_float_conv
|
||||
x = self.head(z) * audio_masks_float_conv
|
||||
|
||||
return x, audio_lengths
|
||||
|
||||
def remove_parametrizations(self):
|
||||
if hasattr(self.backbone, "remove_parametrizations"):
|
||||
self.backbone.remove_parametrizations()
|
||||
|
||||
if hasattr(self.head, "remove_parametrizations"):
|
||||
self.head.remove_parametrizations()
|
||||
|
||||
@property
|
||||
def device(self):
|
||||
return next(self.parameters()).device
|
||||
@@ -0,0 +1,116 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from vector_quantize_pytorch import GroupedResidualFSQ
|
||||
|
||||
from .firefly import ConvNeXtBlock, FishConvNet, FishTransConvNet
|
||||
|
||||
|
||||
@dataclass
|
||||
class FSQResult:
|
||||
z: torch.Tensor
|
||||
codes: torch.Tensor
|
||||
latents: torch.Tensor
|
||||
|
||||
|
||||
class DownsampleFiniteScalarQuantize(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
input_dim: int = 512,
|
||||
n_codebooks: int = 9,
|
||||
n_groups: int = 1,
|
||||
levels: tuple[int] = (8, 5, 5, 5), # Approximate 2**10
|
||||
downsample_factor: tuple[int] = (2, 2),
|
||||
downsample_dims: tuple[int] | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
if downsample_dims is None:
|
||||
downsample_dims = [input_dim for _ in range(len(downsample_factor))]
|
||||
|
||||
all_dims = (input_dim,) + tuple(downsample_dims)
|
||||
|
||||
self.residual_fsq = GroupedResidualFSQ(
|
||||
dim=all_dims[-1],
|
||||
levels=levels,
|
||||
num_quantizers=n_codebooks,
|
||||
groups=n_groups,
|
||||
)
|
||||
|
||||
self.downsample_factor = downsample_factor
|
||||
self.downsample_dims = downsample_dims
|
||||
|
||||
self.downsample = nn.Sequential(
|
||||
*[
|
||||
nn.Sequential(
|
||||
FishConvNet(
|
||||
all_dims[idx],
|
||||
all_dims[idx + 1],
|
||||
kernel_size=factor,
|
||||
stride=factor,
|
||||
),
|
||||
ConvNeXtBlock(dim=all_dims[idx + 1]),
|
||||
)
|
||||
for idx, factor in enumerate(downsample_factor)
|
||||
]
|
||||
)
|
||||
|
||||
self.upsample = nn.Sequential(
|
||||
*[
|
||||
nn.Sequential(
|
||||
FishTransConvNet(
|
||||
all_dims[idx + 1],
|
||||
all_dims[idx],
|
||||
kernel_size=factor,
|
||||
stride=factor,
|
||||
),
|
||||
ConvNeXtBlock(dim=all_dims[idx]),
|
||||
)
|
||||
for idx, factor in reversed(list(enumerate(downsample_factor)))
|
||||
]
|
||||
)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, (nn.Conv1d, nn.Linear)):
|
||||
nn.init.trunc_normal_(m.weight, std=0.02)
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
def forward(self, z) -> FSQResult:
|
||||
original_shape = z.shape
|
||||
z = self.downsample(z)
|
||||
quantized, indices = self.residual_fsq(z.mT)
|
||||
result = FSQResult(
|
||||
z=quantized.mT,
|
||||
codes=indices.mT,
|
||||
latents=z,
|
||||
)
|
||||
result.z = self.upsample(result.z)
|
||||
|
||||
# Pad or crop z to match original shape
|
||||
diff = original_shape[-1] - result.z.shape[-1]
|
||||
left = diff // 2
|
||||
right = diff - left
|
||||
|
||||
if diff > 0:
|
||||
result.z = F.pad(result.z, (left, right))
|
||||
elif diff < 0:
|
||||
result.z = result.z[..., left:-right]
|
||||
|
||||
return result
|
||||
|
||||
def encode(self, z):
|
||||
z = self.downsample(z)
|
||||
_, indices = self.residual_fsq(z.mT)
|
||||
indices = rearrange(indices, "g b l r -> b (g r) l")
|
||||
return indices
|
||||
|
||||
def decode(self, indices: torch.Tensor):
|
||||
indices = rearrange(indices, "b (g r) l -> g b l r", g=self.residual_fsq.groups)
|
||||
z_q = self.residual_fsq.get_output_from_indices(indices)
|
||||
z_q = self.upsample(z_q.mT)
|
||||
return z_q
|
||||
@@ -0,0 +1,94 @@
|
||||
import matplotlib
|
||||
import torch
|
||||
from matplotlib import pyplot as plt
|
||||
|
||||
matplotlib.use("Agg")
|
||||
|
||||
|
||||
def convert_pad_shape(pad_shape):
|
||||
l = pad_shape[::-1]
|
||||
pad_shape = [item for sublist in l for item in sublist]
|
||||
return pad_shape
|
||||
|
||||
|
||||
def sequence_mask(length, max_length=None):
|
||||
if max_length is None:
|
||||
max_length = length.max()
|
||||
x = torch.arange(max_length, dtype=length.dtype, device=length.device)
|
||||
return x.unsqueeze(0) < length.unsqueeze(1)
|
||||
|
||||
|
||||
def init_weights(m, mean=0.0, std=0.01):
|
||||
classname = m.__class__.__name__
|
||||
if classname.find("Conv") != -1:
|
||||
m.weight.data.normal_(mean, std)
|
||||
|
||||
|
||||
def get_padding(kernel_size, dilation=1):
|
||||
return int((kernel_size * dilation - dilation) / 2)
|
||||
|
||||
|
||||
def plot_mel(data, titles=None):
|
||||
fig, axes = plt.subplots(len(data), 1, squeeze=False)
|
||||
|
||||
if titles is None:
|
||||
titles = [None for i in range(len(data))]
|
||||
|
||||
plt.tight_layout()
|
||||
|
||||
for i in range(len(data)):
|
||||
mel = data[i]
|
||||
|
||||
if isinstance(mel, torch.Tensor):
|
||||
mel = mel.float().detach().cpu().numpy()
|
||||
|
||||
axes[i][0].imshow(mel, origin="lower")
|
||||
axes[i][0].set_aspect(2.5, adjustable="box")
|
||||
axes[i][0].set_ylim(0, mel.shape[0])
|
||||
axes[i][0].set_title(titles[i], fontsize="medium")
|
||||
axes[i][0].tick_params(labelsize="x-small", left=False, labelleft=False)
|
||||
axes[i][0].set_anchor("W")
|
||||
|
||||
return fig
|
||||
|
||||
|
||||
def slice_segments(x, ids_str, segment_size=4):
|
||||
ret = torch.zeros_like(x[:, :, :segment_size])
|
||||
for i in range(x.size(0)):
|
||||
idx_str = ids_str[i]
|
||||
idx_end = idx_str + segment_size
|
||||
ret[i] = x[i, :, idx_str:idx_end]
|
||||
|
||||
return ret
|
||||
|
||||
|
||||
def rand_slice_segments(x, x_lengths=None, segment_size=4):
|
||||
b, d, t = x.size()
|
||||
if x_lengths is None:
|
||||
x_lengths = t
|
||||
ids_str_max = torch.clamp(x_lengths - segment_size + 1, min=0)
|
||||
ids_str = (torch.rand([b], device=x.device) * ids_str_max).to(dtype=torch.long)
|
||||
ret = slice_segments(x, ids_str, segment_size)
|
||||
return ret, ids_str
|
||||
|
||||
|
||||
@torch.jit.script
|
||||
def fused_add_tanh_sigmoid_multiply(in_act, n_channels):
|
||||
n_channels_int = n_channels[0]
|
||||
t_act = torch.tanh(in_act[:, :n_channels_int, :])
|
||||
s_act = torch.sigmoid(in_act[:, n_channels_int:, :])
|
||||
acts = t_act * s_act
|
||||
|
||||
return acts
|
||||
|
||||
|
||||
def avg_with_mask(x, mask):
|
||||
assert mask.dtype == torch.float, "Mask should be float"
|
||||
|
||||
if mask.ndim == 2:
|
||||
mask = mask.unsqueeze(1)
|
||||
|
||||
if mask.shape[1] == 1:
|
||||
mask = mask.expand_as(x)
|
||||
|
||||
return (x * mask).sum() / mask.sum()
|
||||