Compare commits
664
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
34a64597f4 | ||
|
|
4998121d6d | ||
|
|
1c25e3ed55 | ||
|
|
e98886f32f | ||
|
|
8619d2942d | ||
|
|
1876244531 | ||
|
|
34f709725e | ||
|
|
c5d6341f98 | ||
|
|
5325aa2cd2 | ||
|
|
74d6fc92a9 | ||
|
|
daaefc5f63 | ||
|
|
3eee42e917 | ||
|
|
28c0a5dce9 | ||
|
|
275dce4700 | ||
|
|
83fc3449e8 | ||
|
|
e96f359236 | ||
|
|
e89273a0da | ||
|
|
407e918676 | ||
|
|
94d86c168b | ||
|
|
5a4c227572 | ||
|
|
554f934d37 | ||
|
|
51b72e7434 | ||
|
|
ebfbf6046a | ||
|
|
892324ae59 | ||
|
|
b66e1a0427 | ||
|
|
02ff44fcd1 | ||
|
|
eb551efe7f | ||
|
|
4af3a99f9e | ||
|
|
b3d16e0284 | ||
|
|
c466fe9d73 | ||
|
|
494129344f | ||
|
|
8b9ecfdbcb | ||
|
|
58440f43de | ||
|
|
c8238e1f7b | ||
|
|
5ff08541f6 | ||
|
|
f771899001 | ||
|
|
c609c28e63 | ||
|
|
c19f72f722 | ||
|
|
43b4778474 | ||
|
|
08f5489406 | ||
|
|
7e6e52e02a | ||
|
|
e4114d8fd4 | ||
|
|
a19ed188c8 | ||
|
|
6d3a71b681 | ||
|
|
81103cd48d | ||
|
|
a3d4a1cef4 | ||
|
|
954dc8e35c | ||
|
|
278e8bb8f3 | ||
|
|
b378896b82 | ||
|
|
1d2a7f233c | ||
|
|
b4aed1721f | ||
|
|
3447ded21b | ||
|
|
0c01223d3d | ||
|
|
f8089afbf1 | ||
|
|
2b0c1e66cb | ||
|
|
79e7dac0ac | ||
|
|
2b3f3cddd4 | ||
|
|
064195363c | ||
|
|
ed1e17f37f | ||
|
|
e32d137aec | ||
|
|
62a5530c21 | ||
|
|
b3be22e1e2 | ||
|
|
7156b8d7d5 | ||
|
|
6dd7980ec2 | ||
|
|
b3c54b9e5b | ||
|
|
333488ac09 | ||
|
|
c7526707cc | ||
|
|
a0ffbe2927 | ||
|
|
faef037f2e | ||
|
|
9c45a15f53 | ||
|
|
e6ff7ae2a3 | ||
|
|
897afb8b94 | ||
|
|
f8178a5a14 | ||
|
|
4c76b3cc7a | ||
|
|
3e0e8cf534 | ||
|
|
635fb2350d | ||
|
|
a52f29374c | ||
|
|
31827c9e0a | ||
|
|
74688e56b5 | ||
|
|
e56b2ae0c4 | ||
|
|
6c6d4b34e7 | ||
|
|
93d6ee49ba | ||
|
|
e678eeb2dc | ||
|
|
432403b12d | ||
|
|
8e899953ff | ||
|
|
fd80dc653d | ||
|
|
9fb510b618 | ||
|
|
97f48ef433 | ||
|
|
5975deed62 | ||
|
|
60c420731d | ||
|
|
a9a66a22a0 | ||
|
|
fd415488df | ||
|
|
a496c1c656 | ||
|
|
2788f9351a | ||
|
|
5209b6ebdd | ||
|
|
db2c68e879 | ||
|
|
682774d730 | ||
|
|
2f4d6caa90 | ||
|
|
dcd8e95c7b | ||
|
|
9e2c9eaa01 | ||
|
|
567f9f66f8 | ||
|
|
48a82aba37 | ||
|
|
4143cd1f95 | ||
|
|
5b3bffcca9 | ||
|
|
37b091f2ef | ||
|
|
70f3407032 | ||
|
|
d22d40c74a | ||
|
|
6fe5521667 | ||
|
|
d1422b646d | ||
|
|
d4e071be20 | ||
|
|
7455bc9e7f | ||
|
|
45c89c46c5 | ||
|
|
7ec6f8f720 | ||
|
|
cd6f8f14f1 | ||
|
|
8824f74609 | ||
|
|
1ff9d45044 | ||
|
|
ffe11218dd | ||
|
|
6c723e75f3 | ||
|
|
4c3a83a0b3 | ||
|
|
d8489bf368 | ||
|
|
f6a6366c58 | ||
|
|
5f8b952ede | ||
|
|
a07aee80c3 | ||
|
|
235b6498b0 | ||
|
|
a89be390e7 | ||
|
|
3b8c76ea64 | ||
|
|
63950da0d4 | ||
|
|
f245e4b4f8 | ||
|
|
1a8f9196ea | ||
|
|
ea372da285 | ||
|
|
fac510d043 | ||
|
|
b6b1a2aa92 | ||
|
|
464c8a421f | ||
|
|
602bf19d93 | ||
|
|
84cf3b6795 | ||
|
|
83a0280a14 | ||
|
|
182f4082f5 | ||
|
|
8d3f651217 | ||
|
|
330717768d | ||
|
|
ef653171dd | ||
|
|
15fbad6dd5 | ||
|
|
c00cc4b531 | ||
|
|
1f423fca06 | ||
|
|
2cc0e17ade | ||
|
|
b1564eb141 | ||
|
|
97f2b7fdaa | ||
|
|
83825e2060 | ||
|
|
da9c1ab7c1 | ||
|
|
7c5b0a6ccd | ||
|
|
5f21e3a979 | ||
|
|
b47585c68d | ||
|
|
eb1669dab7 | ||
|
|
c1eb81f0f9 | ||
|
|
8fb5c38381 | ||
|
|
6e313fad97 | ||
|
|
294993ca78 | ||
|
|
de7fc3150a | ||
|
|
eb1311b9a2 | ||
|
|
d3ea50fc05 | ||
|
|
842435c016 | ||
|
|
394c6444aa | ||
|
|
7ef00d7749 | ||
|
|
16e3afed62 | ||
|
|
3c08cb2fef | ||
|
|
74744be314 | ||
|
|
cc4ba38e1a | ||
|
|
39e00fd1c4 | ||
|
|
adb2a20a3d | ||
|
|
a6eb95f471 | ||
|
|
f0667d8db2 | ||
|
|
20b395624b | ||
|
|
40e0423a1e | ||
|
|
7e42eef228 | ||
|
|
08ae0a8379 | ||
|
|
b50b51ad92 | ||
|
|
a4300446d7 | ||
|
|
ebdfdc49e7 | ||
|
|
fbf13b7688 | ||
|
|
3a8c34efd0 | ||
|
|
ae5679185c | ||
|
|
067870f28f | ||
|
|
535cb6330f | ||
|
|
b08681f697 | ||
|
|
d3ea240481 | ||
|
|
40bb7a55d7 | ||
|
|
27c7046472 | ||
|
|
8d62c7273e | ||
|
|
9dd3c8c494 | ||
|
|
46566a713d | ||
|
|
fb21d9938a | ||
|
|
be284a0274 | ||
|
|
c58d27384a | ||
|
|
c421a6ed11 | ||
|
|
336403dd98 | ||
|
|
23aeb03f91 | ||
|
|
1a0872ac8e | ||
|
|
0360c1b6be | ||
|
|
559fbf3c36 | ||
|
|
183b05173a | ||
|
|
ebefa85284 | ||
|
|
b2976643f8 | ||
|
|
b5b2c0968e | ||
|
|
dc12340706 | ||
|
|
b2a50079ac | ||
|
|
c9a7881dc5 | ||
|
|
71597994d1 | ||
|
|
5b78e8e00f | ||
|
|
858696d591 | ||
|
|
2a8b2328a5 | ||
|
|
f920d640d0 | ||
|
|
1e16f6be29 | ||
|
|
194ee0307a | ||
|
|
06eedf5ea3 | ||
|
|
a19488e33a | ||
|
|
60773f4860 | ||
|
|
da48eca111 | ||
|
|
e3e81fd4f1 | ||
|
|
91be347ad2 | ||
|
|
b3b0d60536 | ||
|
|
6e5515df28 | ||
|
|
732c672cb4 | ||
|
|
e2ae46b718 | ||
|
|
9cd5a906fd | ||
|
|
27a98335b9 | ||
|
|
65ec02d036 | ||
|
|
c21ea81533 | ||
|
|
8aebbc64c0 | ||
|
|
de3f7d8a9f | ||
|
|
9b2951969f | ||
|
|
cfd61fdbb5 | ||
|
|
78272b0941 | ||
|
|
aa6d088e91 | ||
|
|
594d98265b | ||
|
|
8dc49b8a6b | ||
|
|
5e3a3c6f78 | ||
|
|
e767ef3a2e | ||
|
|
e92a28ba49 | ||
|
|
8358b3014c | ||
|
|
3287e525c5 | ||
|
|
c59023066e | ||
|
|
350138480f | ||
|
|
f1086e192c | ||
|
|
1323daa7b3 | ||
|
|
8a6db0b399 | ||
|
|
8fa3e614a1 | ||
|
|
bec0e85238 | ||
|
|
098ecbbd5d | ||
|
|
8d3cd692e0 | ||
|
|
0183e89cb7 | ||
|
|
58b64d2662 | ||
|
|
f58e57d8f5 | ||
|
|
cdc626b14e | ||
|
|
8349763576 | ||
|
|
02ad56275d | ||
|
|
461a4b2c97 | ||
|
|
65a3f0fe8a | ||
|
|
0cf62a3e0c | ||
|
|
e2492578c7 | ||
|
|
d581ba621a | ||
|
|
d65994469a | ||
|
|
193b0398a8 | ||
|
|
e08fe61f2b | ||
|
|
c3cd4da606 | ||
|
|
a4fe6a634d | ||
|
|
933141c5ea | ||
|
|
3f92ec2c53 | ||
|
|
71541c766a | ||
|
|
0623ee9a36 | ||
|
|
4e54046c5f | ||
|
|
326b6cbbe6 | ||
|
|
71eb96c306 | ||
|
|
a20766c262 | ||
|
|
72aef089ff | ||
|
|
9b011137b3 | ||
|
|
14dacb1ca1 | ||
|
|
fd3198d7db | ||
|
|
1556d341ed | ||
|
|
3794a468fa | ||
|
|
b21d386205 | ||
|
|
9b72cb3e6f | ||
|
|
5b5159e8fc | ||
|
|
113bb8c5df | ||
|
|
4ad069ffdc | ||
|
|
421aa9b66e | ||
|
|
8c40057697 | ||
|
|
791a48a088 | ||
|
|
b85795c263 | ||
|
|
bd895c4d90 | ||
|
|
4731a994b0 | ||
|
|
bf4364b50a | ||
|
|
1817a15c17 | ||
|
|
8bc9abb45b | ||
|
|
2dc20a2cd1 | ||
|
|
7d9f03db56 | ||
|
|
5e5c6a64dc | ||
|
|
9906697152 | ||
|
|
1f7edca9fd | ||
|
|
a7c3aebf27 | ||
|
|
029d7eba30 | ||
|
|
22b86b0e32 | ||
|
|
0e8162191f | ||
|
|
8139823ff5 | ||
|
|
b37a6ff3a7 | ||
|
|
172e1bcc58 | ||
|
|
22f61c0272 | ||
|
|
cd99a6759a | ||
|
|
8cd4d9671e | ||
|
|
459369d3b6 | ||
|
|
15113e2e7a | ||
|
|
9016110a90 | ||
|
|
154070c2d0 | ||
|
|
ed570f7b18 | ||
|
|
492dbab0d6 | ||
|
|
0044ec6623 | ||
|
|
e1fc7291d5 | ||
|
|
d3bbaa7a55 | ||
|
|
d58cc4504a | ||
|
|
d30e3653f2 | ||
|
|
b08bf8d907 | ||
|
|
d86ff4f933 | ||
|
|
79bf57e77c | ||
|
|
3afaf0c5fe | ||
|
|
fcd94894da | ||
|
|
a3cf68c883 | ||
|
|
5df5af0629 | ||
|
|
deaea23969 | ||
|
|
3e6abd244d | ||
|
|
760029dbfe | ||
|
|
f08fa0b9d5 | ||
|
|
c335b0b70c | ||
|
|
cdd41c2aa3 | ||
|
|
d02d16e6e2 | ||
|
|
7434308e92 | ||
|
|
511ea03617 | ||
|
|
8d5f0281ae | ||
|
|
68ec69bc6c | ||
|
|
ef0157a168 | ||
|
|
20c3149b62 | ||
|
|
35f5e4cff8 | ||
|
|
282f973d9b | ||
|
|
4b8b80835b | ||
|
|
a554dcc696 | ||
|
|
a4a1b93bbc | ||
|
|
cea0c041df | ||
|
|
8137d40a92 | ||
|
|
e202525517 | ||
|
|
8084949b14 | ||
|
|
b3f6d06941 | ||
|
|
b7f0d770b8 | ||
|
|
ee0b749421 | ||
|
|
0c86a4ff11 | ||
|
|
2aaf448a24 | ||
|
|
a50abb7b49 | ||
|
|
b88dfd6d1e | ||
|
|
5dc588f47c | ||
|
|
f48b3da80a | ||
|
|
9462239286 | ||
|
|
c70648b740 | ||
|
|
45187a0d22 | ||
|
|
b88ae0d76d | ||
|
|
c564dc569d | ||
|
|
55a595ece8 | ||
|
|
c7a9f094c8 | ||
|
|
2dbd3bd700 | ||
|
|
c19b099618 | ||
|
|
7db5b81522 | ||
|
|
be1eb20b95 | ||
|
|
ce836de93d | ||
|
|
44eb12b157 | ||
|
|
5f9ff24a3d | ||
|
|
ec81e2427c | ||
|
|
f78a599fa3 | ||
|
|
6c12429fd2 | ||
|
|
c6d7ed5c3f | ||
|
|
b3fb8cf3c5 | ||
|
|
ebaa135355 | ||
|
|
211c7409bc | ||
|
|
dfef484db5 | ||
|
|
4d476bc759 | ||
|
|
a4b273f255 | ||
|
|
6b51d0ccfd | ||
|
|
ff36d291d8 | ||
|
|
73ad660277 | ||
|
|
0817d81290 | ||
|
|
f2ac960431 | ||
|
|
81cc4190fd | ||
|
|
1f34030723 | ||
|
|
757268eccd | ||
|
|
83cf1dbd8e | ||
|
|
b3c17829e9 | ||
|
|
862f627607 | ||
|
|
73161c3248 | ||
|
|
1d1a6a5b63 | ||
|
|
085062a53f | ||
|
|
516082e06c | ||
|
|
054bf4bfb0 | ||
|
|
510434d9a1 | ||
|
|
befdbd5943 | ||
|
|
daee16c393 | ||
|
|
04034c2e50 | ||
|
|
d2a8e64a83 | ||
|
|
99bbe2ecbe | ||
|
|
210690e431 | ||
|
|
17c430042e | ||
|
|
4b8c154784 | ||
|
|
99a789bd5d | ||
|
|
d82c3b9e60 | ||
|
|
08d8ef6e72 | ||
|
|
51581271b4 | ||
|
|
014e5b4933 | ||
|
|
18a652afff | ||
|
|
8aa4e60cf4 | ||
|
|
9750ead38b | ||
|
|
66f085252e | ||
|
|
f8b0701a4b | ||
|
|
1c72bea528 | ||
|
|
c8811aa2d5 | ||
|
|
07d818282d | ||
|
|
e4ea15e66a | ||
|
|
59e2050bc3 | ||
|
|
a0f4a001a3 | ||
|
|
78b9cb8775 | ||
|
|
ff617764b7 | ||
|
|
f8ca07a1d8 | ||
|
|
61306c7510 | ||
|
|
429391840f | ||
|
|
aa3ff1e9c3 | ||
|
|
e0dac4fb40 | ||
|
|
486218b252 | ||
|
|
21383acca7 | ||
|
|
e4949048e0 | ||
|
|
c39a13d6b0 | ||
|
|
ddf7e82b61 | ||
|
|
064a503c31 | ||
|
|
91409c89cb | ||
|
|
330e46e9e8 | ||
|
|
76e172e08c | ||
|
|
896d5e3657 | ||
|
|
a360da1ec7 | ||
|
|
46126baadf | ||
|
|
2eba66af88 | ||
|
|
a5d891a8bc | ||
|
|
d1c7015eaa | ||
|
|
56ed84acbd | ||
|
|
9fbd96d66c | ||
|
|
0e2a86b2ac | ||
|
|
59030ac60c | ||
|
|
2b17f4fac2 | ||
|
|
523670a478 | ||
|
|
1365640065 | ||
|
|
59e2407366 | ||
|
|
1321e1f977 | ||
|
|
14e217a698 | ||
|
|
d4aa360c50 | ||
|
|
93b624e62d | ||
|
|
f423fc7538 | ||
|
|
a7165b506e | ||
|
|
740363460f | ||
|
|
aeb3081341 | ||
|
|
fb62b8ff7e | ||
|
|
74494966bd | ||
|
|
b57709108a | ||
|
|
1f4d4cbdfd | ||
|
|
995454b518 | ||
|
|
3d3ea0c818 | ||
|
|
6dae0b1c0c | ||
|
|
557b062028 | ||
|
|
28b7fc4bae | ||
|
|
61f9c5bced | ||
|
|
f3274235d8 | ||
|
|
81c455e878 | ||
|
|
443bab1364 | ||
|
|
bfe4f15694 | ||
|
|
ee3810a585 | ||
|
|
d6d54da131 | ||
|
|
f034826c6d | ||
|
|
74ee558f49 | ||
|
|
75938cfaeb | ||
|
|
9a0d0bc843 | ||
|
|
5ebb382d85 | ||
|
|
c08f985c8c | ||
|
|
c891f85022 | ||
|
|
ad54bd8f94 | ||
|
|
aef95b043c | ||
|
|
5834bc0a8b | ||
|
|
29e167fecf | ||
|
|
efad4348d5 | ||
|
|
3b4e881b11 | ||
|
|
e415e25897 | ||
|
|
e2b563ab2d | ||
|
|
6df13b07bb | ||
|
|
39e398e815 | ||
|
|
e780368bb6 | ||
|
|
8a987de2a8 | ||
|
|
a60217616e | ||
|
|
fa401d1d3f | ||
|
|
81f5e207d0 | ||
|
|
7774b0d1eb | ||
|
|
82cc72a739 | ||
|
|
39848c3dff | ||
|
|
c8ba21979e | ||
|
|
38c2c2ccaf | ||
|
|
c42a86310d | ||
|
|
b0a9e72008 | ||
|
|
7573495d70 | ||
|
|
130fcd397a | ||
|
|
41b14c6a13 | ||
|
|
f751508193 | ||
|
|
8eef9df1b4 | ||
|
|
61643ba9e8 | ||
|
|
80d857ac86 | ||
|
|
99fe427ac8 | ||
|
|
cd2ffd7229 | ||
|
|
825fe9eada | ||
|
|
e3f002e280 | ||
|
|
a23d24fd11 | ||
|
|
2fd41ff003 | ||
|
|
bc4e4f6fa6 | ||
|
|
eaee962fa9 | ||
|
|
b07538bf8a | ||
|
|
21fad0d34a | ||
|
|
77a79e67b6 | ||
|
|
cf2e4f48d0 | ||
|
|
9c406ca0d7 | ||
|
|
7e7dcbe406 | ||
|
|
9711db6896 | ||
|
|
e37dae81ac | ||
|
|
20ca4c089c | ||
|
|
518ce56f09 | ||
|
|
997738438f | ||
|
|
b76c71e634 | ||
|
|
203105e504 | ||
|
|
76ce11567a | ||
|
|
d6980acdd1 | ||
|
|
2448aa7c8c | ||
|
|
bb167c3a70 | ||
|
|
85232cc966 | ||
|
|
ad5632cc50 | ||
|
|
9c8b50a29a | ||
|
|
68cf1dd673 | ||
|
|
ae766cb6dd | ||
|
|
2e84d17e50 | ||
|
|
9f4a42baf9 | ||
|
|
7a03fe4ccc | ||
|
|
af8ade6655 | ||
|
|
af49b2a9e9 | ||
|
|
0ca3104739 | ||
|
|
781c8ca81d | ||
|
|
2c7fa77589 | ||
|
|
4cd6d293ca | ||
|
|
5149e05147 | ||
|
|
fbd9b89c01 | ||
|
|
a0a8e4fb76 | ||
|
|
d57553d520 | ||
|
|
8e0e8ba91e | ||
|
|
28ea2cc1b2 | ||
|
|
4a07be0b0e | ||
|
|
ff69d06080 | ||
|
|
390d7f1896 | ||
|
|
ca9bbce122 | ||
|
|
0fde7b4edd | ||
|
|
74115d65c4 | ||
|
|
a261d31f3a | ||
|
|
da2c84cf50 | ||
|
|
3dc3c31fda | ||
|
|
d14563a9d4 | ||
|
|
784c246e9c | ||
|
|
11996dc9b4 | ||
|
|
a03025a8f3 | ||
|
|
035214520e | ||
|
|
93e98f1ef5 | ||
|
|
8b4ad5bf4f | ||
|
|
afdd297bd5 | ||
|
|
8aa1181ab2 | ||
|
|
b64293c723 | ||
|
|
d681598e6e | ||
|
|
cb86e18190 | ||
|
|
9d9bad199c | ||
|
|
b6bc9b1455 | ||
|
|
5caa97e9a9 | ||
|
|
f507bb4947 | ||
|
|
0379dbb349 | ||
|
|
c521999f80 | ||
|
|
fcc8009570 | ||
|
|
5a7e9c8822 | ||
|
|
c716b3281e | ||
|
|
94b71a5ea0 | ||
|
|
e90373b161 | ||
|
|
7cf1bcc93d | ||
|
|
cca84f1c6c | ||
|
|
1ebf531e1e | ||
|
|
c9adc3e9a5 | ||
|
|
3f19169f55 | ||
|
|
006ad5353a | ||
|
|
589da07fec | ||
|
|
b409ef399a | ||
|
|
8944956249 | ||
|
|
921d927a27 | ||
|
|
06423d785d | ||
|
|
38bf4afdfa | ||
|
|
7a78ef2a97 | ||
|
|
4549eeeab3 | ||
|
|
da08e767db | ||
|
|
0118dd8f1d | ||
|
|
4efdbc2aca | ||
|
|
140eb267d9 | ||
|
|
1b99cf2789 | ||
|
|
ae2646e213 | ||
|
|
36e0f8813a | ||
|
|
0622ea11dc | ||
|
|
df625f66cc | ||
|
|
107f140b56 | ||
|
|
2fd71909a4 | ||
|
|
2b272863b0 | ||
|
|
1bed11aee3 | ||
|
|
5dc0b3cb41 | ||
|
|
80d394d64a | ||
|
|
f58d14ba24 | ||
|
|
d3acf63c4e | ||
|
|
193d7ef7a0 | ||
|
|
e47f894f69 | ||
|
|
7cee5443e4 | ||
|
|
edfe14e9fa | ||
|
|
31f5ce25d4 | ||
|
|
5a80da7ec1 | ||
|
|
ec0dd141e1 | ||
|
|
54d3069702 | ||
|
|
abffb358f9 | ||
|
|
e1459ff3d0 | ||
|
|
f2a826a2f8 | ||
|
|
29409555f3 | ||
|
|
d878709456 | ||
|
|
3810227fd6 | ||
|
|
8b3920706a | ||
|
|
932c7fd80b | ||
|
|
cfb34d011b | ||
|
|
0147265637 | ||
|
|
cf40ce1e7a | ||
|
|
4f608035e2 | ||
|
|
70ee26f0ec | ||
|
|
e14c09b3da | ||
|
|
3ecbc2eab6 | ||
|
|
21b257c25e | ||
|
|
458d1d7770 | ||
|
|
34c1d8d2d0 | ||
|
|
192abaeb9e | ||
|
|
8b2f93325f | ||
|
|
2ca658809a | ||
|
|
86332d9c0a | ||
|
|
f26658b7de | ||
|
|
5a56c37298 | ||
|
|
e01856047d | ||
|
|
c14ef52685 | ||
|
|
c62e40c802 | ||
|
|
370be675e8 | ||
|
|
e72ef94a7f | ||
|
|
e777c2e617 | ||
|
|
e4eee7bc20 | ||
|
|
b2dbc2099e | ||
|
|
24006388be | ||
|
|
cf5bafaee6 | ||
|
|
028c52fd9f | ||
|
|
b7dfe57c17 | ||
|
|
a02832797a |
+29
-112
@@ -1,12 +1,17 @@
|
||||
ucf101_stride4x4x4
|
||||
__pycache__
|
||||
*.mp4
|
||||
.ipynb_checkpoints
|
||||
*.pth
|
||||
UCF-101/
|
||||
results/
|
||||
build/
|
||||
fastvideo.egg-info/
|
||||
wandb/
|
||||
.idea
|
||||
*.ipynb
|
||||
*.jpg
|
||||
!examples/dataset/lingbotworld2/image.jpg
|
||||
*.mp3
|
||||
*.safetensors
|
||||
*.mp4
|
||||
*.png
|
||||
@@ -15,123 +20,35 @@ wandb/
|
||||
*.pt
|
||||
cache_dir/
|
||||
wandb/
|
||||
venv/
|
||||
.venv/
|
||||
test*
|
||||
sample_video*
|
||||
sample_image*
|
||||
512*
|
||||
720*
|
||||
1024*
|
||||
debug*
|
||||
private*
|
||||
caption*
|
||||
*deepspeed*
|
||||
revised*
|
||||
129f*
|
||||
all*
|
||||
read*
|
||||
YSH*
|
||||
*pick*
|
||||
*ysh*
|
||||
hw*
|
||||
257f*
|
||||
513f*
|
||||
taming*
|
||||
221hw*
|
||||
65x512x512
|
||||
runs/
|
||||
samples/
|
||||
Miniconda3-latest-Linux-x86_64.sh
|
||||
*validation/
|
||||
data/
|
||||
outputs/
|
||||
outputs_video
|
||||
checkpoints/
|
||||
sbatch.sh
|
||||
*.out
|
||||
env
|
||||
*.o
|
||||
**/build/
|
||||
**.pyc
|
||||
**.txt
|
||||
*.log
|
||||
weights/
|
||||
logs/
|
||||
/Z-Image/
|
||||
official_weights/
|
||||
converted_weights/
|
||||
|
||||
# SSIM test outputs
|
||||
fastvideo/tests/ssim/generated_videos/
|
||||
**/.cache/**
|
||||
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
dist/
|
||||
*.egg-info/
|
||||
*.egg
|
||||
eggs/
|
||||
.eggs/
|
||||
|
||||
# MkDocs documentation
|
||||
site/
|
||||
docs/getting_started/examples/
|
||||
docs/examples/
|
||||
docs/inference/examples/
|
||||
docs/training/examples/
|
||||
docs/distillation/examples/
|
||||
!requirements-mkdocs.txt
|
||||
|
||||
# VSCode
|
||||
.vscode/
|
||||
|
||||
# DS Store
|
||||
.DS_Store
|
||||
|
||||
# vim swap files
|
||||
*.swo
|
||||
*.swp
|
||||
|
||||
# Python pickle files
|
||||
*.pkl
|
||||
|
||||
# Reference videos (negations must come after the catch-all on line below)
|
||||
|
||||
# Static images
|
||||
!docs/assets/images/**/*.png
|
||||
!comfyui/assets/**/*.png
|
||||
!comfyui/assets/**/*.gif
|
||||
!assets/images/**/*.png
|
||||
!assets/images/**/*.jpg
|
||||
!assets/images/**/*.jpeg
|
||||
!assets/images/**/*.gif
|
||||
!assets/videos/**/*.mp4
|
||||
|
||||
dmd_t2v_output/
|
||||
preprocess_output_text/
|
||||
|
||||
# SvelteKit / Node artifacts under apps/fastvideo_studio/: see apps/fastvideo_studio/.gitignore
|
||||
|
||||
# Next.js / Node artifacts under apps/dreamverse/web/
|
||||
apps/dreamverse/web/node_modules/
|
||||
apps/dreamverse/web/.next/
|
||||
apps/dreamverse/web/out/
|
||||
apps/dreamverse/web/coverage/
|
||||
apps/dreamverse/web/test-results/
|
||||
apps/dreamverse/web/playwright-report/
|
||||
apps/dreamverse/web/.env.local
|
||||
apps/dreamverse/web/.env.development.local
|
||||
apps/dreamverse/web/.env.test.local
|
||||
|
||||
# Generated by apps/dreamverse/scripts/install_native_ffmpeg.sh — host-specific
|
||||
apps/dreamverse/scripts/ffmpeg-env.sh
|
||||
apps/dreamverse/web/.env.production.local
|
||||
|
||||
# Unignore migrated Dreamverse product assets — root .gitignore globally
|
||||
# ignores *.png/*.jpg/*.mp4/*.gif, but apps/dreamverse/web/public/ MUST
|
||||
# be tracked (logo, icons, k2.png, etc.).
|
||||
!apps/dreamverse/web/public/**/*.png
|
||||
!apps/dreamverse/web/public/**/*.jpg
|
||||
!apps/dreamverse/web/public/**/*.jpeg
|
||||
!apps/dreamverse/web/public/**/*.mp4
|
||||
!apps/dreamverse/web/public/**/*.gif
|
||||
!apps/dreamverse/web/prompts/**/*.png
|
||||
!apps/dreamverse/web/prompts/**/*.jpg
|
||||
!apps/dreamverse/web/prompts/**/*.jpeg
|
||||
!apps/dreamverse/web/prompts/**/*.mp4
|
||||
!apps/dreamverse/web/prompts/**/*.gif
|
||||
!apps/dreamverse/gpu-pool.svg
|
||||
!apps/dreamverse/gpu-pool.drawio
|
||||
|
||||
.claude/
|
||||
.codex/
|
||||
.agents/tmp/
|
||||
.sisyphus/
|
||||
openspec/
|
||||
fastvideo/tests/ssim/reference_videos/**
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.mp4
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.png
|
||||
|
||||
# Editor logs and local Python version pins (accidentally committed)
|
||||
*.nvimlog
|
||||
.nvimlog
|
||||
.python-version
|
||||
|
||||
@@ -1,10 +0,0 @@
|
||||
# fastvideo2/rl_rewards is vendored byte-identical from upstream (gated by
|
||||
# sha256 in tests) — formatters must not touch it
|
||||
exclude: ^fastvideo2/rl_rewards/
|
||||
repos:
|
||||
- repo: https://github.com/pre-commit/pre-commit-hooks
|
||||
rev: v5.0.0
|
||||
hooks:
|
||||
- id: trailing-whitespace
|
||||
- id: end-of-file-fixer
|
||||
- id: check-merge-conflict
|
||||
@@ -1,89 +0,0 @@
|
||||
# Repository guidelines (branch `will/v2.1`)
|
||||
|
||||
This branch is the fastvideo2 MVP — one package, one model (Wan2.1), four
|
||||
surfaces. Read `README.md` first.
|
||||
|
||||
## Layout
|
||||
|
||||
| Path | Role |
|
||||
|---|---|
|
||||
| `fastvideo2/card.py` | Frozen data cards, `derive()`, digest, validation (stdlib-only) |
|
||||
| `fastvideo2/loop.py` | Driven-loop protocol + `LoopRunner` (stdlib; NVTX lazily) |
|
||||
| `fastvideo2/pipeline.py` | Stage list with enforced `reads`/`writes` (stdlib-only) |
|
||||
| `fastvideo2/loading.py` | Checkpoint → modules, standalone; component fingerprints |
|
||||
| `fastvideo2/layers/` | Shared model layers (norms, MLP, rotary, attention) — torch-only, checkpoint-key compatible, cast semantics preserved (anchor-proven) |
|
||||
| `fastvideo2/engine.py` | One-shot runner: request → outputs + identity-chained trace |
|
||||
| `fastvideo2/verify.py` | Gates T0–T3 + evidence ledger |
|
||||
| `fastvideo2/registry.py` | The only catalog: name → (card, pipeline builder) |
|
||||
| `fastvideo2/wan21/` | The family's logic: card constant, loop, pipeline, vendored `model.py`, `reference.py` (the executable spec) |
|
||||
| `fastvideo2/wan21/gates/` | The family's measurement side: goldens, anchor adapters, official capture shim, comparison CLI, diagnostics — never imported by logic code |
|
||||
| `fastvideo2/evidence/` | Append-only ledger + blessed baselines (see its README) |
|
||||
| `fastvideo2/tests/` | T0 contract tests — CPU, no torch, no weights |
|
||||
|
||||
## Invariants (enforced by review; violating them is the bug)
|
||||
|
||||
1. **Cards are pure data.** No callables, no live objects, no deploy-local
|
||||
paths. If it can't round-trip through JSON, it doesn't belong on a card.
|
||||
2. **Import direction is one-way:** `card` → `loop`/`pipeline` → `engine` →
|
||||
`verify`. Family packages depend on core, never the reverse.
|
||||
`wan21/reference.py` imports only the vendored official model file
|
||||
(`wan21/model.py`, itself standalone) — never core/runtime modules — and
|
||||
nothing outside `verify.py` may import the reference.
|
||||
3. **Loop modules import torch-free** (torch inside methods) so contracts
|
||||
validate anywhere; `import fastvideo2` must work without torch installed.
|
||||
4. **Model-specific inputs are typed** (`WanForwardInputs`); never add an
|
||||
untyped passthrough kwarg to a forward call.
|
||||
5. **Evidence is append-only and human-owned.** Agents run `verify` and commit
|
||||
the records; agents do not edit tolerances, re-bless baselines, or touch
|
||||
`reference.py` to make a failing gate pass — say so instead.
|
||||
6. **One catalog.** New servable ⇒ card constant + registry entry. No parallel
|
||||
model lists.
|
||||
7. **Official implementations are the numerics authority.** Where fidelity is
|
||||
the requirement, run the authors' modeling code: the Wan DiT is vendored
|
||||
from the pinned official commit (`wan21/model.py`, provenance in its
|
||||
header). Restructuring vendored code (e.g. extracting `layers/`) is
|
||||
allowed ONLY when the anchor stays bitwise 0.0 — the gate, not "verbatim",
|
||||
is the equivalence guarantee. Two invariants for any extraction:
|
||||
checkpoint keys unchanged (Sequential indices preserved) and cast/dtype
|
||||
semantics unchanged (fp32 islands, promotion order, fp64 RoPE). Ports
|
||||
(diffusers, etc.) may serve components only with anchor certification,
|
||||
never on trust; the official repo never becomes a dependency or submodule;
|
||||
when conventions conflict, official wins.
|
||||
8. **One environment for goldens and gates.** The supported env is python 3.12
|
||||
+ torch 2.12 (the fastvideo cluster venv). Goldens are captured with that
|
||||
same env — official code rides `PYTHONPATH`, its extra deps go to a pip
|
||||
`--target` dir — so anchor deltas measure implementation differences, never
|
||||
torch/kernel version differences.
|
||||
|
||||
## Commands
|
||||
|
||||
```bash
|
||||
pytest # T0, runs on a laptop
|
||||
python -m fastvideo2 verify <model> --tier N # gates; appends evidence
|
||||
python -m fastvideo2 describe <model> # card JSON + digest
|
||||
python -m fastvideo2 generate <model> --prompt ... # one request
|
||||
```
|
||||
|
||||
GPU work runs on dlcluster via the `run-fastvideo-dlcluster` skill from the
|
||||
main FastVideo checkout (sync this branch with `git push origin HEAD`, then
|
||||
run inside the branch clone at `/mnt/fv21` — do not disturb `/mnt/FastVideo`'s
|
||||
checkout). One-time per environment, install the package editable with no
|
||||
dependency changes (torch etc. already live in the venv) — after this,
|
||||
scripts and `fastvideo2 <cmd>` work from any directory, no PYTHONPATH:
|
||||
|
||||
```bash
|
||||
/mnt/FastVideo/.venv/bin/pip install -e /mnt/fv21 --no-deps -q
|
||||
```
|
||||
|
||||
Cluster runs append to `fastvideo2/evidence/` and those files get
|
||||
fetched and committed locally, so before every cluster `git pull`, reset that
|
||||
tree or the pull conflicts:
|
||||
|
||||
```bash
|
||||
git checkout -- fastvideo2/evidence; git clean -qfd fastvideo2/evidence; git pull
|
||||
```
|
||||
|
||||
## Commit style
|
||||
|
||||
Short subject with a tag prefix (`[feat]: ...`, `[fix]: ...`, `[docs]: ...`).
|
||||
Do not add AI co-author trailers.
|
||||
@@ -1,187 +1,21 @@
|
||||
Apache License
|
||||
Version 2.0, January 2004
|
||||
http://www.apache.org/licenses/
|
||||
MIT License
|
||||
|
||||
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
|
||||
Copyright (c) 2024 PKU-YUAN's Group (袁粒课题组-北大信工) and Rabbitpre AI
|
||||
|
||||
1. Definitions.
|
||||
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:
|
||||
|
||||
"License" shall mean the terms and conditions for use, reproduction,
|
||||
and distribution as defined by Sections 1 through 9 of this document.
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
"Licensor" shall mean the copyright owner or entity authorized by
|
||||
the copyright owner that is granting the License.
|
||||
|
||||
"Legal Entity" shall mean the union of the acting entity and all
|
||||
other entities that control, are controlled by, or are under common
|
||||
control with that entity. For the purposes of this definition,
|
||||
"control" means (i) the power, direct or indirect, to cause the
|
||||
direction or management of such entity, whether by contract or
|
||||
otherwise, or (ii) ownership of fifty percent (50%) or more of the
|
||||
outstanding shares, or (iii) beneficial ownership of such entity.
|
||||
|
||||
"You" (or "Your") shall mean an individual or Legal Entity
|
||||
exercising permissions granted by this License.
|
||||
|
||||
"Source" form shall mean the preferred form for making modifications,
|
||||
including but not limited to software source code, documentation
|
||||
source, and configuration files.
|
||||
|
||||
"Object" form shall mean any form resulting from mechanical
|
||||
transformation or translation of a Source form, including but
|
||||
not limited to compiled object code, generated documentation,
|
||||
and conversions to other media types.
|
||||
|
||||
"Work" shall mean the work of authorship, whether in Source or
|
||||
Object form, made available under the License, as indicated by a
|
||||
copyright notice that is included in or attached to the work
|
||||
(an example is provided in the Appendix below).
|
||||
|
||||
"Derivative Works" shall mean any work, whether in Source or Object
|
||||
form, that is based on (or derived from) the Work and for which the
|
||||
editorial revisions, annotations, elaborations, or other modifications
|
||||
represent, as a whole, an original work of authorship. For the purposes
|
||||
of this License, Derivative Works shall not include works that remain
|
||||
separable from, or merely link (or bind by name) to the interfaces of,
|
||||
the Work and Derivative Works thereof.
|
||||
|
||||
"Contribution" shall mean any work of authorship, including
|
||||
the original version of the Work and any modifications or additions
|
||||
to that Work or Derivative Works thereof, that is intentionally
|
||||
submitted to Licensor for inclusion in the Work by the copyright owner
|
||||
or by an individual or Legal Entity authorized to submit on behalf of
|
||||
the copyright owner. For the purposes of this definition, "submitted"
|
||||
means any form of electronic, verbal, or written communication sent
|
||||
to the Licensor or its representatives, including but not limited to
|
||||
communication on electronic mailing lists, source code control systems,
|
||||
and issue tracking systems that are managed by, or on behalf of, the
|
||||
Licensor for the purpose of discussing and improving the Work, but
|
||||
excluding communication that is conspicuously marked or otherwise
|
||||
designated in writing by the copyright owner as "Not a Contribution."
|
||||
|
||||
"Contributor" shall mean Licensor and any individual or Legal Entity
|
||||
on behalf of whom a Contribution has been received by Licensor and
|
||||
subsequently incorporated within the Work.
|
||||
|
||||
2. Grant of Copyright License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
copyright license to reproduce, prepare Derivative Works of,
|
||||
publicly display, publicly perform, sublicense, and distribute the
|
||||
Work and such Derivative Works in Source or Object form.
|
||||
|
||||
3. Grant of Patent License. Subject to the terms and conditions of
|
||||
this License, each Contributor hereby grants to You a perpetual,
|
||||
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
|
||||
(except as stated in this section) patent license to make, have made,
|
||||
use, offer to sell, sell, import, and otherwise transfer the Work,
|
||||
where such license applies only to those patent claims licensable
|
||||
by such Contributor that are necessarily infringed by their
|
||||
Contribution(s) alone or by combination of their Contribution(s)
|
||||
with the Work to which such Contribution(s) was submitted. If You
|
||||
institute patent litigation against any entity (including a
|
||||
cross-claim or counterclaim in a lawsuit) alleging that the Work
|
||||
or a Contribution incorporated within the Work constitutes direct
|
||||
or contributory patent infringement, then any patent licenses
|
||||
granted to You under this License for that Work shall terminate
|
||||
as of the date such litigation is filed.
|
||||
|
||||
4. Redistribution. You may reproduce and distribute copies of the
|
||||
Work or Derivative Works thereof in any medium, with or without
|
||||
modifications, and in Source or Object form, provided that You
|
||||
meet the following conditions:
|
||||
|
||||
(a) You must give any other recipients of the Work or
|
||||
Derivative Works a copy of this License; and
|
||||
|
||||
(b) You must cause any modified files to carry prominent notices
|
||||
stating that You changed the files; and
|
||||
|
||||
(c) You must retain, in the Source form of any Derivative Works
|
||||
that You distribute, all copyright, patent, trademark, and
|
||||
attribution notices from the Source form of the Work,
|
||||
excluding those notices that do not pertain to any part of
|
||||
the Derivative Works; and
|
||||
|
||||
(d) If the Work includes a "NOTICE" text file as part of its
|
||||
distribution, then any Derivative Works that You distribute must
|
||||
include a readable copy of the attribution notices contained
|
||||
within such NOTICE file, excluding those notices that do not
|
||||
pertain to any part of the Derivative Works, in at least one
|
||||
of the following places: within a NOTICE text file distributed
|
||||
as part of the Derivative Works; within the Source form or
|
||||
documentation, if provided along with the Derivative Works; or,
|
||||
within a display generated by the Derivative Works, if and
|
||||
wherever such third-party notices normally appear. The contents
|
||||
of the NOTICE file are for informational purposes only and
|
||||
do not modify the License. You may add Your own attribution
|
||||
notices within Derivative Works that You distribute, alongside
|
||||
or as an addendum to the NOTICE text from the Work, provided
|
||||
that such additional attribution notices cannot be construed
|
||||
as modifying the License.
|
||||
|
||||
You may add Your own copyright statement to Your modifications and
|
||||
may provide additional or different license terms and conditions
|
||||
for use, reproduction, or distribution of Your modifications, or
|
||||
for any such Derivative Works as a whole, provided Your use,
|
||||
reproduction, and distribution of the Work otherwise complies with
|
||||
the conditions stated in this License.
|
||||
|
||||
5. Submission of Contributions. Unless You explicitly state otherwise,
|
||||
any Contribution intentionally submitted for inclusion in the Work
|
||||
by You to the Licensor shall be under the terms and conditions of
|
||||
this License, without any additional terms or conditions.
|
||||
Notwithstanding the above, nothing herein shall supersede or modify
|
||||
the terms of any separate license agreement you may have executed
|
||||
with Licensor regarding such Contributions.
|
||||
|
||||
6. Trademarks. This License does not grant permission to use the trade
|
||||
names, trademarks, service marks, or product names of the Licensor,
|
||||
except as required for reasonable and customary use in describing the
|
||||
origin of the Work and reproducing the content of the NOTICE file.
|
||||
|
||||
7. Disclaimer of Warranty. Unless required by applicable law or
|
||||
agreed to in writing, Licensor provides the Work (and each
|
||||
Contributor provides its Contributions) on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
|
||||
implied, including, without limitation, any warranties or conditions
|
||||
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
|
||||
PARTICULAR PURPOSE. You are solely responsible for determining the
|
||||
appropriateness of using or redistributing the Work and assume any
|
||||
risks associated with Your exercise of permissions under this License.
|
||||
|
||||
8. Limitation of Liability. In no event and under no legal theory,
|
||||
whether in tort (including negligence), contract, or otherwise,
|
||||
unless required by applicable law (such as deliberate and grossly
|
||||
negligent acts) or agreed to in writing, shall any Contributor be
|
||||
liable to You for damages, including any direct, indirect, special,
|
||||
incidental, or consequential damages of any character arising as a
|
||||
result of this License or out of the use or inability to use the
|
||||
Work (including but not limited to damages for loss of goodwill,
|
||||
work stoppage, computer failure or malfunction, or any and all
|
||||
other commercial damages or losses), even if such Contributor
|
||||
has been advised of the possibility of such damages.
|
||||
|
||||
9. Accepting Warranty or Additional Liability. While redistributing
|
||||
the Work or Derivative Works thereof, You may choose to offer,
|
||||
and charge a fee for, acceptance of support, warranty, indemnity,
|
||||
or other liability obligations and/or rights consistent with this
|
||||
License. However, in accepting such obligations, You may act only
|
||||
on Your own behalf and on Your sole responsibility, not on behalf
|
||||
of any other Contributor, and only if You agree to indemnify,
|
||||
defend, and hold each Contributor harmless for any liability
|
||||
incurred by, or claims asserted against, such Contributor by reason
|
||||
of your accepting any such warranty or additional liability.
|
||||
|
||||
END OF TERMS AND CONDITIONS
|
||||
|
||||
APPENDIX: How to apply the Apache License to your work.
|
||||
|
||||
To apply the Apache License to your work, attach the following
|
||||
boilerplate notice, with the fields enclosed by brackets "[]"
|
||||
replaced with your own identifying information. (Don't include
|
||||
the brackets!) The text should be enclosed in the appropriate
|
||||
comment syntax for the file format. We also recommend that a
|
||||
file or class name and description of purpose be included on the
|
||||
same "printed page" as the copyright notice for easier
|
||||
identification within third-party archives.
|
||||
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,72 +1,108 @@
|
||||
# fastvideo2 — the v2.1 MVP (branch `will/v2.1`)
|
||||
# Fast Video
|
||||
This is currently based on Open-Sora-1.2.0: https://github.com/PKU-YuanGroup/Open-Sora-Plan/tree/294993ca78bf65dec1c3b6fb25541432c545eda9
|
||||
|
||||
A from-scratch, deliberately small substrate for the FastVideo big bet:
|
||||
**post-training → inference-optimized serving**, designed so that both kinds of
|
||||
agents — the ones that *build* the framework and the ones that will *operate*
|
||||
video models inside products — get inspectable contracts, a ground-truth
|
||||
oracle, machine-readable verification, and an identity-chained runtime.
|
||||
|
||||
This branch is a clean slate: everything except `LICENSE` was removed, and the
|
||||
MVP supports exactly one model, **Wan2.1-T2V-1.3B**, end to end.
|
||||
|
||||
## The four surfaces
|
||||
|
||||
| Surface | Where | What it guarantees |
|
||||
|---|---|---|
|
||||
| **Contracts** | `fastvideo2/card.py`, `pipeline.py`, `loop.py` | Cards are frozen *data* (no callables): JSON round-trip, stable content digest — the identity used by deploy configs, trainers, and RL environment manifests alike. Pipeline stage edges (`reads`/`writes`) are enforced at run time. Loop classes carry a `semantics` id and provenance pins it: distilled weights cannot silently run under a base sampler. |
|
||||
| **Reference** | `fastvideo2/wan21/reference.py` | The complete model in one standalone eager file — the textbook an agent copies, and the oracle the production path is measured against. Never imported by production code. |
|
||||
| **Verifier** | `fastvideo2/verify.py`, `fastvideo2/evidence/` | Tiered gates: T0 contracts (CPU, seconds) → T1 component fingerprints vs a blessed baseline → T2 trajectory parity vs the reference, with tolerance calibrated by measured run-to-run self-noise (the determinism contract) → T3 decoded-output parity + anti-degeneracy anchors. Every run appends typed records (card digest + env fingerprint) to the evidence ledger. |
|
||||
| **Trace** | `fastvideo2/engine.py`, `loop.py` | Every unit of work is named `request/stage/loop.step`; the same identity chain lands in the returned trace (typed timings) and in nested NVTX ranges, so Nsight correlates kernels to model-level identity for free. |
|
||||
|
||||
## Quickstart
|
||||
|
||||
```bash
|
||||
# machine-readable capability discovery
|
||||
python -m fastvideo2 describe wan2.1-t2v-1.3b
|
||||
|
||||
# contracts only — CPU, no weights, no torch
|
||||
python -m fastvideo2 verify wan2.1-t2v-1.3b --tier 0
|
||||
pytest # the same contracts, as tests
|
||||
|
||||
# GPU: bless the component baseline once, then gate against it
|
||||
python -m fastvideo2 verify wan2.1-t2v-1.3b --tier 1 --bless
|
||||
python -m fastvideo2 verify wan2.1-t2v-1.3b --tier 3
|
||||
|
||||
# generate — CLI or the SDK (the handle is the loaded card, modality-neutral)
|
||||
python -m fastvideo2 generate wan2.1-t2v-1.3b --prompt "a cat surfing a wave" --out cat.mp4
|
||||
python -c '
|
||||
import fastvideo2 as fv2
|
||||
model = fv2.load("wan2.1-t2v-1.3b") # -> Model (capabilities from the card)
|
||||
model.generate("a cat surfing a wave", seed=7).save("cat.mp4")'
|
||||
|
||||
# the oracle, standalone (this file works copied out of the repo)
|
||||
python -m fastvideo2.wan21.reference --prompt "a cat surfing a wave" --out ref.mp4
|
||||
## Envrironment
|
||||
Change the index-url cuda version according to your system.
|
||||
```
|
||||
conda create -n fastvideo python=3.10.12
|
||||
conda activate fastvideo
|
||||
pip3 install torch==2.5.0 torchvision --index-url https://download.pytorch.org/whl/cu121
|
||||
pip install git+https://github.com/huggingface/diffusers.git@76b7d86a9a5c0c2186efa09c4a67b5f5666ac9e3
|
||||
pip install packaging ninja && pip install flash-attn==2.7.0.post2 --no-build-isolation
|
||||
```
|
||||
|
||||
Weights resolve from the HF cache (`Wan-AI/Wan2.1-T2V-1.3B-Diffusers`) or an
|
||||
explicit `--root`; components are stock diffusers/transformers modules, so
|
||||
there is no conversion step and `load_component()` works standalone in a REPL.
|
||||
```
|
||||
pip install -e . && pip install -e ".[train]"
|
||||
sudo apt-get update && apt install screen && pip install watch gpustat
|
||||
```
|
||||
|
||||
## Design lineage (what this MVP encodes)
|
||||
## Prepare Data & Models
|
||||
We've prepared some debug data to facilitate development. To make sure the training pipeline is correct, train on the debug data and make sure the model overfit on it (feed it the same text prompt and see if the output video is the same as the training data)
|
||||
|
||||
- **Cards as declared constants; variants as `derive()` diffs** — no builder
|
||||
functions, no factory bags, no toy backends welded into production cards.
|
||||
- **The card digest is the axle artifact**: the same identity a deploy config
|
||||
points at, a trainer stamps provenance into, and an RL environment manifest
|
||||
pins (`substitution: exact | bounded | quality-changing` is already on
|
||||
`Provenance` for the post-training flywheel).
|
||||
- **Typed conditioning** (`WanForwardInputs`): a new control channel is a new
|
||||
field the forward must consume — never a silently dropped kwarg.
|
||||
- **Verification is the product**: gates fail closed, evidence is append-only
|
||||
data, baselines and tolerances are human-owned.
|
||||
- **One loop, runtime-visible**: the driven-loop contract is what sessions,
|
||||
interleaved serving, and RL rollout branching will consume next; the engine
|
||||
stays a deliberately dumb one-shot runner until those consumers land.
|
||||
```
|
||||
mkdir data && mkdir data/outputs/
|
||||
python scripts/download_hf.py --repo_id=Stealths-Video/mochi_diffuser --local_dir=data/mochi --repo_type=model
|
||||
python scripts/download_hf.py --repo_id=Stealths-Video/Merge-30k-Data --local_dir=data/Merge-30k-Data --repo_type=dataset
|
||||
python scripts/download_hf.py --repo_id=Stealths-Video/validation_embeddings --local_dir=data/validation_embeddings --repo_type=dataset
|
||||
cd data/Merge-30k-Data
|
||||
cat Merged30K.tar.gz.part.* > Merged30K.tar.gz
|
||||
rm Merged30K.tar.gz.part.*
|
||||
tar --use-compress-program="pigz --processes 64" -xvf Merged30K.tar.gz
|
||||
mv ephemeral/hao.zhang/codefolder/FastVideo-OSP/data/Merged-30K-Data/* .
|
||||
rm -r ephemeral
|
||||
rm Merged30K.tar.gz
|
||||
cd ../..
|
||||
```
|
||||
|
||||
## Scope and non-goals (MVP)
|
||||
## Things Learned
|
||||
1. shift8 clear but got structural artifacts
|
||||
2. lq, 0.025 vague
|
||||
3. adv not really helpful
|
||||
4. shift8 euler steps 50 v.s. 100 very similar
|
||||
5. 为啥image不会越distill越炸
|
||||
6. EMA, 大batchsize, 1.5,2.5,3.5,4.5
|
||||
7. Must have schedule
|
||||
8. phase 1, 2 learning rate 5e-6不行
|
||||
|
||||
In: Wan2.1 T2V, bidirectional, single GPU, one-shot generation, tiers T0–T3.
|
||||
Out (next, in order): causal/self-forcing students + sessions with forkable
|
||||
state, the post-training flywheel emitting derived cards + evidence, the RL
|
||||
environment server (`reset/step/branch`) over the same contracts, additional
|
||||
model families via `derive()` and new recipe packages.
|
||||
## Experiments
|
||||
Scripts are located at scripts/experiment_N.sh
|
||||
|
||||
1. pcm_linear_quadratic, euler_steps 50, 0.025
|
||||
2. pcm_linear_quadratic, euler_steps 50, 0.05
|
||||
3. shift 8, euler_steps 100
|
||||
4. shift 8, euler_steps 50
|
||||
5. shift 8, euler_steps 100, adv
|
||||
6. pcm_linear_quadratic, euler_steps 50, 0.025, adv
|
||||
7. pcm_linear_quadratic, euler_steps 50, 0.05, multiphase 125
|
||||
8. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75
|
||||
9. pcm_linear_quadratic, euler_steps 50, 0.05, range 0.75
|
||||
10. pcm_linear_quadratic, euler_steps 50, 0.05, batchsize 32
|
||||
11. pcm_linear_quadratic, euler_steps 50, learning rate,1e-7
|
||||
12. shift1, euler_steps 50
|
||||
|
||||
13. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 1
|
||||
14. 4.5 cfg, validation no cfg, pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75
|
||||
15. pcm_linear_quadratic, euler_steps 50, 0.15, linear_range 0.75
|
||||
16. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75 ema 0.95, decay 0.0
|
||||
|
||||
|
||||
17. no cfg, validation no cfg, pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75
|
||||
18. shift16, euler_steps 50
|
||||
|
||||
19. 4step_infer_shift16_euler_50
|
||||
20. 4step_infer_shift12_euler_50
|
||||
21. 4step_infer_lq_euler_50_thresh0.1_lrg_0.75
|
||||
22. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 1, lr 1e-7
|
||||
23. lq_euler_50_thres0.1_lrg_0.75_bs_64
|
||||
24. lq_euler_50_thres0.1_lrg_0.75_lr5e-7
|
||||
|
||||
|
||||
|
||||
25. shift1_euler_50_0.75_phase1
|
||||
26. kill
|
||||
27. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 1, ema 0.95, cfg 4.5
|
||||
|
||||
28. lq_euler_50_thresh0.1_lrg_0.75_phase1_ema0.95
|
||||
29. lq_euler_50_thres0.1_lrg_0.75_phase_ema0.95_cfg7
|
||||
30. lq_euler_50_thresh0.1_lrg_0.75_phase1_ema0.98_cfg4.5
|
||||
31. lq_euler_50_thresh0.1_lrg_0.75_phase1_lr_3e-7
|
||||
32. lq_euler_50_thresh0.15_lrg_0.75_phase1_ema0.95_cfg4.5
|
||||
33. lq_euler_50_thres0.1_linear_range_0.75_repro
|
||||
34. lq_euler_50_thres0.1_lrg_0.75_reproduc
|
||||
|
||||
35. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 1, learning rate 5e-6
|
||||
36. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 2, learning rate 1e-6
|
||||
37. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 2, learning rate 5e-6
|
||||
38. lq_euler_50_thres0.1_linear_range_0.75, learning rate 5e-6
|
||||
39. lq_euler_50_thres0.1_linear_range_0.75, learning rate 1e-5
|
||||
40. lq_euler_50_thres0.1_lrg_0.75_phase1_lr1e-6_repro
|
||||
|
||||
|
||||
41. lq_euler_50_thres0.1_lrg_0.75_reproduce
|
||||
42. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 4, learning rate 1e-6
|
||||
43. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 1, learning rate 1e-6, cfg 6.0
|
||||
44. lq_euler_50_thres0.1_lrg_0.75_phase1_lr_5e-6_test_norm
|
||||
45. lq_euler_50_thres0.1_lrg_0.75_phase1_lr_5e-6_pred_decay_0.1_latent14
|
||||
46-48. lq_euler_50_thres0.1_lrg_0.75_phase1_lr1e-6, l2 or l1, decay weight 0.1 to 0.001
|
||||
|
||||
49.
|
||||
@@ -1,10 +0,0 @@
|
||||
# Examples
|
||||
|
||||
- `generate_t2v.py` — text-to-video via the SDK (`fv2.load(...)` →
|
||||
`model.generate(...)` → `result.save(...)`). Card defaults, everything
|
||||
overridable by flag. The standalone reference implementation (no SDK, no
|
||||
runtime) lives at `fastvideo2/wan21/reference.py`.
|
||||
|
||||
The `--model` flag takes any catalog id — e.g. the 3-step FastWan students
|
||||
(`fastwan-qad-fp8-1.3b`, `fastwan-t2v-1.3b`); their step count, sampler, and
|
||||
sparsity/quant recipe come from the card, so no other flags change.
|
||||
@@ -1,47 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Text-to-video with the fastvideo2 SDK — the canonical example.
|
||||
|
||||
python examples/generate_t2v.py --prompt "a cat surfing a wave" --out cat.mp4
|
||||
|
||||
Loads the card resident once, generates with card defaults (50 steps, 81
|
||||
frames, 480x832 — override anything via flags), saves an mp4. Requires a CUDA
|
||||
box; weights resolve from the HF cache on first use.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
|
||||
import fastvideo2 as fv2
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__)
|
||||
p.add_argument("--prompt",
|
||||
default="a golden retriever puppy running through a sprinkler "
|
||||
"on a sunny lawn, water droplets sparkling, slow motion, cinematic")
|
||||
p.add_argument("--model", default="wan2.1-t2v-1.3b",
|
||||
help="a model id from the catalog (see `python -m fastvideo2 describe`)")
|
||||
p.add_argument("--out", default="out.mp4")
|
||||
p.add_argument("--seed", type=int, default=42)
|
||||
p.add_argument("--num-steps", dest="num_steps", type=int, default=None,
|
||||
help="unset -> the card's default")
|
||||
p.add_argument("--num-frames", dest="num_frames", type=int, default=None)
|
||||
p.add_argument("--guidance-scale", dest="guidance_scale", type=float, default=None)
|
||||
args = p.parse_args()
|
||||
|
||||
model = fv2.load(args.model)
|
||||
print(model)
|
||||
|
||||
overrides = {k: getattr(args, k) for k in ("seed", "num_steps", "num_frames", "guidance_scale")
|
||||
if getattr(args, k) is not None}
|
||||
result = model.generate(args.prompt, **overrides)
|
||||
|
||||
steps = [t for t in result.trace if "/denoise." in t["label"]]
|
||||
print(f"video {result.video.shape} | {len(steps)} denoise steps, "
|
||||
f"{sum(t['seconds'] for t in steps) / max(len(steps), 1):.2f}s/step, "
|
||||
f"{result.seconds:.1f}s total")
|
||||
print(f"saved -> {result.save(args.out)}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,85 @@
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from torchvision import transforms
|
||||
from torchvision.transforms import Lambda
|
||||
from fastvideo.dataset.t2v_datasets import T2V_dataset
|
||||
from fastvideo.dataset.latent_datasets import LatentDataset
|
||||
from fastvideo.dataset.transform import Normalize255, TemporalRandomCrop,CenterCropResizeVideo
|
||||
|
||||
def getdataset(args):
|
||||
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
|
||||
norm_fun = Lambda(lambda x: 2. * x - 1.)
|
||||
resize_topcrop = [CenterCropResizeVideo((args.max_height, args.max_width), top_crop=True), ]
|
||||
resize = [CenterCropResizeVideo((args.max_height, args.max_width)), ]
|
||||
transform = transforms.Compose([
|
||||
# Normalize255(),
|
||||
*resize,
|
||||
# RandomHorizontalFlipVideo(p=0.5), # in case their caption have position decription
|
||||
# norm_fun
|
||||
])
|
||||
transform_topcrop = transforms.Compose([
|
||||
Normalize255(),
|
||||
*resize_topcrop,
|
||||
# RandomHorizontalFlipVideo(p=0.5), # in case their caption have position decription
|
||||
norm_fun
|
||||
])
|
||||
# tokenizer = AutoTokenizer.from_pretrained("/storage/ongoing/new/Open-Sora-Plan/cache_dir/mt5-xxl", cache_dir=args.cache_dir)
|
||||
tokenizer = AutoTokenizer.from_pretrained(args.text_encoder_name, cache_dir=args.cache_dir)
|
||||
if args.dataset == 't2v':
|
||||
return T2V_dataset(args, transform=transform, temporal_sample=temporal_sample, tokenizer=tokenizer,
|
||||
transform_topcrop=transform_topcrop)
|
||||
|
||||
raise NotImplementedError(args.dataset)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from accelerate import Accelerator
|
||||
from fastvideo.dataset.t2v_datasets import dataset_prog
|
||||
import random
|
||||
from tqdm import tqdm
|
||||
args = type('args', (),
|
||||
{
|
||||
'ae': 'CausalVAEModel_4x8x8',
|
||||
'dataset': 't2v',
|
||||
'attention_mode': 'xformers',
|
||||
'use_rope': True,
|
||||
'text_max_length': 300,
|
||||
'max_height': 320,
|
||||
'max_width': 240,
|
||||
'num_frames': 1,
|
||||
'use_image_num': 0,
|
||||
'interpolation_scale_t': 1,
|
||||
'interpolation_scale_h': 1,
|
||||
'interpolation_scale_w': 1,
|
||||
'cache_dir': '../cache_dir',
|
||||
'image_data': '/storage/ongoing/new/Open-Sora-Plan-bak/7.14bak/scripts/train_data/image_data.txt',
|
||||
'video_data': '1',
|
||||
'train_fps': 24,
|
||||
'drop_short_ratio': 1.0,
|
||||
'use_img_from_vid': False,
|
||||
'speed_factor': 1.0,
|
||||
'cfg': 0.1,
|
||||
'text_encoder_name': 'google/mt5-xxl',
|
||||
'dataloader_num_workers': 10,
|
||||
|
||||
}
|
||||
)
|
||||
accelerator = Accelerator()
|
||||
dataset = getdataset(args)
|
||||
num = len(dataset_prog.img_cap_list)
|
||||
zero = 0
|
||||
for idx in tqdm(range(num)):
|
||||
image_data = dataset_prog.img_cap_list[idx]
|
||||
caps = [i['cap'] if isinstance(i['cap'], list) else [i['cap']] for i in image_data]
|
||||
try:
|
||||
caps = [[random.choice(i)] for i in caps]
|
||||
except Exception as e:
|
||||
print(e)
|
||||
# import ipdb;ipdb.set_trace()
|
||||
print(image_data)
|
||||
zero += 1
|
||||
continue
|
||||
assert caps[0] is not None and len(caps[0]) > 0
|
||||
print(num, zero)
|
||||
import ipdb;ipdb.set_trace()
|
||||
print('end')
|
||||
@@ -0,0 +1,83 @@
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
|
||||
class LatentDataset(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
json_path,
|
||||
num_latent_t,
|
||||
cfg_rate,
|
||||
):
|
||||
# data_merge_path: video_dir, latent_dir, prompt_embed_dir, json_path
|
||||
self.json_path = json_path
|
||||
self.cfg_rate = cfg_rate
|
||||
self.datase_dir_path = os.path.dirname(json_path)
|
||||
self.video_dir = os.path.join(self.datase_dir_path, "video")
|
||||
self.latent_dir = os.path.join(self.datase_dir_path, "latent")
|
||||
self.prompt_embed_dir = os.path.join(self.datase_dir_path, "prompt_embed")
|
||||
self.prompt_attention_mask_dir = os.path.join(self.datase_dir_path, "prompt_attention_mask")
|
||||
with open(self.json_path, 'r') as f:
|
||||
self.data_anno = json.load(f)
|
||||
# json.load(f) already keeps the order
|
||||
# self.data_anno = sorted(self.data_anno, key=lambda x: x['latent_path'])
|
||||
self.num_latent_t = num_latent_t
|
||||
# just zero embeddings [256, 4096]
|
||||
self.uncond_prompt_embed = torch.zeros(256, 4096).to(torch.float32)
|
||||
# 256 zeros
|
||||
self.uncond_prompt_mask = torch.zeros(256).bool()
|
||||
self.lengths = [data_item['length'] if "length" in data_item else 1 for data_item in self.data_anno]
|
||||
def __getitem__(self, idx):
|
||||
latent_file = self.data_anno[idx]["latent_path"]
|
||||
prompt_embed_file = self.data_anno[idx]["prompt_embed_path"]
|
||||
prompt_attention_mask_file = self.data_anno[idx]["prompt_attention_mask"]
|
||||
# load
|
||||
latent = torch.load(os.path.join(self.latent_dir, latent_file), map_location="cpu", weights_only=True)
|
||||
# TODO: Hack
|
||||
|
||||
latent = latent.squeeze(0)[:, -self.num_latent_t:]
|
||||
if random.random() < self.cfg_rate:
|
||||
prompt_embed = self.uncond_prompt_embed
|
||||
prompt_attention_mask = self.uncond_prompt_mask
|
||||
else:
|
||||
prompt_embed = torch.load(os.path.join(self.prompt_embed_dir, prompt_embed_file), map_location="cpu", weights_only=True)
|
||||
prompt_attention_mask = torch.load(os.path.join(self.prompt_attention_mask_dir, prompt_attention_mask_file), map_location="cpu", weights_only=True)
|
||||
return latent, prompt_embed, prompt_attention_mask
|
||||
|
||||
def __len__(self):
|
||||
return len(self.data_anno)
|
||||
|
||||
def latent_collate_function(batch):
|
||||
# return latent, prompt, latent_attn_mask, text_attn_mask
|
||||
# latent_attn_mask: # b t h w
|
||||
# text_attn_mask: b 1 l
|
||||
# needs to check if the latent/prompt' size and apply padding & attn mask
|
||||
latents, prompt_embeds, prompt_attention_masks = zip(*batch)
|
||||
# calculate max shape
|
||||
max_t = max([latent.shape[1] for latent in latents])
|
||||
max_h = max([latent.shape[2] for latent in latents])
|
||||
max_w = max([latent.shape[3] for latent in latents])
|
||||
|
||||
# padding
|
||||
latents = [torch.nn.functional.pad(latent, (0, max_t - latent.shape[1], 0, max_h - latent.shape[2], 0, max_w - latent.shape[3])) for latent in latents]
|
||||
# attn mask
|
||||
latent_attn_mask = torch.ones(len(latents), max_t, max_h, max_w)
|
||||
# set to 0 if padding
|
||||
for i, latent in enumerate(latents):
|
||||
latent_attn_mask[i, latent.shape[1]:, :, :] = 0
|
||||
latent_attn_mask[i, :, latent.shape[2]:, :] = 0
|
||||
latent_attn_mask[i, :, :, latent.shape[3]:] = 0
|
||||
|
||||
prompt_embeds = torch.stack(prompt_embeds, dim=0)
|
||||
prompt_attention_masks = torch.stack(prompt_attention_masks, dim=0)
|
||||
latents = torch.stack(latents, dim=0)
|
||||
return latents, prompt_embeds, latent_attn_mask, prompt_attention_masks
|
||||
|
||||
if __name__ == "__main__":
|
||||
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt", num_latent_t=28)
|
||||
dataloader = torch.utils.data.DataLoader(dataset, batch_size=2, shuffle=False, collate_fn=latent_collate_function)
|
||||
for latent, prompt_embed, latent_attn_mask, prompt_attention_mask in dataloader:
|
||||
print(latent.shape, prompt_embed.shape, latent_attn_mask.shape, prompt_attention_mask.shape)
|
||||
import pdb; pdb.set_trace()
|
||||
@@ -0,0 +1,312 @@
|
||||
import json
|
||||
import os, io, csv, math, random
|
||||
import numpy as np
|
||||
from einops import rearrange
|
||||
from decord import VideoReader
|
||||
from os.path import join as opj
|
||||
from collections import Counter
|
||||
|
||||
import torch
|
||||
from torch.utils.data.dataset import Dataset
|
||||
from torch.utils.data import DataLoader, Dataset, get_worker_info
|
||||
from tqdm import tqdm
|
||||
from PIL import Image
|
||||
from accelerate.logging import get_logger
|
||||
from fastvideo.utils.dataset_utils import DecordInit
|
||||
from fastvideo.utils.utils import text_preprocessing
|
||||
import torchvision
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class SingletonMeta(type):
|
||||
"""
|
||||
这是一个元类,用于创建单例类。
|
||||
"""
|
||||
_instances = {}
|
||||
|
||||
def __call__(cls, *args, **kwargs):
|
||||
if cls not in cls._instances:
|
||||
instance = super().__call__(*args, **kwargs)
|
||||
cls._instances[cls] = instance
|
||||
return cls._instances[cls]
|
||||
|
||||
|
||||
class DataSetProg(metaclass=SingletonMeta):
|
||||
def __init__(self):
|
||||
self.cap_list = []
|
||||
self.elements = []
|
||||
self.num_workers = 1
|
||||
self.n_elements = 0
|
||||
self.worker_elements = dict()
|
||||
self.n_used_elements = dict()
|
||||
|
||||
def set_cap_list(self, num_workers, cap_list, n_elements):
|
||||
self.num_workers = num_workers
|
||||
self.cap_list = cap_list
|
||||
self.n_elements = n_elements
|
||||
self.elements = list(range(n_elements))
|
||||
random.shuffle(self.elements)
|
||||
print(f"n_elements: {len(self.elements)}", flush=True)
|
||||
|
||||
for i in range(self.num_workers):
|
||||
self.n_used_elements[i] = 0
|
||||
per_worker = int(math.ceil(len(self.elements) / float(self.num_workers)))
|
||||
start = i * per_worker
|
||||
end = min(start + per_worker, len(self.elements))
|
||||
self.worker_elements[i] = self.elements[start: end]
|
||||
|
||||
def get_item(self, work_info):
|
||||
if work_info is None:
|
||||
worker_id = 0
|
||||
else:
|
||||
worker_id = work_info.id
|
||||
|
||||
idx = self.worker_elements[worker_id][self.n_used_elements[worker_id] % len(self.worker_elements[worker_id])]
|
||||
self.n_used_elements[worker_id] += 1
|
||||
return idx
|
||||
|
||||
|
||||
dataset_prog = DataSetProg()
|
||||
|
||||
def filter_resolution(h, w, max_h_div_w_ratio=17/16, min_h_div_w_ratio=8 / 16):
|
||||
if h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
|
||||
class T2V_dataset(Dataset):
|
||||
def __init__(self, args, transform, temporal_sample, tokenizer, transform_topcrop):
|
||||
self.data = args.data_merge_path
|
||||
self.num_frames = args.num_frames
|
||||
self.train_fps = args.train_fps
|
||||
self.use_image_num = args.use_image_num
|
||||
self.transform = transform
|
||||
self.transform_topcrop = transform_topcrop
|
||||
self.temporal_sample = temporal_sample
|
||||
self.tokenizer = tokenizer
|
||||
self.text_max_length = args.text_max_length
|
||||
self.cfg = args.cfg
|
||||
self.speed_factor = args.speed_factor
|
||||
self.max_height = args.max_height
|
||||
self.max_width = args.max_width
|
||||
self.drop_short_ratio = args.drop_short_ratio
|
||||
assert self.speed_factor >= 1
|
||||
self.v_decoder = DecordInit()
|
||||
self.video_length_tolerance_range = args.video_length_tolerance_range
|
||||
self.support_Chinese = True
|
||||
if not ('mt5' in args.text_encoder_name):
|
||||
self.support_Chinese = False
|
||||
|
||||
cap_list = self.get_cap_list()
|
||||
|
||||
assert len(cap_list) > 0
|
||||
cap_list, self.sample_num_frames = self.define_frame_index(cap_list)
|
||||
self.lengths = self.sample_num_frames
|
||||
|
||||
n_elements = len(cap_list)
|
||||
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list, n_elements)
|
||||
|
||||
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
|
||||
|
||||
def set_checkpoint(self, n_used_elements):
|
||||
for i in range(len(dataset_prog.n_used_elements)):
|
||||
dataset_prog.n_used_elements[i] = n_used_elements
|
||||
|
||||
def __len__(self):
|
||||
return dataset_prog.n_elements
|
||||
|
||||
def __getitem__(self, idx):
|
||||
try:
|
||||
data = self.get_data(idx)
|
||||
return data
|
||||
except Exception as e:
|
||||
logger.info(f'Error with {e}')
|
||||
if idx in dataset_prog.cap_list:
|
||||
logger.info(f"Caught an exception! {dataset_prog.cap_list[idx]}")
|
||||
return self.__getitem__(random.randint(0, self.__len__() - 1))
|
||||
|
||||
def get_data(self, idx):
|
||||
path = dataset_prog.cap_list[idx]['path']
|
||||
if path.endswith('.mp4'):
|
||||
return self.get_video(idx)
|
||||
else:
|
||||
return self.get_image(idx)
|
||||
|
||||
def get_video(self, idx):
|
||||
video_path = dataset_prog.cap_list[idx]['path']
|
||||
assert os.path.exists(video_path), f"file {video_path} do not exist!"
|
||||
frame_indices = dataset_prog.cap_list[idx]['sample_frame_index']
|
||||
torchvision_video, _, metadata = torchvision.io.read_video(video_path, output_format="TCHW")
|
||||
video = torchvision_video[frame_indices]
|
||||
video = self.transform(video)
|
||||
video = rearrange(video, 't c h w -> c t h w')
|
||||
video = video.to(torch.uint8)
|
||||
assert video.dtype == torch.uint8
|
||||
|
||||
h, w = video.shape[-2:]
|
||||
assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}'
|
||||
|
||||
video = video.float() / 127.5 - 1.0
|
||||
|
||||
text = dataset_prog.cap_list[idx]['cap']
|
||||
if not isinstance(text, list):
|
||||
text = [text]
|
||||
text = [random.choice(text)]
|
||||
|
||||
text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
padding='max_length',
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors='pt'
|
||||
)
|
||||
input_ids = text_tokens_and_mask['input_ids']
|
||||
cond_mask = text_tokens_and_mask['attention_mask']
|
||||
return dict(pixel_values=video, text=text, input_ids=input_ids, cond_mask=cond_mask, path=video_path)
|
||||
|
||||
def get_image(self, idx):
|
||||
image_data = dataset_prog.cap_list[idx] # [{'path': path, 'cap': cap}, ...]
|
||||
|
||||
image = Image.open(image_data['path']).convert('RGB') # [h, w, c]
|
||||
image = torch.from_numpy(np.array(image)) # [h, w, c]
|
||||
image = rearrange(image, 'h w c -> c h w').unsqueeze(0) # [1 c h w]
|
||||
# for i in image:
|
||||
# h, w = i.shape[-2:]
|
||||
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
|
||||
|
||||
image = self.transform_topcrop(image) if 'human_images' in image_data['path'] else self.transform(image) # [1 C H W] -> num_img [1 C H W]
|
||||
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
|
||||
|
||||
image = image.float() / 127.5 - 1.0
|
||||
|
||||
caps = image_data['cap'] if isinstance(image_data['cap'], list) else [image_data['cap']]
|
||||
caps = [random.choice(caps)]
|
||||
text = text_preprocessing(caps, support_Chinese=self.support_Chinese)
|
||||
input_ids, cond_mask = [], []
|
||||
text = text if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
padding='max_length',
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors='pt'
|
||||
)
|
||||
input_ids = text_tokens_and_mask['input_ids'] # 1, l
|
||||
cond_mask = text_tokens_and_mask['attention_mask'] # 1, l
|
||||
return dict(pixel_values=image, text=text, input_ids=input_ids, cond_mask=cond_mask, path=image_data['path'])
|
||||
|
||||
def define_frame_index(self, cap_list):
|
||||
|
||||
new_cap_list = []
|
||||
sample_num_frames = []
|
||||
cnt_too_long = 0
|
||||
cnt_too_short = 0
|
||||
cnt_no_cap = 0
|
||||
cnt_no_resolution = 0
|
||||
cnt_resolution_mismatch = 0
|
||||
cnt_movie = 0
|
||||
cnt_img = 0
|
||||
for i in cap_list:
|
||||
path = i['path']
|
||||
cap = i.get('cap', None)
|
||||
# ======no caption=====
|
||||
if cap is None:
|
||||
cnt_no_cap += 1
|
||||
continue
|
||||
if path.endswith('.mp4'):
|
||||
# ======no fps and duration=====
|
||||
duration = i.get('duration', None)
|
||||
fps = i.get('fps', None)
|
||||
if fps is None or duration is None:
|
||||
continue
|
||||
|
||||
# ======resolution mismatch=====
|
||||
resolution = i.get('resolution', None)
|
||||
if resolution is None:
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
else:
|
||||
if resolution.get('height', None) is None or resolution.get('width', None) is None:
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
height, width = i['resolution']['height'], i['resolution']['width']
|
||||
aspect = self.max_height / self.max_width
|
||||
hw_aspect_thr = 1.5
|
||||
is_pick = filter_resolution(height, width, max_h_div_w_ratio=hw_aspect_thr*aspect,
|
||||
min_h_div_w_ratio=1/hw_aspect_thr*aspect)
|
||||
if not is_pick:
|
||||
print("resolution mismatch")
|
||||
cnt_resolution_mismatch += 1
|
||||
continue
|
||||
|
||||
# import ipdb;ipdb.set_trace()
|
||||
i['num_frames'] = math.ceil(fps * duration)
|
||||
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
|
||||
if i['num_frames'] / fps > self.video_length_tolerance_range * (self.num_frames / self.train_fps * self.speed_factor): # too long video is not suitable for this training stage (self.num_frames)
|
||||
cnt_too_long += 1
|
||||
continue
|
||||
|
||||
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
|
||||
frame_interval = fps / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, i['num_frames'], frame_interval).astype(int)
|
||||
|
||||
# comment out it to enable dynamic frames training
|
||||
if len(frame_indices) < self.num_frames and random.random() < self.drop_short_ratio:
|
||||
cnt_too_short += 1
|
||||
continue
|
||||
|
||||
# too long video will be temporal-crop randomly
|
||||
if len(frame_indices) > self.num_frames:
|
||||
begin_index, end_index = self.temporal_sample(len(frame_indices))
|
||||
frame_indices = frame_indices[begin_index: end_index]
|
||||
# frame_indices = frame_indices[:self.num_frames] # head crop
|
||||
i['sample_frame_index'] = frame_indices.tolist()
|
||||
new_cap_list.append(i)
|
||||
i['sample_num_frames'] = len(i['sample_frame_index']) # will use in dataloader(group sampler)
|
||||
sample_num_frames.append(i['sample_num_frames'])
|
||||
elif path.endswith('.jpg'): # image
|
||||
cnt_img += 1
|
||||
new_cap_list.append(i)
|
||||
i['sample_num_frames'] = 1
|
||||
sample_num_frames.append(i['sample_num_frames'])
|
||||
else:
|
||||
raise NameError(f"Unknown file extention {path.split('.')[-1]}, only support .mp4 for video and .jpg for image")
|
||||
# import ipdb;ipdb.set_trace()
|
||||
logger.info(f'no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, '
|
||||
f'no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, '
|
||||
f'Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, '
|
||||
f'before filter: {len(cap_list)}, after filter: {len(new_cap_list)}')
|
||||
return new_cap_list, sample_num_frames
|
||||
|
||||
def decord_read(self, path, frame_indices):
|
||||
decord_vr = self.v_decoder(path)
|
||||
video_data = decord_vr.get_batch(frame_indices).asnumpy()
|
||||
video_data = torch.from_numpy(video_data)
|
||||
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
|
||||
return video_data
|
||||
|
||||
def read_jsons(self, data):
|
||||
cap_lists = []
|
||||
with open(data, 'r') as f:
|
||||
folder_anno = [i.strip().split(',') for i in f.readlines() if len(i.strip()) > 0]
|
||||
print(folder_anno)
|
||||
for folder, anno in folder_anno:
|
||||
with open(anno, 'r') as f:
|
||||
sub_list = json.load(f)
|
||||
logger.info(f'Building {anno}...')
|
||||
for i in range(len(sub_list)):
|
||||
sub_list[i]['path'] = opj(folder, sub_list[i]['path'])
|
||||
cap_lists += sub_list
|
||||
return cap_lists
|
||||
|
||||
|
||||
def get_cap_list(self):
|
||||
cap_lists = self.read_jsons(self.data)
|
||||
return cap_lists
|
||||
@@ -0,0 +1,591 @@
|
||||
import torch
|
||||
import random
|
||||
import numbers
|
||||
from torchvision.transforms import RandomCrop, RandomResizedCrop
|
||||
|
||||
|
||||
def _is_tensor_video_clip(clip):
|
||||
if not torch.is_tensor(clip):
|
||||
raise TypeError("clip should be Tensor. Got %s" % type(clip))
|
||||
|
||||
if not clip.ndimension() == 4:
|
||||
raise ValueError("clip should be 4D. Got %dD" % clip.dim())
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def center_crop_arr(pil_image, image_size):
|
||||
"""
|
||||
Center cropping implementation from ADM.
|
||||
https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126
|
||||
"""
|
||||
while min(*pil_image.size) >= 2 * image_size:
|
||||
pil_image = pil_image.resize(
|
||||
tuple(x // 2 for x in pil_image.size), resample=Image.BOX
|
||||
)
|
||||
|
||||
scale = image_size / min(*pil_image.size)
|
||||
pil_image = pil_image.resize(
|
||||
tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC
|
||||
)
|
||||
|
||||
arr = np.array(pil_image)
|
||||
crop_y = (arr.shape[0] - image_size) // 2
|
||||
crop_x = (arr.shape[1] - image_size) // 2
|
||||
return Image.fromarray(arr[crop_y: crop_y + image_size, crop_x: crop_x + image_size])
|
||||
|
||||
|
||||
def crop(clip, i, j, h, w):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
"""
|
||||
if len(clip.size()) != 4:
|
||||
raise ValueError("clip should be a 4D tensor")
|
||||
return clip[..., i: i + h, j: j + w]
|
||||
|
||||
|
||||
def resize(clip, target_size, interpolation_mode):
|
||||
if len(target_size) != 2:
|
||||
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
|
||||
return torch.nn.functional.interpolate(clip, size=target_size, mode=interpolation_mode, align_corners=True, antialias=True)
|
||||
|
||||
|
||||
def resize_scale(clip, target_size, interpolation_mode):
|
||||
if len(target_size) != 2:
|
||||
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
|
||||
H, W = clip.size(-2), clip.size(-1)
|
||||
scale_ = target_size[0] / min(H, W)
|
||||
return torch.nn.functional.interpolate(clip, scale_factor=scale_, mode=interpolation_mode, align_corners=True, antialias=True)
|
||||
|
||||
|
||||
def resized_crop(clip, i, j, h, w, size, interpolation_mode="bilinear"):
|
||||
"""
|
||||
Do spatial cropping and resizing to the video clip
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
i (int): i in (i,j) i.e coordinates of the upper left corner.
|
||||
j (int): j in (i,j) i.e coordinates of the upper left corner.
|
||||
h (int): Height of the cropped region.
|
||||
w (int): Width of the cropped region.
|
||||
size (tuple(int, int)): height and width of resized clip
|
||||
Returns:
|
||||
clip (torch.tensor): Resized and cropped clip. Size is (T, C, H, W)
|
||||
"""
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
clip = crop(clip, i, j, h, w)
|
||||
clip = resize(clip, size, interpolation_mode)
|
||||
return clip
|
||||
|
||||
|
||||
def center_crop(clip, crop_size):
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
th, tw = crop_size
|
||||
if h < th or w < tw:
|
||||
raise ValueError("height and width must be no smaller than crop_size")
|
||||
|
||||
i = int(round((h - th) / 2.0))
|
||||
j = int(round((w - tw) / 2.0))
|
||||
return crop(clip, i, j, th, tw)
|
||||
|
||||
|
||||
def center_crop_using_short_edge(clip):
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
if h < w:
|
||||
th, tw = h, h
|
||||
i = 0
|
||||
j = int(round((w - tw) / 2.0))
|
||||
else:
|
||||
th, tw = w, w
|
||||
i = int(round((h - th) / 2.0))
|
||||
j = 0
|
||||
return crop(clip, i, j, th, tw)
|
||||
|
||||
|
||||
|
||||
def center_crop_th_tw(clip, th, tw, top_crop):
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
|
||||
# import ipdb;ipdb.set_trace()
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
tr = th / tw
|
||||
if h / w > tr:
|
||||
new_h = int(w * tr)
|
||||
new_w = w
|
||||
else:
|
||||
new_h = h
|
||||
new_w = int(h / tr)
|
||||
|
||||
i = 0 if top_crop else int(round((h - new_h) / 2.0))
|
||||
j = int(round((w - new_w) / 2.0))
|
||||
return crop(clip, i, j, new_h, new_w)
|
||||
|
||||
def random_shift_crop(clip):
|
||||
'''
|
||||
Slide along the long edge, with the short edge as crop size
|
||||
'''
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
|
||||
if h <= w:
|
||||
long_edge = w
|
||||
short_edge = h
|
||||
else:
|
||||
long_edge = h
|
||||
short_edge = w
|
||||
|
||||
th, tw = short_edge, short_edge
|
||||
|
||||
i = torch.randint(0, h - th + 1, size=(1,)).item()
|
||||
j = torch.randint(0, w - tw + 1, size=(1,)).item()
|
||||
return crop(clip, i, j, th, tw)
|
||||
|
||||
|
||||
def normalize_video(clip):
|
||||
"""
|
||||
Convert tensor data type from uint8 to float, divide value by 255.0 and
|
||||
permute the dimensions of clip tensor
|
||||
Args:
|
||||
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
|
||||
Return:
|
||||
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
|
||||
"""
|
||||
_is_tensor_video_clip(clip)
|
||||
if not clip.dtype == torch.uint8:
|
||||
raise TypeError("clip tensor should have data type uint8. Got %s" % str(clip.dtype))
|
||||
# return clip.float().permute(3, 0, 1, 2) / 255.0
|
||||
return clip.float() / 255.0
|
||||
|
||||
|
||||
def normalize(clip, mean, std, inplace=False):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
|
||||
mean (tuple): pixel RGB mean. Size is (3)
|
||||
std (tuple): pixel standard deviation. Size is (3)
|
||||
Returns:
|
||||
normalized clip (torch.tensor): Size is (T, C, H, W)
|
||||
"""
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
if not inplace:
|
||||
clip = clip.clone()
|
||||
mean = torch.as_tensor(mean, dtype=clip.dtype, device=clip.device)
|
||||
# print(mean)
|
||||
std = torch.as_tensor(std, dtype=clip.dtype, device=clip.device)
|
||||
clip.sub_(mean[:, None, None, None]).div_(std[:, None, None, None])
|
||||
return clip
|
||||
|
||||
|
||||
def hflip(clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be normalized. Size is (T, C, H, W)
|
||||
Returns:
|
||||
flipped clip (torch.tensor): Size is (T, C, H, W)
|
||||
"""
|
||||
if not _is_tensor_video_clip(clip):
|
||||
raise ValueError("clip should be a 4D torch.tensor")
|
||||
return clip.flip(-1)
|
||||
|
||||
|
||||
class RandomCropVideo:
|
||||
def __init__(self, size):
|
||||
if isinstance(size, numbers.Number):
|
||||
self.size = (int(size), int(size))
|
||||
else:
|
||||
self.size = size
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: randomly cropped video clip.
|
||||
size is (T, C, OH, OW)
|
||||
"""
|
||||
i, j, h, w = self.get_params(clip)
|
||||
return crop(clip, i, j, h, w)
|
||||
|
||||
def get_params(self, clip):
|
||||
h, w = clip.shape[-2:]
|
||||
th, tw = self.size
|
||||
|
||||
if h < th or w < tw:
|
||||
raise ValueError(f"Required crop size {(th, tw)} is larger than input image size {(h, w)}")
|
||||
|
||||
if w == tw and h == th:
|
||||
return 0, 0, h, w
|
||||
|
||||
i = torch.randint(0, h - th + 1, size=(1,)).item()
|
||||
j = torch.randint(0, w - tw + 1, size=(1,)).item()
|
||||
|
||||
return i, j, th, tw
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size})"
|
||||
|
||||
|
||||
class SpatialStrideCropVideo:
|
||||
def __init__(self, stride):
|
||||
self.stride = stride
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: cropped video clip by stride.
|
||||
size is (T, C, OH, OW)
|
||||
"""
|
||||
i, j, h, w = self.get_params(clip)
|
||||
return crop(clip, i, j, h, w)
|
||||
|
||||
def get_params(self, clip):
|
||||
h, w = clip.shape[-2:]
|
||||
|
||||
th, tw = h // self.stride * self.stride, w // self.stride * self.stride
|
||||
|
||||
return 0, 0, th, tw # from top-left
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size})"
|
||||
|
||||
class LongSideResizeVideo:
|
||||
'''
|
||||
First use the long side,
|
||||
then resize to the specified size
|
||||
'''
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
skip_low_resolution=False,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
self.size = size
|
||||
self.skip_low_resolution = skip_low_resolution
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: scale resized video clip.
|
||||
size is (T, C, 512, *) or (T, C, *, 512)
|
||||
"""
|
||||
_, _, h, w = clip.shape
|
||||
if self.skip_low_resolution and max(h, w) <= self.size:
|
||||
return clip
|
||||
if h > w:
|
||||
w = int(w * self.size / h)
|
||||
h = self.size
|
||||
else:
|
||||
h = int(h * self.size / w)
|
||||
w = self.size
|
||||
resize_clip = resize(clip, target_size=(h, w),
|
||||
interpolation_mode=self.interpolation_mode)
|
||||
return resize_clip
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
|
||||
|
||||
class CenterCropResizeVideo:
|
||||
'''
|
||||
First use the short side for cropping length,
|
||||
center crop video, then resize to the specified size
|
||||
'''
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
top_crop=False,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
self.top_crop = top_crop
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: scale resized / center cropped video clip.
|
||||
size is (T, C, crop_size, crop_size)
|
||||
"""
|
||||
# clip_center_crop = center_crop_using_short_edge(clip)
|
||||
clip_center_crop = center_crop_th_tw(clip, self.size[0], self.size[1], top_crop=self.top_crop)
|
||||
# import ipdb;ipdb.set_trace()
|
||||
clip_center_crop_resize = resize(clip_center_crop, target_size=self.size,
|
||||
interpolation_mode=self.interpolation_mode)
|
||||
return clip_center_crop_resize
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
|
||||
|
||||
|
||||
class UCFCenterCropVideo:
|
||||
'''
|
||||
First scale to the specified size in equal proportion to the short edge,
|
||||
then center cropping
|
||||
'''
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: scale resized / center cropped video clip.
|
||||
size is (T, C, crop_size, crop_size)
|
||||
"""
|
||||
clip_resize = resize_scale(clip=clip, target_size=self.size, interpolation_mode=self.interpolation_mode)
|
||||
clip_center_crop = center_crop(clip_resize, self.size)
|
||||
return clip_center_crop
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
|
||||
|
||||
|
||||
class KineticsRandomCropResizeVideo:
|
||||
'''
|
||||
Slide along the long edge, with the short edge as crop size. And resie to the desired size.
|
||||
'''
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip):
|
||||
clip_random_crop = random_shift_crop(clip)
|
||||
clip_resize = resize(clip_random_crop, self.size, self.interpolation_mode)
|
||||
return clip_resize
|
||||
|
||||
|
||||
class CenterCropVideo:
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
|
||||
self.interpolation_mode = interpolation_mode
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Video clip to be cropped. Size is (T, C, H, W)
|
||||
Returns:
|
||||
torch.tensor: center cropped video clip.
|
||||
size is (T, C, crop_size, crop_size)
|
||||
"""
|
||||
clip_center_crop = center_crop(clip, self.size)
|
||||
return clip_center_crop
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(size={self.size}, interpolation_mode={self.interpolation_mode}"
|
||||
|
||||
|
||||
class Normalize:
|
||||
"""
|
||||
Normalize the video clip by mean subtraction and division by standard deviation
|
||||
Args:
|
||||
mean (3-tuple): pixel RGB mean
|
||||
std (3-tuple): pixel RGB standard deviation
|
||||
inplace (boolean): whether do in-place normalization
|
||||
"""
|
||||
|
||||
def __init__(self, mean, std, inplace=False):
|
||||
self.mean = mean
|
||||
self.std = std
|
||||
self.inplace = inplace
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): video clip must be normalized. Size is (C, T, H, W)
|
||||
"""
|
||||
return normalize(clip, self.mean, self.std, self.inplace)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(mean={self.mean}, std={self.std}, inplace={self.inplace})"
|
||||
|
||||
|
||||
class Normalize255:
|
||||
"""
|
||||
Convert tensor data type from uint8 to float, divide value by 255.0 and
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor, dtype=torch.uint8): Size is (T, C, H, W)
|
||||
Return:
|
||||
clip (torch.tensor, dtype=torch.float): Size is (T, C, H, W)
|
||||
"""
|
||||
return normalize_video(clip)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return self.__class__.__name__
|
||||
|
||||
|
||||
class RandomHorizontalFlipVideo:
|
||||
"""
|
||||
Flip the video clip along the horizontal direction with a given probability
|
||||
Args:
|
||||
p (float): probability of the clip being flipped. Default value is 0.5
|
||||
"""
|
||||
|
||||
def __init__(self, p=0.5):
|
||||
self.p = p
|
||||
|
||||
def __call__(self, clip):
|
||||
"""
|
||||
Args:
|
||||
clip (torch.tensor): Size is (T, C, H, W)
|
||||
Return:
|
||||
clip (torch.tensor): Size is (T, C, H, W)
|
||||
"""
|
||||
if random.random() < self.p:
|
||||
clip = hflip(clip)
|
||||
return clip
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}(p={self.p})"
|
||||
|
||||
|
||||
# ------------------------------------------------------------
|
||||
# --------------------- Sampling ---------------------------
|
||||
# ------------------------------------------------------------
|
||||
class TemporalRandomCrop(object):
|
||||
"""Temporally crop the given frame indices at a random location.
|
||||
|
||||
Args:
|
||||
size (int): Desired length of frames will be seen in the model.
|
||||
"""
|
||||
|
||||
def __init__(self, size):
|
||||
self.size = size
|
||||
|
||||
def __call__(self, total_frames):
|
||||
rand_end = max(0, total_frames - self.size - 1)
|
||||
begin_index = random.randint(0, rand_end)
|
||||
end_index = min(begin_index + self.size, total_frames)
|
||||
return begin_index, end_index
|
||||
|
||||
class DynamicSampleDuration(object):
|
||||
"""Temporally crop the given frame indices at a random location.
|
||||
|
||||
Args:
|
||||
size (int): Desired length of frames will be seen in the model.
|
||||
"""
|
||||
|
||||
def __init__(self, t_stride, extra_1):
|
||||
self.t_stride = t_stride
|
||||
self.extra_1 = extra_1
|
||||
|
||||
def __call__(self, t, h, w):
|
||||
if self.extra_1:
|
||||
t = t - 1
|
||||
truncate_t_list = list(range(t+1))[t//2:][::self.t_stride] # need half at least
|
||||
truncate_t = random.choice(truncate_t_list)
|
||||
if self.extra_1:
|
||||
truncate_t = truncate_t + 1
|
||||
return 0, truncate_t
|
||||
|
||||
if __name__ == '__main__':
|
||||
from torchvision import transforms
|
||||
import torchvision.io as io
|
||||
import numpy as np
|
||||
from torchvision.utils import save_image
|
||||
import os
|
||||
|
||||
vframes, aframes, info = io.read_video(
|
||||
filename='./v_Archery_g01_c03.avi',
|
||||
pts_unit='sec',
|
||||
output_format='TCHW'
|
||||
)
|
||||
|
||||
trans = transforms.Compose([
|
||||
Normalize255(),
|
||||
RandomHorizontalFlipVideo(),
|
||||
UCFCenterCropVideo(512),
|
||||
# NormalizeVideo(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True)
|
||||
])
|
||||
|
||||
target_video_len = 32
|
||||
frame_interval = 1
|
||||
total_frames = len(vframes)
|
||||
print(total_frames)
|
||||
|
||||
temporal_sample = TemporalRandomCrop(target_video_len * frame_interval)
|
||||
|
||||
# Sampling video frames
|
||||
start_frame_ind, end_frame_ind = temporal_sample(total_frames)
|
||||
# print(start_frame_ind)
|
||||
# print(end_frame_ind)
|
||||
assert end_frame_ind - start_frame_ind >= target_video_len
|
||||
frame_indice = np.linspace(start_frame_ind, end_frame_ind - 1, target_video_len, dtype=int)
|
||||
print(frame_indice)
|
||||
|
||||
select_vframes = vframes[frame_indice]
|
||||
print(select_vframes.shape)
|
||||
print(select_vframes.dtype)
|
||||
|
||||
select_vframes_trans = trans(select_vframes)
|
||||
print(select_vframes_trans.shape)
|
||||
print(select_vframes_trans.dtype)
|
||||
|
||||
select_vframes_trans_int = ((select_vframes_trans * 0.5 + 0.5) * 255).to(dtype=torch.uint8)
|
||||
print(select_vframes_trans_int.dtype)
|
||||
print(select_vframes_trans_int.permute(0, 2, 3, 1).shape)
|
||||
|
||||
io.write_video('./test.avi', select_vframes_trans_int.permute(0, 2, 3, 1), fps=8)
|
||||
|
||||
for i in range(target_video_len):
|
||||
save_image(select_vframes_trans[i], os.path.join('./test000', '%04d.png' % i), normalize=True,
|
||||
value_range=(-1, 1))
|
||||
@@ -0,0 +1,719 @@
|
||||
import argparse
|
||||
from email.policy import strict
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, \
|
||||
destroy_sequence_parallel_group, get_sequence_parallel_state, nccl_info
|
||||
from fastvideo.utils.communications import sp_parallel_dataloader_wrapper, broadcast
|
||||
from fastvideo.model.mochi_latents_utils import normalize_mochi_dit_input
|
||||
from fastvideo.utils.validation import log_validation
|
||||
import time
|
||||
from torch.utils.data import DataLoader
|
||||
import torch
|
||||
from torch.distributed.fsdp import ShardingStrategy
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
StateDictType,
|
||||
FullStateDictConfig,
|
||||
)
|
||||
from fastvideo.model.pipeline_mochi import linear_quadratic_schedule
|
||||
import json
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from fastvideo.utils.dataset_utils import LengthGroupedSampler
|
||||
import wandb
|
||||
from accelerate.utils import set_seed
|
||||
from tqdm.auto import tqdm
|
||||
from fastvideo.fsdp_util import get_dit_fsdp_kwargs, apply_fsdp_checkpointing
|
||||
import diffusers
|
||||
from diffusers import (
|
||||
FlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
|
||||
from copy import deepcopy
|
||||
from diffusers.optimization import get_scheduler
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
from diffusers.utils import check_min_version
|
||||
from fastvideo.dataset.latent_datasets import LatentDataset, latent_collate_function
|
||||
import torch.distributed as dist
|
||||
from safetensors.torch import save_file, load_file
|
||||
from peft import LoraConfig, inject_adapter_in_model
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
)
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
check_min_version("0.31.0")
|
||||
import time
|
||||
from collections import deque
|
||||
|
||||
import sys
|
||||
import pdb
|
||||
#ForkedPdb().set_trace()
|
||||
class ForkedPdb(pdb.Pdb):
|
||||
"""A Pdb subclass that may be used
|
||||
from a forked multiprocessing child
|
||||
|
||||
"""
|
||||
def interaction(self, *args, **kwargs):
|
||||
_stdin = sys.stdin
|
||||
try:
|
||||
sys.stdin = open('/dev/stdin')
|
||||
pdb.Pdb.interaction(self, *args, **kwargs)
|
||||
finally:
|
||||
sys.stdin = _stdin
|
||||
|
||||
def main_print(content):
|
||||
if int(os.environ['LOCAL_RANK']) <= 0:
|
||||
print(content)
|
||||
|
||||
def save_checkpoint(transformer: MochiTransformer3DModel, rank, output_dir, step):
|
||||
main_print(f"--> saving checkpoint at step {step}")
|
||||
with FSDP.state_dict_type(
|
||||
transformer, StateDictType.FULL_STATE_DICT, FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
):
|
||||
cpu_state = transformer.state_dict()
|
||||
#todo move to get_state_dict
|
||||
if rank <= 0:
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
|
||||
save_file(cpu_state, weight_path)
|
||||
config_dict = dict(transformer.config)
|
||||
config_path = os.path.join(save_dir, "config.json")
|
||||
# save dict as json
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
main_print(f"--> checkpoint saved at step {step}")
|
||||
|
||||
|
||||
def reshard_fsdp(model):
|
||||
for m in FSDP.fsdp_modules(model):
|
||||
if m._has_params and m.sharding_strategy is not ShardingStrategy.NO_SHARD:
|
||||
torch.distributed.fsdp._runtime_utils._reshard(m, m._handle, True)
|
||||
|
||||
def get_norm(model_pred, norms, gradient_accumulation_steps):
|
||||
fro_norm = torch.linalg.matrix_norm(model_pred, ord="fro") / gradient_accumulation_steps
|
||||
largest_singular_value = torch.linalg.matrix_norm(model_pred, ord=2) / gradient_accumulation_steps
|
||||
absolute_mean = torch.mean(torch.abs(model_pred)) / gradient_accumulation_steps
|
||||
absolute_max = torch.max(torch.abs(model_pred)) / gradient_accumulation_steps
|
||||
dist.all_reduce(fro_norm, op=dist.ReduceOp.AVG)
|
||||
dist.all_reduce(largest_singular_value, op=dist.ReduceOp.AVG)
|
||||
dist.all_reduce(absolute_mean, op=dist.ReduceOp.AVG)
|
||||
norms["fro"] += torch.mean(fro_norm).item()
|
||||
norms["largest singular value"] += torch.mean(largest_singular_value).item()
|
||||
norms["absolute mean"] += absolute_mean.item()
|
||||
norms["absolute max"] += absolute_max.item()
|
||||
|
||||
def train_one_step_mochi(transformer, teacher_transformer, ema_transformer, optimizer, lr_scheduler,loader, noise_scheduler, solver,noise_random_generator, gradient_accumulation_steps, sp_size, precondition_outputs, max_grad_norm, uncond_prompt_embed, uncond_prompt_mask, num_euler_timesteps, multiphase, not_apply_cfg_solver, distill_cfg, ema_decay, pred_decay_weight, pred_decay_type):
|
||||
total_loss = 0.0
|
||||
optimizer.zero_grad()
|
||||
model_pred_norm = {"fro": 0.0, "largest singular value": 0.0, "absolute mean": 0.0, "absolute max": 0.0}
|
||||
for _ in range(gradient_accumulation_steps):
|
||||
latents, encoder_hidden_states, latents_attention_mask, encoder_attention_mask = next(loader)
|
||||
model_input = normalize_mochi_dit_input(latents)
|
||||
noise = torch.randn_like(model_input)
|
||||
bsz = model_input.shape[0]
|
||||
index = torch.randint(
|
||||
0, num_euler_timesteps, (bsz,), device=model_input.device
|
||||
).long()
|
||||
if sp_size > 1:
|
||||
broadcast(index)
|
||||
# Add noise according to flow matching.
|
||||
# sigmas = get_sigmas(start_timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
|
||||
sigmas = extract_into_tensor(solver.sigmas, index, model_input.shape)
|
||||
sigmas_prev = extract_into_tensor(
|
||||
solver.sigmas_prev, index, model_input.shape
|
||||
)
|
||||
|
||||
timesteps = (
|
||||
sigmas * noise_scheduler.config.num_train_timesteps
|
||||
).view(-1)
|
||||
# if squeeze to [], unsqueeze to [1]
|
||||
|
||||
timesteps_prev = (
|
||||
sigmas_prev * noise_scheduler.config.num_train_timesteps
|
||||
).view(-1)
|
||||
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
|
||||
|
||||
# Predict the noise residual
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
|
||||
model_pred = transformer(
|
||||
noisy_model_input,
|
||||
encoder_hidden_states,
|
||||
timesteps,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict= False
|
||||
)[0]
|
||||
|
||||
# if accelerator.is_main_process:
|
||||
model_pred, end_index = solver.euler_style_multiphase_pred(
|
||||
noisy_model_input, model_pred, index, multiphase
|
||||
)
|
||||
with torch.no_grad():
|
||||
w = distill_cfg
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
cond_teacher_output = teacher_transformer(
|
||||
noisy_model_input,
|
||||
encoder_hidden_states,
|
||||
timesteps,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict= False
|
||||
)[0].float()
|
||||
if not_apply_cfg_solver:
|
||||
uncond_teacher_output = cond_teacher_output
|
||||
else:
|
||||
# Get teacher model prediction on noisy_latents and unconditional embedding
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
uncond_teacher_output = teacher_transformer(
|
||||
noisy_model_input,
|
||||
uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
|
||||
timesteps,
|
||||
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
|
||||
return_dict= False
|
||||
)[0].float()
|
||||
teacher_output = cond_teacher_output + w * (
|
||||
cond_teacher_output - uncond_teacher_output
|
||||
)
|
||||
x_prev = solver.euler_step(
|
||||
noisy_model_input, teacher_output, index
|
||||
)
|
||||
|
||||
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
|
||||
with torch.no_grad():
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
if ema_transformer is not None:
|
||||
target_pred = ema_transformer(
|
||||
x_prev.float(),
|
||||
encoder_hidden_states,
|
||||
timesteps_prev,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict= False
|
||||
)[0]
|
||||
else:
|
||||
target_pred = transformer(
|
||||
x_prev.float(),
|
||||
encoder_hidden_states,
|
||||
timesteps_prev,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict= False
|
||||
)[0]
|
||||
|
||||
target, end_index = solver.euler_style_multiphase_pred(
|
||||
x_prev, target_pred, index, multiphase, True
|
||||
)
|
||||
|
||||
|
||||
huber_c = 0.001
|
||||
# loss = loss.mean()
|
||||
loss = torch.mean(
|
||||
torch.sqrt(
|
||||
(model_pred.float() - target.float()) ** 2 + huber_c**2
|
||||
)
|
||||
- huber_c
|
||||
) / gradient_accumulation_steps
|
||||
if pred_decay_weight > 0:
|
||||
if pred_decay_type == "l1":
|
||||
pred_decay_loss = torch.mean(torch.sqrt(model_pred.float() ** 2 )) * pred_decay_weight / gradient_accumulation_steps
|
||||
loss += pred_decay_loss
|
||||
elif pred_decay_type == "l2":
|
||||
# essnetially k2?
|
||||
pred_decay_loss = torch.mean(model_pred.float() ** 2 ) * pred_decay_weight / gradient_accumulation_steps
|
||||
loss += pred_decay_loss
|
||||
else:
|
||||
assert NotImplementedError("pred_decay_type is not implemented")
|
||||
|
||||
# calculate model_pred norm and mean
|
||||
get_norm(model_pred.detach().float(), model_pred_norm, gradient_accumulation_steps)
|
||||
loss.backward()
|
||||
|
||||
avg_loss = loss.detach().clone()
|
||||
dist.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
|
||||
# dist.all_reduce(pred_decay_loss.detach(), op=dist.ReduceOp.AVG)
|
||||
total_loss += avg_loss.item()
|
||||
|
||||
|
||||
# update ema
|
||||
if ema_transformer is not None:
|
||||
reshard_fsdp(ema_transformer)
|
||||
for p_averaged, p_model in zip(ema_transformer.parameters(), transformer.parameters()):
|
||||
with torch.no_grad():
|
||||
p_averaged.copy_(torch.lerp(p_averaged.detach(), p_model.detach(), 1 - ema_decay))
|
||||
|
||||
grad_norm = transformer.clip_grad_norm_(max_grad_norm)
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
|
||||
|
||||
return total_loss, grad_norm.item(), model_pred_norm
|
||||
|
||||
def get_lora_model(transformer, lora_config):
|
||||
transformer.requires_grad_(False)
|
||||
transformer = inject_adapter_in_model(lora_config, transformer)
|
||||
return transformer
|
||||
|
||||
def save_lora_checkpoint(
|
||||
transformer: MochiTransformer3DModel,
|
||||
optimizer,
|
||||
rank,
|
||||
output_dir,
|
||||
step
|
||||
):
|
||||
main_print(f"--> saving LoRA checkpoint at step {step}")
|
||||
with FSDP.state_dict_type(
|
||||
transformer,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
):
|
||||
full_state_dict = transformer.state_dict()
|
||||
lora_state_dict = {
|
||||
k: v for k, v in full_state_dict.items()
|
||||
if 'lora' in k.lower()
|
||||
}
|
||||
lora_optim_state = FSDP.optim_state_dict(
|
||||
transformer,
|
||||
optimizer,
|
||||
)
|
||||
if rank <= 0:
|
||||
save_dir = os.path.join(output_dir, f"lora-checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
weight_path = os.path.join(save_dir, "lora_weights.safetensors")
|
||||
save_file(lora_state_dict, weight_path)
|
||||
optim_path = os.path.join(save_dir, "lora_optimizer.pt")
|
||||
torch.save(lora_optim_state, optim_path)
|
||||
lora_config = {
|
||||
'step': step,
|
||||
'lora_params': {
|
||||
'lora_rank': transformer.config.lora_rank,
|
||||
'lora_alpha': transformer.config.lora_alpha,
|
||||
'target_modules': transformer.config.lora_target_modules
|
||||
}
|
||||
}
|
||||
config_path = os.path.join(save_dir, "lora_config.json")
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(lora_config, f, indent=4)
|
||||
main_print(f"--> LoRA checkpoint saved at step {step}")
|
||||
|
||||
def resume_lora_training(
|
||||
transformer,
|
||||
checkpoint_dir,
|
||||
optimizer
|
||||
):
|
||||
weight_path = os.path.join(checkpoint_dir, "lora_weights.safetensors")
|
||||
lora_weights = load_file(weight_path)
|
||||
config_path = os.path.join(checkpoint_dir, "lora_config.json")
|
||||
with open(config_path, "r") as f:
|
||||
config_dict = json.load(f)
|
||||
with FSDP.state_dict_type(
|
||||
transformer,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
):
|
||||
current_state = transformer.state_dict()
|
||||
current_state.update(lora_weights)
|
||||
transformer.load_state_dict(current_state, strict=False)
|
||||
optim_path = os.path.join(checkpoint_dir, "lora_optimizer.pt")
|
||||
optimizer_state_dict = torch.load(optim_path, weights_only=False)
|
||||
optim_state = FSDP.optim_state_dict_to_load(
|
||||
model=transformer,
|
||||
optim=optimizer,
|
||||
optim_state_dict=optimizer_state_dict
|
||||
)
|
||||
optimizer.load_state_dict(optim_state)
|
||||
step = config_dict['step']
|
||||
main_print(f"--> Successfully resuming LoRA training from step {step}")
|
||||
return transformer, optimizer, step
|
||||
|
||||
def main(args):
|
||||
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
|
||||
local_rank = int(os.environ['LOCAL_RANK'])
|
||||
rank = int(os.environ['RANK'])
|
||||
world_size = int(os.environ['WORLD_SIZE'])
|
||||
dist.init_process_group("nccl")
|
||||
torch.cuda.set_device(local_rank)
|
||||
device = torch.cuda.current_device()
|
||||
initialize_sequence_parallel_state(args.sp_size)
|
||||
|
||||
# If passed along, set the training seed now. On GPU...
|
||||
if args.seed is not None:
|
||||
# TODO: t within the same seq parallel group should be the same. Noise should be different.
|
||||
set_seed(args.seed + rank)
|
||||
# We use different seeds for the noise generation in each process to ensure that the noise is different in a batch.
|
||||
noise_random_generator = None
|
||||
|
||||
# Handle the repository creation
|
||||
if rank <=0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weigths to half-precision
|
||||
# as these weights are only used for inference, keeping weights in full precision is not required.
|
||||
|
||||
# Create model:
|
||||
|
||||
main_print(f"--> loading model from {args.pretrained_model_name_or_path}")
|
||||
# keep the master weight to float32
|
||||
if args.dit_model_name_or_path:
|
||||
transformer = transformer = MochiTransformer3DModel.from_pretrained(
|
||||
args.dit_model_name_or_path,
|
||||
torch_dtype = torch.float32,
|
||||
#torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
else:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(
|
||||
args.pretrained_model_name_or_path,
|
||||
subfolder="transformer",
|
||||
torch_dtype = torch.float32,
|
||||
#torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
teacher_transformer = deepcopy(transformer)
|
||||
if args.use_ema:
|
||||
ema_transformer = deepcopy(transformer)
|
||||
else:
|
||||
ema_transformer = None
|
||||
if args.use_lora:
|
||||
lora_config = LoraConfig(
|
||||
r=args.lora_rank,
|
||||
lora_alpha=args.lora_alpha,
|
||||
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
|
||||
init_lora_weights=True,
|
||||
)
|
||||
transformer = get_lora_model(transformer, lora_config)
|
||||
|
||||
main_print(f" Total training parameters = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e6} M")
|
||||
main_print(f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}")
|
||||
fsdp_kwargs = get_dit_fsdp_kwargs(args.fsdp_sharding_startegy, args.use_lora, args.use_cpu_offload)
|
||||
|
||||
if args.use_lora:
|
||||
transformer.config.lora_rank = args.lora_rank
|
||||
transformer.config.lora_alpha = args.lora_alpha
|
||||
transformer.config.lora_target_modules = ["to_k", "to_q", "to_v", "to_out.0"]
|
||||
transformer._no_split_modules = ["MochiTransformerBlock"]
|
||||
fsdp_kwargs['auto_wrap_policy'] = fsdp_kwargs['auto_wrap_policy'](transformer)
|
||||
|
||||
|
||||
transformer = FSDP(
|
||||
transformer,
|
||||
**fsdp_kwargs,
|
||||
)
|
||||
teacher_transformer = FSDP(
|
||||
teacher_transformer,
|
||||
**fsdp_kwargs,
|
||||
)
|
||||
if args.use_ema:
|
||||
ema_transformer = FSDP(
|
||||
ema_transformer,
|
||||
**fsdp_kwargs,
|
||||
)
|
||||
main_print(f"--> model loaded")
|
||||
|
||||
if args.gradient_checkpointing:
|
||||
apply_fsdp_checkpointing(transformer, args.selective_checkpointing)
|
||||
apply_fsdp_checkpointing(teacher_transformer, args.selective_checkpointing)
|
||||
if args.use_ema:
|
||||
apply_fsdp_checkpointing(ema_transformer, args.selective_checkpointing)
|
||||
# Set model as trainable.
|
||||
transformer.train()
|
||||
teacher_transformer.requires_grad_(False)
|
||||
if args.use_ema:
|
||||
ema_transformer.requires_grad_(False)
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
|
||||
if args.scheduler_type == "pcm_linear_quadratic":
|
||||
linear_steps = int(noise_scheduler.config.num_train_timesteps * args.linear_range)
|
||||
sigmas = linear_quadratic_schedule(noise_scheduler.config.num_train_timesteps, args.linear_quadratic_threshold, linear_steps)
|
||||
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
|
||||
else:
|
||||
sigmas = noise_scheduler.sigmas
|
||||
solver = EulerSolver(
|
||||
sigmas.numpy()[::-1],
|
||||
noise_scheduler.config.num_train_timesteps,
|
||||
euler_timesteps=args.num_euler_timesteps,
|
||||
)
|
||||
solver.to(device)
|
||||
params_to_optimize = transformer.parameters()
|
||||
params_to_optimize = list(filter(lambda p: p.requires_grad, params_to_optimize))
|
||||
|
||||
optimizer = torch.optim.AdamW(
|
||||
params_to_optimize,
|
||||
lr=args.learning_rate,
|
||||
betas=(0.9,0.999),
|
||||
weight_decay=args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
|
||||
init_steps = 0
|
||||
if args.resume_from_lora_checkpoint:
|
||||
transformer, optimizer, init_steps = resume_lora_training(
|
||||
transformer, args.resume_from_lora_checkpoint, optimizer
|
||||
)
|
||||
main_print(f"optimizer: {optimizer}")
|
||||
|
||||
#todo add lr scheduler
|
||||
lr_scheduler = get_scheduler(
|
||||
args.lr_scheduler,
|
||||
optimizer=optimizer,
|
||||
num_warmup_steps=args.lr_warmup_steps * world_size,
|
||||
num_training_steps=args.max_train_steps * world_size,
|
||||
num_cycles=args.lr_num_cycles,
|
||||
power=args.lr_power,
|
||||
last_epoch=init_steps - 1,
|
||||
)
|
||||
|
||||
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t, args.cfg)
|
||||
uncond_prompt_embed = train_dataset.uncond_prompt_embed
|
||||
uncond_prompt_mask = train_dataset.uncond_prompt_mask
|
||||
sampler = LengthGroupedSampler(
|
||||
args.train_batch_size,
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
lengths=train_dataset.lengths,
|
||||
group_frame=args.group_frame,
|
||||
group_resolution=args.group_resolution,
|
||||
) if (args.group_frame or args.group_resolution) else DistributedSampler(train_dataset, rank=rank, num_replicas=world_size, shuffle=False)
|
||||
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
collate_fn=latent_collate_function,
|
||||
pin_memory=True,
|
||||
batch_size=args.train_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
drop_last=True,
|
||||
)
|
||||
|
||||
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
|
||||
|
||||
if rank <= 0:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
# Train!
|
||||
total_batch_size = args.train_batch_size * world_size * args.gradient_accumulation_steps / args.sp_size * args.train_sp_batch_size
|
||||
main_print("***** Running training *****")
|
||||
main_print(f" Num examples = {len(train_dataset)}")
|
||||
main_print(f" Dataloader size = {len(train_dataloader)}")
|
||||
main_print(f" Num Epochs = {args.num_train_epochs}")
|
||||
main_print(f" Resume training from step {init_steps}")
|
||||
main_print(f" Instantaneous batch size per device = {args.train_batch_size}")
|
||||
main_print(f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}")
|
||||
main_print(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
|
||||
main_print(f" Total optimization steps = {args.max_train_steps}")
|
||||
main_print(f" Total training parameters per FSDP shard = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e9} B")
|
||||
# print dtype
|
||||
main_print(f" Master weight dtype: {transformer.parameters().__next__().dtype}")
|
||||
|
||||
|
||||
# Potentially load in the weights and states from a previous save
|
||||
if args.resume_from_checkpoint:
|
||||
assert NotImplementedError("resume_from_checkpoint is not supported now.")
|
||||
# TODO
|
||||
|
||||
|
||||
progress_bar = tqdm(
|
||||
range(0, args.max_train_steps),
|
||||
initial=init_steps,
|
||||
desc="Steps",
|
||||
# Only show the progress bar once on each machine.
|
||||
disable= local_rank > 0,
|
||||
)
|
||||
|
||||
loader = sp_parallel_dataloader_wrapper(train_dataloader, device, args.train_batch_size, args.sp_size, args.train_sp_batch_size)
|
||||
|
||||
step_times = deque(maxlen=100)
|
||||
|
||||
#todo future
|
||||
for i in range(init_steps):
|
||||
next(loader)
|
||||
# log_validation(args, transformer, device,
|
||||
# torch.bfloat16, 0, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,ema=False)
|
||||
def get_num_phases(multi_phased_distill_schedule, step):
|
||||
# step-phase,step-phase
|
||||
multi_phases = multi_phased_distill_schedule.split(",")
|
||||
phase = multi_phases[-1].split("-")[-1]
|
||||
for step_phases in multi_phases:
|
||||
phase_step, phase = step_phases.split("-")
|
||||
if step <= int(phase_step):
|
||||
return int(phase)
|
||||
return phase
|
||||
for step in range(init_steps + 1, args.max_train_steps+1):
|
||||
start_time = time.time()
|
||||
assert args.multi_phased_distill_schedule is not None
|
||||
num_phases = get_num_phases(args.multi_phased_distill_schedule, step)
|
||||
|
||||
loss, grad_norm, pred_norm = train_one_step_mochi(transformer,teacher_transformer, ema_transformer, optimizer, lr_scheduler, loader, noise_scheduler,solver, noise_random_generator, args.gradient_accumulation_steps, args.sp_size, args.precondition_outputs, args.max_grad_norm, uncond_prompt_embed, uncond_prompt_mask, args.num_euler_timesteps, num_phases, args.not_apply_cfg_solver,args.distill_cfg, args.ema_decay , args.pred_decay_weight, args.pred_decay_type)
|
||||
|
||||
step_time = time.time() - start_time
|
||||
step_times.append(step_time)
|
||||
avg_step_time = sum(step_times) / len(step_times)
|
||||
|
||||
progress_bar.set_postfix({
|
||||
"loss": f"{loss:.4f}",
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
"grad_norm": grad_norm,
|
||||
"phases": num_phases,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if rank <= 0:
|
||||
wandb.log({
|
||||
"train_loss": loss,
|
||||
"learning_rate": lr_scheduler.get_last_lr()[0],
|
||||
"step_time": step_time,
|
||||
"avg_step_time": avg_step_time,
|
||||
"grad_norm": grad_norm ,
|
||||
"pred_fro_norm": pred_norm["fro"],
|
||||
"pred_largest_singular_value": pred_norm["largest singular value"],
|
||||
"pred_absolute_mean": pred_norm["absolute mean"],
|
||||
"pred_absolute_max": pred_norm["absolute max"],
|
||||
}, step=step)
|
||||
if step % args.checkpointing_steps == 0:
|
||||
if args.use_lora:
|
||||
# Save LoRA weights
|
||||
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, step)
|
||||
else:
|
||||
# Your existing checkpoint saving code
|
||||
if args.use_ema:
|
||||
save_checkpoint(ema_transformer, rank, args.output_dir, step)
|
||||
else:
|
||||
save_checkpoint(transformer, rank, args.output_dir, step)
|
||||
dist.barrier()
|
||||
if args.log_validation and step % args.validation_steps == 0:
|
||||
log_validation(args, transformer, device,
|
||||
torch.bfloat16, step, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,linear_range=args.linear_range, ema=False)
|
||||
if args.use_ema:
|
||||
log_validation(args, ema_transformer, device,
|
||||
torch.bfloat16, step, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,linear_range=args.linear_range, ema=True)
|
||||
|
||||
if args.use_lora:
|
||||
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, args.max_train_steps)
|
||||
else:
|
||||
save_checkpoint(transformer, rank, args.output_dir, args.max_train_steps)
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
destroy_sequence_parallel_group()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--data_json_path", type=str, required=True)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument("--dataloader_num_workers", type=int, default=10, help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.")
|
||||
parser.add_argument("--train_batch_size", type=int, default=16, help="Batch size (per device) for the training dataloader.")
|
||||
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--pretrained_model_name_or_path", type=str)
|
||||
parser.add_argument("--dit_model_name_or_path", type=str)
|
||||
parser.add_argument("--cache_dir", type=str, default='./cache_dir')
|
||||
parser.add_argument('--enable_stable_fp32', action='store_true') # TODO
|
||||
|
||||
# diffusion setting
|
||||
parser.add_argument("--ema_decay", type=float, default=0.95)
|
||||
parser.add_argument("--ema_start_step", type=int, default=0)
|
||||
parser.add_argument('--cfg', type=float, default=0.1)
|
||||
parser.add_argument("--precondition_outputs", action="store_true", help="Whether to precondition the outputs of the model.")
|
||||
|
||||
# validation & logs
|
||||
parser.add_argument("--validation_prompt_dir", type=str)
|
||||
parser.add_argument("--validation_sampling_steps", type=str, default="64")
|
||||
parser.add_argument('--validation_guidance_scale', type=str, default="4.5")
|
||||
|
||||
parser.add_argument('--validation_steps', type=float, default=64)
|
||||
parser.add_argument("--log_validation", action="store_true")
|
||||
parser.add_argument("--tracker_project_name", type=str, default=None)
|
||||
parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.")
|
||||
parser.add_argument("--output_dir", type=str, default=None, help="The output directory where the model predictions and checkpoints will be written.")
|
||||
parser.add_argument("--checkpoints_total_limit", type=int, default=None, help=("Max number of checkpoints to store."))
|
||||
parser.add_argument("--checkpointing_steps", type=int, default=500,
|
||||
help=(
|
||||
"Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
|
||||
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
|
||||
" training using `--resume_from_checkpoint`."
|
||||
),
|
||||
)
|
||||
parser.add_argument("--shift", type=float, default=1.0 )
|
||||
parser.add_argument("--resume_from_checkpoint", type=str, default=None,
|
||||
help=(
|
||||
"Whether training should be resumed from a previous checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
|
||||
),
|
||||
)
|
||||
parser.add_argument("--resume_from_lora_checkpoint", type=str, default=None,
|
||||
help=(
|
||||
"Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
|
||||
),
|
||||
)
|
||||
parser.add_argument("--logging_dir", type=str, default="logs",
|
||||
help=(
|
||||
"[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
|
||||
),
|
||||
)
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
parser.add_argument("--num_train_epochs", type=int, default=100)
|
||||
parser.add_argument("--max_train_steps", type=int, default=None, help="Total number of training steps to perform. If provided, overrides num_train_epochs.")
|
||||
parser.add_argument("--gradient_accumulation_steps", type=int, default=1, help="Number of updates steps to accumulate before performing a backward/update pass.")
|
||||
parser.add_argument("--learning_rate", type=float, default=1e-4, help="Initial learning rate (after the potential warmup period) to use.")
|
||||
parser.add_argument("--scale_lr", action="store_true", default=False, help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.")
|
||||
parser.add_argument("--lr_warmup_steps", type=int, default=10, help="Number of steps for the warmup in the lr scheduler.")
|
||||
parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.")
|
||||
parser.add_argument("--gradient_checkpointing", action="store_true", help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.")
|
||||
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
|
||||
parser.add_argument("--allow_tf32", action="store_true",
|
||||
help=(
|
||||
"Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
|
||||
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
|
||||
),
|
||||
)
|
||||
parser.add_argument("--mixed_precision", type=str, default=None, choices=["no", "fp16", "bf16"],
|
||||
help=(
|
||||
"Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
|
||||
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
|
||||
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
|
||||
),
|
||||
)
|
||||
parser.add_argument("--use_cpu_offload", action="store_true", help="Whether to use CPU offload for param & gradient & optimizer states.")
|
||||
|
||||
parser.add_argument("--sp_size", type=int, default=1, help="For sequence parallel")
|
||||
parser.add_argument("--train_sp_batch_size", type=int, default=1, help="Batch size for sequence parallel training")
|
||||
|
||||
parser.add_argument("--use_lora", action="store_true", default=False, help="Whether to use LoRA for finetuning.")
|
||||
parser.add_argument("--lora_alpha", type=int, default=256, help="Alpha parameter for LoRA.")
|
||||
parser.add_argument("--lora_rank", type=int, default=128, help="LoRA rank parameter. ")
|
||||
parser.add_argument("--fsdp_sharding_startegy", default="full")
|
||||
|
||||
# lr_scheduler
|
||||
parser.add_argument("--lr_scheduler", type=str, default="constant",
|
||||
help=(
|
||||
'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
|
||||
' "constant", "constant_with_warmup"]'
|
||||
),
|
||||
)
|
||||
parser.add_argument("--num_euler_timesteps", type=int, default=100)
|
||||
parser.add_argument("--lr_num_cycles", type=int, default=1, help="Number of cycles in the learning rate scheduler.")
|
||||
parser.add_argument("--lr_power", type=float, default=1.0, help="Power factor of the polynomial scheduler.",)
|
||||
parser.add_argument("--not_apply_cfg_solver", action="store_true", help="Whether to apply the cfg_solver.")
|
||||
parser.add_argument("--distill_cfg", type=float, default=3.0, help="Distillation coefficient.")
|
||||
# ["euler_linear_quadratic", "pcm", "pcm_linear_qudratic"]
|
||||
parser.add_argument("--scheduler_type", type=str, default="pcm", help="The scheduler type to use.")
|
||||
parser.add_argument("--linear_quadratic_threshold", type=float, default=0.025, help="Threshold for linear quadratic scheduler.")
|
||||
parser.add_argument("--linear_range", type=float, default=0.5, help="Range for linear quadratic scheduler.")
|
||||
parser.add_argument("--weight_decay", type=float, default=0.001, help="Weight decay to apply.")
|
||||
parser.add_argument("--use_ema", action="store_true", help="Whether to use EMA.")
|
||||
parser.add_argument("--multi_phased_distill_schedule", type=str, default=None)
|
||||
parser.add_argument("--finetune_weight", type=float, default=0.0)
|
||||
parser.add_argument("--pred_decay_weight", type=float, default=0.0)
|
||||
parser.add_argument("--pred_decay_type", default="l1")
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,105 @@
|
||||
from typing import Any, Dict, Optional, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
|
||||
from diffusers.models.attention import JointTransformerBlock
|
||||
from diffusers.models.attention_processor import Attention, AttentionProcessor
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.models.normalization import AdaLayerNormContinuous
|
||||
from diffusers.utils import (
|
||||
USE_PEFT_BACKEND,
|
||||
is_torch_version,
|
||||
logging,
|
||||
scale_lora_layers,
|
||||
unscale_lora_layers,
|
||||
)
|
||||
from diffusers.models.embeddings import CombinedTimestepTextProjEmbeddings, PatchEmbed
|
||||
from diffusers.models.transformers.transformer_2d import Transformer2DModelOutput
|
||||
from diffusers.models.transformers.transformer_sd3 import SD3Transformer2DModel
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
|
||||
class DiscriminatorHead(nn.Module):
|
||||
def __init__(self, input_channel, output_channel=1):
|
||||
super().__init__()
|
||||
inner_channel = 1024
|
||||
self.conv1 = nn.Sequential(
|
||||
nn.Conv2d(input_channel, inner_channel, 1, 1, 0),
|
||||
nn.GroupNorm(32, inner_channel),
|
||||
nn.LeakyReLU(
|
||||
inplace=True
|
||||
), # use LeakyReLu instead of GELU shown in the paper to save memory
|
||||
)
|
||||
self.conv2 = nn.Sequential(
|
||||
nn.Conv2d(inner_channel, inner_channel, 1, 1, 0),
|
||||
nn.GroupNorm(32, inner_channel),
|
||||
nn.LeakyReLU(
|
||||
inplace=True
|
||||
), # use LeakyReLu instead of GELU shown in the paper to save memory
|
||||
)
|
||||
|
||||
self.conv_out = nn.Conv2d(inner_channel, output_channel, 1, 1, 0)
|
||||
|
||||
def forward(self, x):
|
||||
b, twh, c = x.shape
|
||||
t = twh // (30 * 53)
|
||||
x = x.view(-1, 30 *53, c)
|
||||
x = x.permute(0, 2, 1)
|
||||
x = x.view(b*t, c, 30, 53)
|
||||
x = self.conv1(x)
|
||||
x = self.conv2(x) + x
|
||||
x = self.conv_out(x)
|
||||
return x
|
||||
|
||||
|
||||
class Discriminator(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
stride = 8,
|
||||
num_h_per_head=1,
|
||||
adapter_channel_dims=[3072],
|
||||
):
|
||||
super().__init__()
|
||||
adapter_channel_dims = adapter_channel_dims * (48 // stride)
|
||||
self.stride = stride
|
||||
self.num_h_per_head = num_h_per_head
|
||||
self.head_num = len(adapter_channel_dims)
|
||||
self.heads = nn.ModuleList(
|
||||
[
|
||||
nn.ModuleList(
|
||||
[
|
||||
DiscriminatorHead(adapter_channel)
|
||||
for _ in range(self.num_h_per_head)
|
||||
]
|
||||
)
|
||||
for adapter_channel in adapter_channel_dims
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
|
||||
def forward(self, features):
|
||||
outputs = []
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
return module(*inputs)
|
||||
return custom_forward
|
||||
assert len(features) // self.stride == len(self.heads)
|
||||
for i in range(0, len(features), self.stride):
|
||||
for h in self.heads[i//self.stride]:
|
||||
# out = torch.utils.checkpoint.checkpoint(
|
||||
# create_custom_forward(h),
|
||||
# features[i],
|
||||
# use_reentrant=False
|
||||
# )
|
||||
out=h(features[i])
|
||||
outputs.append(out)
|
||||
return outputs
|
||||
|
||||
|
||||
@@ -0,0 +1,308 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
from fastvideo.model.pipeline_mochi import linear_quadratic_schedule
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
@dataclass
|
||||
class PCMFMSchedulerOutput(BaseOutput):
|
||||
prev_sample: torch.FloatTensor
|
||||
|
||||
def extract_into_tensor(a, t, x_shape):
|
||||
b, *_ = t.shape
|
||||
out = a.gather(-1, t)
|
||||
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
|
||||
|
||||
class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
|
||||
_compatibles = []
|
||||
order = 1
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
num_train_timesteps: int = 1000,
|
||||
shift: float = 1.0,
|
||||
pcm_timesteps: int = 50,
|
||||
linear_quadratic=False,
|
||||
linear_quadratic_threshold=0.025,
|
||||
linear_range=0.5,
|
||||
):
|
||||
|
||||
if linear_quadratic:
|
||||
linear_steps = int(num_train_timesteps * linear_range)
|
||||
sigmas = linear_quadratic_schedule(num_train_timesteps, linear_quadratic_threshold, linear_steps)
|
||||
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
|
||||
else:
|
||||
timesteps = np.linspace(
|
||||
1, num_train_timesteps, num_train_timesteps, dtype=np.float32
|
||||
)[::-1].copy()
|
||||
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
|
||||
sigmas = timesteps / num_train_timesteps
|
||||
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
|
||||
self.euler_timesteps = (
|
||||
np.arange(1, pcm_timesteps + 1) * (num_train_timesteps // pcm_timesteps)
|
||||
).round().astype(np.int64) - 1
|
||||
self.sigmas = sigmas.numpy()[::-1][self.euler_timesteps]
|
||||
self.sigmas = torch.from_numpy((self.sigmas[::-1].copy()))
|
||||
self.timesteps = self.sigmas * num_train_timesteps
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
self.sigmas = self.sigmas.to("cpu") # to avoid too much CPU/GPU communication
|
||||
self.sigma_min = self.sigmas[-1].item()
|
||||
self.sigma_max = self.sigmas[0].item()
|
||||
|
||||
@property
|
||||
def step_index(self):
|
||||
"""
|
||||
The index counter for current timestep. It will increase 1 after each scheduler step.
|
||||
"""
|
||||
return self._step_index
|
||||
|
||||
@property
|
||||
def begin_index(self):
|
||||
"""
|
||||
The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
|
||||
"""
|
||||
return self._begin_index
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
|
||||
def set_begin_index(self, begin_index: int = 0):
|
||||
"""
|
||||
Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
|
||||
|
||||
Args:
|
||||
begin_index (`int`):
|
||||
The begin index for the scheduler.
|
||||
"""
|
||||
self._begin_index = begin_index
|
||||
|
||||
def scale_noise(
|
||||
self,
|
||||
sample: torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
noise: Optional[torch.FloatTensor] = None,
|
||||
) -> torch.FloatTensor:
|
||||
"""
|
||||
Forward process in flow-matching
|
||||
|
||||
Args:
|
||||
sample (`torch.FloatTensor`):
|
||||
The input sample.
|
||||
timestep (`int`, *optional*):
|
||||
The current timestep in the diffusion chain.
|
||||
|
||||
Returns:
|
||||
`torch.FloatTensor`:
|
||||
A scaled input sample.
|
||||
"""
|
||||
if self.step_index is None:
|
||||
self._init_step_index(timestep)
|
||||
|
||||
sigma = self.sigmas[self.step_index]
|
||||
sample = sigma * noise + (1.0 - sigma) * sample
|
||||
|
||||
return sample
|
||||
|
||||
def _sigma_to_t(self, sigma):
|
||||
return sigma * self.config.num_train_timesteps
|
||||
|
||||
def set_timesteps(
|
||||
self, num_inference_steps: int, device: Union[str, torch.device] = None
|
||||
):
|
||||
"""
|
||||
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||
|
||||
Args:
|
||||
num_inference_steps (`int`):
|
||||
The number of diffusion steps used when generating samples with a pre-trained model.
|
||||
device (`str` or `torch.device`, *optional*):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
"""
|
||||
self.num_inference_steps = num_inference_steps
|
||||
inference_indices = np.linspace(
|
||||
0, self.config.pcm_timesteps, num=num_inference_steps, endpoint=False
|
||||
)
|
||||
inference_indices = np.floor(inference_indices).astype(np.int64)
|
||||
inference_indices = torch.from_numpy(inference_indices).long()
|
||||
|
||||
self.sigmas_ = self.sigmas[inference_indices]
|
||||
timesteps = self.sigmas_ * self.config.num_train_timesteps
|
||||
self.timesteps = timesteps.to(device=device)
|
||||
self.sigmas_ = torch.cat(
|
||||
[self.sigmas_, torch.zeros(1, device=self.sigmas_.device)]
|
||||
)
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
|
||||
def index_for_timestep(self, timestep, schedule_timesteps=None):
|
||||
if schedule_timesteps is None:
|
||||
schedule_timesteps = self.timesteps
|
||||
|
||||
indices = (schedule_timesteps == timestep).nonzero()
|
||||
|
||||
# The sigma index that is taken for the **very** first `step`
|
||||
# is always the second index (or the last index if there is only 1)
|
||||
# This way we can ensure we don't accidentally skip a sigma in
|
||||
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
|
||||
pos = 1 if len(indices) > 1 else 0
|
||||
|
||||
return indices[pos].item()
|
||||
|
||||
def _init_step_index(self, timestep):
|
||||
if self.begin_index is None:
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
timestep = timestep.to(self.timesteps.device)
|
||||
self._step_index = self.index_for_timestep(timestep)
|
||||
else:
|
||||
self._step_index = self._begin_index
|
||||
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
sample: torch.FloatTensor,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
return_dict: bool = True,
|
||||
) -> Union[PCMFMSchedulerOutput, Tuple]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
|
||||
process from the learned model outputs (most often the predicted noise).
|
||||
|
||||
Args:
|
||||
model_output (`torch.FloatTensor`):
|
||||
The direct output from learned diffusion model.
|
||||
timestep (`float`):
|
||||
The current discrete timestep in the diffusion chain.
|
||||
sample (`torch.FloatTensor`):
|
||||
A current instance of a sample created by the diffusion process.
|
||||
s_churn (`float`):
|
||||
s_tmin (`float`):
|
||||
s_tmax (`float`):
|
||||
s_noise (`float`, defaults to 1.0):
|
||||
Scaling factor for noise added to the sample.
|
||||
generator (`torch.Generator`, *optional*):
|
||||
A random number generator.
|
||||
return_dict (`bool`):
|
||||
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
|
||||
tuple.
|
||||
|
||||
Returns:
|
||||
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
|
||||
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
|
||||
returned, otherwise a tuple is returned where the first element is the sample tensor.
|
||||
"""
|
||||
|
||||
if (
|
||||
isinstance(timestep, int)
|
||||
or isinstance(timestep, torch.IntTensor)
|
||||
or isinstance(timestep, torch.LongTensor)
|
||||
):
|
||||
raise ValueError(
|
||||
(
|
||||
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
|
||||
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||
" one of the `scheduler.timesteps` as a timestep."
|
||||
),
|
||||
)
|
||||
|
||||
if self.step_index is None:
|
||||
self._init_step_index(timestep)
|
||||
|
||||
sample = sample.to(torch.float32)
|
||||
|
||||
sigma = self.sigmas_[self.step_index]
|
||||
|
||||
denoised = sample - model_output * sigma
|
||||
derivative = (sample - denoised) / sigma
|
||||
|
||||
dt = self.sigmas_[self.step_index + 1] - sigma
|
||||
prev_sample = sample + derivative * dt
|
||||
prev_sample = prev_sample.to(model_output.dtype)
|
||||
self._step_index += 1
|
||||
|
||||
if not return_dict:
|
||||
return (prev_sample,)
|
||||
|
||||
return PCMFMSchedulerOutput(prev_sample=prev_sample)
|
||||
|
||||
def __len__(self):
|
||||
return self.config.num_train_timesteps
|
||||
|
||||
class EulerSolver:
|
||||
def __init__(self, sigmas, timesteps=1000, euler_timesteps=50):
|
||||
self.step_ratio = timesteps // euler_timesteps
|
||||
self.euler_timesteps = (
|
||||
np.arange(1, euler_timesteps + 1) * self.step_ratio
|
||||
).round().astype(np.int64) - 1
|
||||
self.euler_timesteps_prev = np.asarray([0] + self.euler_timesteps[:-1].tolist())
|
||||
self.sigmas = sigmas[self.euler_timesteps]
|
||||
self.sigmas_prev = np.asarray(
|
||||
[sigmas[0]] + sigmas[self.euler_timesteps[:-1]].tolist()
|
||||
) # either use sigma0 or 0
|
||||
|
||||
self.euler_timesteps = torch.from_numpy(self.euler_timesteps).long()
|
||||
self.euler_timesteps_prev = torch.from_numpy(self.euler_timesteps_prev).long()
|
||||
self.sigmas = torch.from_numpy(self.sigmas)
|
||||
self.sigmas_prev = torch.from_numpy(self.sigmas_prev)
|
||||
|
||||
def to(self, device):
|
||||
self.euler_timesteps = self.euler_timesteps.to(device)
|
||||
self.euler_timesteps_prev = self.euler_timesteps_prev.to(device)
|
||||
|
||||
self.sigmas = self.sigmas.to(device)
|
||||
self.sigmas_prev = self.sigmas_prev.to(device)
|
||||
return self
|
||||
|
||||
def euler_step(self, sample, model_pred, timestep_index):
|
||||
sigma = extract_into_tensor(self.sigmas, timestep_index, model_pred.shape)
|
||||
sigma_prev = extract_into_tensor(
|
||||
self.sigmas_prev, timestep_index, model_pred.shape
|
||||
)
|
||||
x_prev = sample + (sigma_prev - sigma) * model_pred
|
||||
return x_prev
|
||||
|
||||
def euler_style_multiphase_pred(
|
||||
self,
|
||||
sample,
|
||||
model_pred,
|
||||
timestep_index,
|
||||
multiphase,
|
||||
is_target=False,
|
||||
):
|
||||
|
||||
inference_indices = np.linspace(
|
||||
0, len(self.euler_timesteps), num=multiphase, endpoint=False
|
||||
)
|
||||
inference_indices = np.floor(inference_indices).astype(np.int64)
|
||||
inference_indices = (
|
||||
torch.from_numpy(inference_indices).long().to(self.euler_timesteps.device)
|
||||
)
|
||||
expanded_timestep_index = timestep_index.unsqueeze(1).expand(
|
||||
-1, inference_indices.size(0)
|
||||
)
|
||||
valid_indices_mask = expanded_timestep_index >= inference_indices
|
||||
last_valid_index = valid_indices_mask.flip(dims=[1]).long().argmax(dim=1)
|
||||
last_valid_index = inference_indices.size(0) - 1 - last_valid_index
|
||||
timestep_index_end = inference_indices[last_valid_index]
|
||||
|
||||
if is_target:
|
||||
sigma = extract_into_tensor(self.sigmas_prev, timestep_index, sample.shape)
|
||||
else:
|
||||
sigma = extract_into_tensor(self.sigmas, timestep_index, sample.shape)
|
||||
sigma_prev = extract_into_tensor(
|
||||
self.sigmas_prev, timestep_index_end, sample.shape
|
||||
)
|
||||
x_prev = sample + (sigma_prev - sigma) * model_pred
|
||||
|
||||
return x_prev, timestep_index_end
|
||||
|
||||
@@ -0,0 +1,687 @@
|
||||
import argparse
|
||||
from email.policy import strict
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, \
|
||||
destroy_sequence_parallel_group, get_sequence_parallel_state, nccl_info
|
||||
from fastvideo.utils.communications import sp_parallel_dataloader_wrapper, broadcast
|
||||
from fastvideo.model.mochi_latents_utils import normalize_mochi_dit_input
|
||||
from fastvideo.utils.validation import log_validation
|
||||
import time
|
||||
from torch.utils.data import DataLoader
|
||||
import torch
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
StateDictType,
|
||||
FullStateDictConfig,
|
||||
)
|
||||
|
||||
from fastvideo.model.pipeline_mochi import linear_quadratic_schedule
|
||||
import json
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from fastvideo.utils.dataset_utils import LengthGroupedSampler
|
||||
import wandb
|
||||
from accelerate.utils import set_seed
|
||||
from tqdm.auto import tqdm
|
||||
from fastvideo.fsdp_util import get_dit_fsdp_kwargs, apply_fsdp_checkpointing, get_discriminator_fsdp_kwargs
|
||||
import diffusers
|
||||
from diffusers import (
|
||||
FlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
from fastvideo.distill.discriminator import Discriminator
|
||||
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
|
||||
from copy import deepcopy
|
||||
from diffusers.optimization import get_scheduler
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
from diffusers.utils import check_min_version
|
||||
from fastvideo.dataset.latent_datasets import LatentDataset, latent_collate_function
|
||||
import torch.distributed as dist
|
||||
from peft import LoraConfig, inject_adapter_in_model
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
)
|
||||
from fastvideo.utils.checkpoint import save_checkpoint, save_lora_checkpoint, resume_lora_training, resume_training, save_checkpoint_generator_discriminator, resume_training_generator_discriminator
|
||||
from fastvideo.utils.logging import main_print
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
check_min_version("0.31.0")
|
||||
import time
|
||||
from collections import deque
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def gan_d_loss(
|
||||
discriminator,
|
||||
teacher_transformer,
|
||||
sample_fake,
|
||||
sample_real,
|
||||
timestep,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
weight,
|
||||
):
|
||||
loss = 0.0
|
||||
# collate sample_fake and sample_real
|
||||
with torch.no_grad():
|
||||
fake_features = teacher_transformer(
|
||||
sample_fake,
|
||||
encoder_hidden_states,
|
||||
timestep,
|
||||
encoder_attention_mask,
|
||||
output_attn=True,
|
||||
return_dict= False
|
||||
)[1]
|
||||
real_features = teacher_transformer(
|
||||
sample_real,
|
||||
encoder_hidden_states,
|
||||
timestep,
|
||||
encoder_attention_mask,
|
||||
output_attn=True,
|
||||
return_dict= False
|
||||
)[1]
|
||||
|
||||
fake_outputs = discriminator(
|
||||
fake_features
|
||||
)
|
||||
real_outputs = discriminator(
|
||||
real_features
|
||||
)
|
||||
for fake_output, real_output in zip(fake_outputs, real_outputs):
|
||||
loss += (
|
||||
torch.mean(weight * torch.relu(fake_output.float() + 1))
|
||||
+ torch.mean(weight * torch.relu(1 - real_output.float()))
|
||||
) / (discriminator.head_num * discriminator.num_h_per_head)
|
||||
return loss
|
||||
|
||||
def gan_g_loss(
|
||||
discriminator,
|
||||
teacher_transformer,
|
||||
sample_fake,
|
||||
timestep,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
weight,
|
||||
):
|
||||
loss = 0.0
|
||||
features = teacher_transformer(
|
||||
sample_fake,
|
||||
encoder_hidden_states,
|
||||
timestep,
|
||||
encoder_attention_mask,
|
||||
output_attn=True,
|
||||
return_dict= False
|
||||
)[1]
|
||||
fake_outputs = discriminator(
|
||||
features,
|
||||
)
|
||||
for fake_output in fake_outputs:
|
||||
loss += torch.mean(weight * torch.relu(1 - fake_output.float())) / (
|
||||
discriminator.head_num * discriminator.num_h_per_head
|
||||
)
|
||||
return loss
|
||||
|
||||
def train_one_step_mochi(transformer, teacher_transformer , optimizer, discriminator, discriminator_optimizer,global_step, lr_scheduler,loader, noise_scheduler, solver,noise_random_generator, sp_size, precondition_outputs, max_grad_norm, uncond_prompt_embed, uncond_prompt_mask, num_euler_timesteps, multiphase, not_apply_cfg_solver, distill_cfg, adv_weight):
|
||||
|
||||
|
||||
optimizer.zero_grad()
|
||||
discriminator_optimizer.zero_grad()
|
||||
|
||||
latents, encoder_hidden_states, latents_attention_mask, encoder_attention_mask = next(loader)
|
||||
model_input = normalize_mochi_dit_input(latents)
|
||||
noise = torch.randn_like(model_input)
|
||||
bsz = model_input.shape[0]
|
||||
index = torch.randint(
|
||||
0, num_euler_timesteps, (bsz,), device=model_input.device
|
||||
).long()
|
||||
if sp_size > 1:
|
||||
broadcast(index)
|
||||
# Add noise according to flow matching.
|
||||
# sigmas = get_sigmas(start_timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
|
||||
sigmas = extract_into_tensor(solver.sigmas, index, model_input.shape)
|
||||
sigmas_prev = extract_into_tensor(
|
||||
solver.sigmas_prev, index, model_input.shape
|
||||
)
|
||||
|
||||
timesteps = (
|
||||
sigmas * noise_scheduler.config.num_train_timesteps
|
||||
).view(-1)
|
||||
# if squeeze to [], unsqueeze to [1]
|
||||
|
||||
timesteps_prev = (
|
||||
sigmas_prev * noise_scheduler.config.num_train_timesteps
|
||||
).view(-1)
|
||||
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
|
||||
|
||||
# Predict the noise residual
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
|
||||
model_pred = transformer(
|
||||
noisy_model_input,
|
||||
encoder_hidden_states,
|
||||
timesteps,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict= False
|
||||
)[0]
|
||||
|
||||
# if accelerator.is_main_process:
|
||||
model_pred, end_index = solver.euler_style_multiphase_pred(
|
||||
noisy_model_input, model_pred, index, multiphase
|
||||
)
|
||||
|
||||
weighting = 1.0
|
||||
# # simplified flow matching aka 0-rectified flow matching loss
|
||||
# # target = model_input - noise
|
||||
# target = model_input
|
||||
adv_index = torch.empty_like(end_index)
|
||||
for i in range(end_index.size(0)):
|
||||
adv_index[i] = torch.randint(
|
||||
end_index[i].item(),
|
||||
end_index[i].item()
|
||||
+ num_euler_timesteps // multiphase,
|
||||
(1,),
|
||||
dtype=end_index.dtype,
|
||||
device=end_index.device,
|
||||
)
|
||||
|
||||
sigmas_end = extract_into_tensor(
|
||||
solver.sigmas_prev, end_index, model_input.shape
|
||||
)
|
||||
sigmas_adv = extract_into_tensor(
|
||||
solver.sigmas_prev, adv_index, model_input.shape
|
||||
)
|
||||
timesteps_end = (
|
||||
sigmas_end * noise_scheduler.config.num_train_timesteps
|
||||
).view(-1)
|
||||
timesteps_adv = (
|
||||
sigmas_adv * noise_scheduler.config.num_train_timesteps
|
||||
).view(-1)
|
||||
|
||||
|
||||
with torch.no_grad():
|
||||
w = distill_cfg
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
cond_teacher_output = teacher_transformer(
|
||||
noisy_model_input,
|
||||
encoder_hidden_states,
|
||||
timesteps,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict= False
|
||||
)[0].float()
|
||||
if not_apply_cfg_solver:
|
||||
uncond_teacher_output = cond_teacher_output
|
||||
else:
|
||||
# Get teacher model prediction on noisy_latents and unconditional embedding
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
uncond_teacher_output = teacher_transformer(
|
||||
noisy_model_input,
|
||||
uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
|
||||
timesteps,
|
||||
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
|
||||
return_dict= False
|
||||
)[0].float()
|
||||
teacher_output = cond_teacher_output + w * (
|
||||
cond_teacher_output - uncond_teacher_output
|
||||
)
|
||||
x_prev = solver.euler_step(
|
||||
noisy_model_input, teacher_output, index
|
||||
)
|
||||
|
||||
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
|
||||
with torch.no_grad():
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
target_pred = transformer(
|
||||
x_prev.float(),
|
||||
encoder_hidden_states,
|
||||
timesteps_prev,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict= False
|
||||
)[0]
|
||||
|
||||
target, end_index = solver.euler_style_multiphase_pred(
|
||||
x_prev, target_pred, index, multiphase, True
|
||||
)
|
||||
|
||||
real_adv = (
|
||||
(1 - sigmas_adv) * target
|
||||
+ (sigmas_adv - sigmas_end) * torch.randn_like(target)
|
||||
) / (1 - sigmas_end)
|
||||
fake_adv = (
|
||||
(1 - sigmas_adv) * model_pred
|
||||
+ (sigmas_adv - sigmas_end) * torch.randn_like(model_pred)
|
||||
) / (1 - sigmas_end)
|
||||
|
||||
|
||||
|
||||
|
||||
huber_c = 0.001
|
||||
g_loss = torch.mean(
|
||||
torch.sqrt(
|
||||
(model_pred.float() - target.float()) ** 2
|
||||
+ huber_c**2
|
||||
)
|
||||
- huber_c
|
||||
)
|
||||
discriminator.requires_grad_(False)
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
g_gan_loss = adv_weight * gan_g_loss(
|
||||
discriminator,
|
||||
teacher_transformer,
|
||||
fake_adv.float(),
|
||||
timesteps_adv,
|
||||
encoder_hidden_states.float(),
|
||||
encoder_attention_mask,
|
||||
1.0,
|
||||
)
|
||||
g_loss += g_gan_loss
|
||||
g_loss.backward()
|
||||
|
||||
g_loss = g_loss.detach().clone()
|
||||
dist.all_reduce(g_loss, op=dist.ReduceOp.AVG)
|
||||
|
||||
|
||||
g_grad_norm = transformer.clip_grad_norm_(max_grad_norm).item()
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
optimizer.zero_grad()
|
||||
|
||||
discriminator_optimizer.zero_grad()
|
||||
discriminator.requires_grad_(True)
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
d_loss = gan_d_loss(
|
||||
discriminator,
|
||||
teacher_transformer,
|
||||
fake_adv.detach(),
|
||||
real_adv.detach(),
|
||||
timesteps_adv,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
1.0,
|
||||
)
|
||||
|
||||
d_loss.backward()
|
||||
d_grad_norm = discriminator.clip_grad_norm_(max_grad_norm).item()
|
||||
discriminator_optimizer.step()
|
||||
discriminator_optimizer.zero_grad()
|
||||
|
||||
return g_loss, g_grad_norm, d_loss, d_grad_norm
|
||||
|
||||
def get_lora_model(transformer, lora_config):
|
||||
transformer.requires_grad_(False)
|
||||
transformer = inject_adapter_in_model(lora_config, transformer)
|
||||
return transformer
|
||||
|
||||
|
||||
def main(args):
|
||||
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
|
||||
local_rank = int(os.environ['LOCAL_RANK'])
|
||||
rank = int(os.environ['RANK'])
|
||||
world_size = int(os.environ['WORLD_SIZE'])
|
||||
dist.init_process_group("nccl")
|
||||
torch.cuda.set_device(local_rank)
|
||||
device = torch.cuda.current_device()
|
||||
initialize_sequence_parallel_state(args.sp_size)
|
||||
|
||||
# If passed along, set the training seed now. On GPU...
|
||||
if args.seed is not None:
|
||||
# TODO: t within the same seq parallel group should be the same. Noise should be different.
|
||||
set_seed(args.seed + rank)
|
||||
# We use different seeds for the noise generation in each process to ensure that the noise is different in a batch.
|
||||
noise_random_generator = None
|
||||
|
||||
# Handle the repository creation
|
||||
if rank <=0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weigths to half-precision
|
||||
# as these weights are only used for inference, keeping weights in full precision is not required.
|
||||
|
||||
# Create model:
|
||||
|
||||
main_print(f"--> loading model from {args.pretrained_model_name_or_path}")
|
||||
# keep the master weight to float32
|
||||
if args.dit_model_name_or_path:
|
||||
transformer = transformer = MochiTransformer3DModel.from_pretrained(
|
||||
args.dit_model_name_or_path,
|
||||
torch_dtype = torch.float32,
|
||||
#torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
else:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(
|
||||
args.pretrained_model_name_or_path,
|
||||
subfolder="transformer",
|
||||
torch_dtype = torch.float32,
|
||||
#torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
teacher_transformer = deepcopy(transformer)
|
||||
discriminator = Discriminator(args.discriminator_head_stride)
|
||||
|
||||
|
||||
if args.use_lora:
|
||||
lora_config = LoraConfig(
|
||||
r=args.lora_rank,
|
||||
lora_alpha=args.lora_alpha,
|
||||
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
|
||||
init_lora_weights=True,
|
||||
)
|
||||
transformer = get_lora_model(transformer, lora_config)
|
||||
|
||||
main_print(f" Total transformer parameters = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e6} M")
|
||||
# discriminator
|
||||
main_print(f" Total discriminator parameters = {sum(p.numel() for p in discriminator.parameters() if p.requires_grad) / 1e6} M")
|
||||
main_print(f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}")
|
||||
fsdp_kwargs = get_dit_fsdp_kwargs(args.fsdp_sharding_startegy, args.use_lora, args.use_cpu_offload)
|
||||
discriminator_fsdp_kwargs = get_discriminator_fsdp_kwargs()
|
||||
if args.use_lora:
|
||||
transformer.config.lora_rank = args.lora_rank
|
||||
transformer.config.lora_alpha = args.lora_alpha
|
||||
transformer.config.lora_target_modules = ["to_k", "to_q", "to_v", "to_out.0"]
|
||||
transformer._no_split_modules = ["MochiTransformerBlock"]
|
||||
fsdp_kwargs['auto_wrap_policy'] = fsdp_kwargs['auto_wrap_policy'](transformer)
|
||||
|
||||
|
||||
transformer = FSDP(
|
||||
transformer,
|
||||
**fsdp_kwargs,
|
||||
)
|
||||
teacher_transformer = FSDP(
|
||||
teacher_transformer,
|
||||
**fsdp_kwargs,
|
||||
)
|
||||
discriminator = FSDP(
|
||||
discriminator,
|
||||
**discriminator_fsdp_kwargs,
|
||||
)
|
||||
main_print(f"--> model loaded")
|
||||
|
||||
if args.gradient_checkpointing:
|
||||
apply_fsdp_checkpointing(transformer, args.selective_checkpointing)
|
||||
apply_fsdp_checkpointing(teacher_transformer, args.selective_checkpointing)
|
||||
# Set model as trainable.
|
||||
transformer.train()
|
||||
teacher_transformer.requires_grad_(False)
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
|
||||
if args.scheduler_type == "pcm_linear_quadratic":
|
||||
sigmas = linear_quadratic_schedule(noise_scheduler.config.num_train_timesteps, args.linear_quadratic_threshold)
|
||||
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
|
||||
else:
|
||||
sigmas = noise_scheduler.sigmas
|
||||
solver = EulerSolver(
|
||||
sigmas.numpy()[::-1],
|
||||
noise_scheduler.config.num_train_timesteps,
|
||||
euler_timesteps=args.num_euler_timesteps,
|
||||
)
|
||||
solver.to(device)
|
||||
params_to_optimize = transformer.parameters()
|
||||
params_to_optimize = list(filter(lambda p: p.requires_grad, params_to_optimize))
|
||||
|
||||
optimizer = torch.optim.AdamW(
|
||||
params_to_optimize,
|
||||
lr=args.learning_rate,
|
||||
betas=(0.9,0.999),
|
||||
weight_decay=1e-3,
|
||||
eps=1e-8,
|
||||
)
|
||||
|
||||
discriminator_optimizer = torch.optim.AdamW(
|
||||
discriminator.parameters(),
|
||||
lr=args.discriminator_learning_rate,
|
||||
betas=(0, 0.999),
|
||||
weight_decay=1e-3,
|
||||
eps=1e-8,
|
||||
)
|
||||
|
||||
|
||||
init_steps = 0
|
||||
if args.resume_from_lora_checkpoint:
|
||||
transformer, optimizer, init_steps = resume_lora_training(
|
||||
transformer, args.resume_from_lora_checkpoint, optimizer
|
||||
)
|
||||
elif args.resume_from_checkpoint:
|
||||
transformer, optimizer,discriminator, discriminator_optimizer, init_steps = resume_training_generator_discriminator(
|
||||
transformer, optimizer,discriminator, discriminator_optimizer, args.resume_from_checkpoint, rank
|
||||
)
|
||||
|
||||
main_print(f"optimizer: {optimizer}")
|
||||
|
||||
lr_scheduler = get_scheduler(
|
||||
args.lr_scheduler,
|
||||
optimizer=optimizer,
|
||||
num_warmup_steps=args.lr_warmup_steps * world_size,
|
||||
num_training_steps=args.max_train_steps * world_size,
|
||||
num_cycles=args.lr_num_cycles,
|
||||
power=args.lr_power,
|
||||
last_epoch=init_steps - 1,
|
||||
)
|
||||
|
||||
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t, args.cfg)
|
||||
uncond_prompt_embed = train_dataset.uncond_prompt_embed
|
||||
uncond_prompt_mask = train_dataset.uncond_prompt_mask
|
||||
sampler = LengthGroupedSampler(
|
||||
args.train_batch_size,
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
lengths=train_dataset.lengths,
|
||||
group_frame=args.group_frame,
|
||||
group_resolution=args.group_resolution,
|
||||
) if (args.group_frame or args.group_resolution) else DistributedSampler(train_dataset, rank=rank, num_replicas=world_size, shuffle=False)
|
||||
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
collate_fn=latent_collate_function,
|
||||
pin_memory=True,
|
||||
batch_size=args.train_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
drop_last=True,
|
||||
)
|
||||
assert args.gradient_accumulation_steps == 1
|
||||
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
|
||||
|
||||
if rank <= 0:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
# Train!
|
||||
total_batch_size = args.train_batch_size * world_size * args.gradient_accumulation_steps / args.sp_size * args.train_sp_batch_size
|
||||
main_print("***** Running training *****")
|
||||
main_print(f" Num examples = {len(train_dataset)}")
|
||||
main_print(f" Dataloader size = {len(train_dataloader)}")
|
||||
main_print(f" Num Epochs = {args.num_train_epochs}")
|
||||
main_print(f" Resume training from step {init_steps}")
|
||||
main_print(f" Instantaneous batch size per device = {args.train_batch_size}")
|
||||
main_print(f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}")
|
||||
main_print(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
|
||||
main_print(f" Total optimization steps = {args.max_train_steps}")
|
||||
main_print(f" Total training parameters per FSDP shard = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e9} B")
|
||||
# print dtype
|
||||
main_print(f" Master weight dtype: {transformer.parameters().__next__().dtype}")
|
||||
|
||||
|
||||
|
||||
|
||||
progress_bar = tqdm(
|
||||
range(0, args.max_train_steps),
|
||||
initial=init_steps,
|
||||
desc="Steps",
|
||||
# Only show the progress bar once on each machine.
|
||||
disable= local_rank > 0,
|
||||
)
|
||||
loader = sp_parallel_dataloader_wrapper(train_dataloader, device, args.train_batch_size, args.sp_size, args.train_sp_batch_size)
|
||||
|
||||
step_times = deque(maxlen=100)
|
||||
# log_validation(args, transformer, device,
|
||||
# torch.bfloat16, init_steps, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold, ema=False)
|
||||
|
||||
for i in range(init_steps):
|
||||
_ = next(loader)
|
||||
for step in range(init_steps + 1, args.max_train_steps+1):
|
||||
start_time = time.time()
|
||||
generator_loss, generator_grad_norm, discriminator_loss, discriminator_grad_norm= train_one_step_mochi(transformer,teacher_transformer, optimizer, discriminator, discriminator_optimizer, step,lr_scheduler, loader, noise_scheduler,solver, noise_random_generator , args.sp_size, args.precondition_outputs, args.max_grad_norm, uncond_prompt_embed, uncond_prompt_mask, args.num_euler_timesteps, args.validation_sampling_steps, args.not_apply_cfg_solver,args.distill_cfg, args.adv_weight)
|
||||
|
||||
step_time = time.time() - start_time
|
||||
step_times.append(step_time)
|
||||
avg_step_time = sum(step_times) / len(step_times)
|
||||
|
||||
progress_bar.set_postfix({
|
||||
"g_loss": f"{generator_loss:.4f}",
|
||||
"d_loss": f"{discriminator_loss:.4f}",
|
||||
"g_grad_norm": generator_grad_norm,
|
||||
"d_grad_norm": discriminator_grad_norm,
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if rank <= 0:
|
||||
wandb.log({
|
||||
"generator_loss": generator_loss,
|
||||
"discriminator_loss": discriminator_loss,
|
||||
"generator_grad_norm": generator_grad_norm,
|
||||
"discriminator_grad_norm": discriminator_grad_norm,
|
||||
"learning_rate": lr_scheduler.get_last_lr()[0],
|
||||
"step_time": step_time,
|
||||
"avg_step_time": avg_step_time,
|
||||
}, step=step)
|
||||
if step % args.checkpointing_steps == 0:
|
||||
main_print(f"--> saving checkpoint at step {step}")
|
||||
if args.use_lora:
|
||||
# Save LoRA weights
|
||||
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, step)
|
||||
else:
|
||||
# Your existing checkpoint saving code
|
||||
save_checkpoint_generator_discriminator(transformer, optimizer, discriminator, discriminator_optimizer, rank, args.output_dir, step)
|
||||
main_print(f"--> checkpoint saved at step {step}")
|
||||
dist.barrier()
|
||||
if args.log_validation and step % args.validation_steps == 0:
|
||||
log_validation(args, transformer, device,
|
||||
torch.bfloat16, step, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps,linear_quadratic_threshold=args.linear_quadratic_threshold, ema=False)
|
||||
|
||||
if args.use_lora:
|
||||
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, args.max_train_steps)
|
||||
else:
|
||||
save_checkpoint(transformer, optimizer, rank, args.output_dir, args.max_train_steps)
|
||||
save_checkpoint(discriminator, discriminator_optimizer, rank, args.output_dir, step, discriminator=True)
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
destroy_sequence_parallel_group()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--data_json_path", type=str, required=True)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument("--dataloader_num_workers", type=int, default=10, help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.")
|
||||
parser.add_argument("--train_batch_size", type=int, default=16, help="Batch size (per device) for the training dataloader.")
|
||||
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--pretrained_model_name_or_path", type=str)
|
||||
parser.add_argument("--dit_model_name_or_path", type=str)
|
||||
parser.add_argument("--cache_dir", type=str, default='./cache_dir')
|
||||
|
||||
# diffusion setting
|
||||
parser.add_argument("--ema_decay", type=float, default=0.999)
|
||||
parser.add_argument("--ema_start_step", type=int, default=0)
|
||||
parser.add_argument('--cfg', type=float, default=0.1)
|
||||
parser.add_argument("--precondition_outputs", action="store_true", help="Whether to precondition the outputs of the model.")
|
||||
|
||||
# validation & logs
|
||||
parser.add_argument("--validation_prompt_dir", type=str)
|
||||
parser.add_argument("--validation_sampling_steps", type=int, default=64)
|
||||
parser.add_argument('--validation_guidance_scale', type=float, default=4.5)
|
||||
parser.add_argument('--validation_steps', type=float, default=64)
|
||||
parser.add_argument("--log_validation", action="store_true")
|
||||
parser.add_argument("--tracker_project_name", type=str, default=None)
|
||||
parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.")
|
||||
parser.add_argument("--output_dir", type=str, default=None, help="The output directory where the model predictions and checkpoints will be written.")
|
||||
parser.add_argument("--checkpoints_total_limit", type=int, default=None, help=("Max number of checkpoints to store."))
|
||||
parser.add_argument("--checkpointing_steps", type=int, default=500,
|
||||
help=(
|
||||
"Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
|
||||
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
|
||||
" training using `--resume_from_checkpoint`."
|
||||
),
|
||||
)
|
||||
parser.add_argument("--shift", type=float, default=1.0 )
|
||||
parser.add_argument("--resume_from_checkpoint", type=str, default=None,
|
||||
help=(
|
||||
"Whether training should be resumed from a previous checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
|
||||
),
|
||||
)
|
||||
parser.add_argument("--resume_from_lora_checkpoint", type=str, default=None,
|
||||
help=(
|
||||
"Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
|
||||
),
|
||||
)
|
||||
parser.add_argument("--logging_dir", type=str, default="logs",
|
||||
help=(
|
||||
"[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
|
||||
),
|
||||
)
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
parser.add_argument("--num_train_epochs", type=int, default=100)
|
||||
parser.add_argument("--max_train_steps", type=int, default=None, help="Total number of training steps to perform. If provided, overrides num_train_epochs.")
|
||||
parser.add_argument("--learning_rate", type=float, default=1e-4, help="Initial learning rate (after the potential warmup period) to use.")
|
||||
parser.add_argument("--discriminator_learning_rate", type=float, default=1e-5, help="Initial learning rate (after the potential warmup period) to use.")
|
||||
parser.add_argument("--scale_lr", action="store_true", default=False, help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.")
|
||||
parser.add_argument("--lr_warmup_steps", type=int, default=10, help="Number of steps for the warmup in the lr scheduler.")
|
||||
parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.")
|
||||
parser.add_argument("--gradient_checkpointing", action="store_true", help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.")
|
||||
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
|
||||
parser.add_argument("--allow_tf32", action="store_true",
|
||||
help=(
|
||||
"Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
|
||||
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
|
||||
),
|
||||
)
|
||||
parser.add_argument("--mixed_precision", type=str, default=None, choices=["no", "fp16", "bf16"],
|
||||
help=(
|
||||
"Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
|
||||
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
|
||||
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
|
||||
),
|
||||
)
|
||||
parser.add_argument("--use_cpu_offload", action="store_true", help="Whether to use CPU offload for param & gradient & optimizer states.")
|
||||
|
||||
parser.add_argument("--sp_size", type=int, default=1, help="For sequence parallel")
|
||||
parser.add_argument("--train_sp_batch_size", type=int, default=1, help="Batch size for sequence parallel training")
|
||||
|
||||
parser.add_argument("--use_lora", action="store_true", default=False, help="Whether to use LoRA for finetuning.")
|
||||
parser.add_argument("--lora_alpha", type=int, default=256, help="Alpha parameter for LoRA.")
|
||||
parser.add_argument("--lora_rank", type=int, default=128, help="LoRA rank parameter. ")
|
||||
parser.add_argument("--fsdp_sharding_startegy", default="full")
|
||||
parser.add_argument("--gradient_accumulation_steps", type=int, default=1, help="Number of updates steps to accumulate before performing a backward/update pass.")
|
||||
|
||||
# lr_scheduler
|
||||
parser.add_argument("--lr_scheduler", type=str, default="constant",
|
||||
help=(
|
||||
'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
|
||||
' "constant", "constant_with_warmup"]'
|
||||
),
|
||||
)
|
||||
parser.add_argument("--num_euler_timesteps", type=int, default=100)
|
||||
parser.add_argument("--lr_num_cycles", type=int, default=1, help="Number of cycles in the learning rate scheduler.")
|
||||
parser.add_argument("--lr_power", type=float, default=1.0, help="Power factor of the polynomial scheduler.",)
|
||||
parser.add_argument("--not_apply_cfg_solver", action="store_true", help="Whether to apply the cfg_solver.")
|
||||
parser.add_argument("--distill_cfg", type=float, default=3.0, help="Distillation coefficient.")
|
||||
# ["euler_linear_quadratic", "pcm", "pcm_linear_qudratic"]
|
||||
parser.add_argument("--scheduler_type", type=str, default="pcm", help="The scheduler type to use.")
|
||||
parser.add_argument("--adv_weight", type=float, default=0.1, help="The weight of the adversarial loss.")
|
||||
parser.add_argument("--discriminator_head_stride", type=int, default=2, help="The stride of the discriminator head.")
|
||||
parser.add_argument("--linear_quadratic_threshold", type=float, default=0.025, help="The threshold of the linear quadratic scheduler.")
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,148 @@
|
||||
from sympy import use
|
||||
import torch
|
||||
import os
|
||||
import torch.distributed as dist
|
||||
from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
|
||||
checkpoint_wrapper,
|
||||
CheckpointImpl,
|
||||
apply_activation_checkpointing,
|
||||
)
|
||||
from peft.utils.other import fsdp_auto_wrap_policy
|
||||
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
StateDictType,
|
||||
FullStateDictConfig, # general model non-sharded, non-flattened params
|
||||
LocalStateDictConfig, # flattened params, usable only by FSDP
|
||||
# ShardedStateDictConfig, # un-flattened param but shards, usable by other parallel schemes.
|
||||
)
|
||||
|
||||
from fastvideo.model.modeling_mochi import MochiTransformerBlock
|
||||
|
||||
from functools import partial
|
||||
|
||||
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
|
||||
|
||||
from torch.distributed.fsdp import MixedPrecision, ShardingStrategy
|
||||
import functools
|
||||
|
||||
|
||||
non_reentrant_wrapper = partial(
|
||||
checkpoint_wrapper,
|
||||
checkpoint_impl=CheckpointImpl.NO_REENTRANT,
|
||||
)
|
||||
|
||||
check_fn = lambda submodule: isinstance(submodule, MochiTransformerBlock)
|
||||
|
||||
|
||||
def apply_fsdp_checkpointing(model, p=1):
|
||||
# https://github.com/foundation-model-stack/fms-fsdp/blob/408c7516d69ea9b6bcd4c0f5efab26c0f64b3c2d/fms_fsdp/policies/ac_handler.py#L16
|
||||
"""apply activation checkpointing to model
|
||||
returns None as model is updated directly
|
||||
"""
|
||||
print(f"--> applying fdsp activation checkpointing...")
|
||||
|
||||
block_idx = 0
|
||||
cut_off = 1 / 2
|
||||
# when passing p as a fraction number (e.g. 1/3), it will be interpreted
|
||||
# as a string in argv, thus we need eval("1/3") here for fractions.
|
||||
p = eval(p) if isinstance(p, str) else p
|
||||
|
||||
def selective_checkpointing(submodule):
|
||||
nonlocal block_idx
|
||||
nonlocal cut_off
|
||||
|
||||
if isinstance(submodule, MochiTransformerBlock):
|
||||
block_idx += 1
|
||||
if block_idx * p >= cut_off:
|
||||
cut_off += 1
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
apply_activation_checkpointing(
|
||||
model, checkpoint_wrapper_fn=non_reentrant_wrapper, check_fn=selective_checkpointing
|
||||
)
|
||||
|
||||
|
||||
float32 = MixedPrecision(
|
||||
param_dtype=torch.float32,
|
||||
# Gradient communication precision.
|
||||
reduce_dtype=torch.float32,
|
||||
# Buffer precision.
|
||||
buffer_dtype=torch.float32,
|
||||
cast_forward_inputs=False
|
||||
)
|
||||
|
||||
|
||||
|
||||
def get_dit_fsdp_kwargs(sharding_strategy, use_lora=False, cpu_offload=False):
|
||||
if use_lora:
|
||||
auto_wrap_policy = fsdp_auto_wrap_policy
|
||||
else:
|
||||
auto_wrap_policy = functools.partial(
|
||||
transformer_auto_wrap_policy,
|
||||
transformer_layer_cls={
|
||||
MochiTransformerBlock,
|
||||
},
|
||||
)
|
||||
|
||||
# we use float32 for fsdp but autocast during training
|
||||
mixed_precision = float32
|
||||
|
||||
if sharding_strategy == "full":
|
||||
sharding_strategy = ShardingStrategy.FULL_SHARD
|
||||
elif sharding_strategy == "hybrid_full":
|
||||
sharding_strategy = ShardingStrategy.HYBRID_SHARD
|
||||
elif sharding_strategy == "none":
|
||||
sharding_strategy = ShardingStrategy.NO_SHARD
|
||||
auto_wrap_policy = None
|
||||
elif sharding_strategy == "hybrid_zero2":
|
||||
sharding_strategy = ShardingStrategy._HYBRID_SHARD_ZERO2
|
||||
|
||||
device_id = torch.cuda.current_device()
|
||||
cpu_offload=torch.distributed.fsdp.CPUOffload(offload_params=True) if cpu_offload else None
|
||||
fsdp_kwargs = {
|
||||
"auto_wrap_policy": auto_wrap_policy,
|
||||
"mixed_precision": mixed_precision,
|
||||
"sharding_strategy": sharding_strategy,
|
||||
"device_id": device_id,
|
||||
"limit_all_gathers": True,
|
||||
"cpu_offload": cpu_offload,
|
||||
}
|
||||
|
||||
# Add LoRA-specific settings when LoRA is enabled
|
||||
if use_lora:
|
||||
fsdp_kwargs.update({
|
||||
"use_orig_params": False, # Required for LoRA memory savings
|
||||
"sync_module_states": True,
|
||||
})
|
||||
|
||||
return fsdp_kwargs
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
def get_discriminator_fsdp_kwargs():
|
||||
|
||||
auto_wrap_policy = None
|
||||
|
||||
|
||||
# Use existing mixed precision settings
|
||||
|
||||
mixed_precision = float32
|
||||
sharding_strategy = ShardingStrategy.NO_SHARD
|
||||
device_id = torch.cuda.current_device()
|
||||
fsdp_kwargs = {
|
||||
"auto_wrap_policy": auto_wrap_policy,
|
||||
"mixed_precision": mixed_precision,
|
||||
"sharding_strategy": sharding_strategy,
|
||||
"device_id": device_id,
|
||||
"limit_all_gathers": True,
|
||||
}
|
||||
|
||||
return fsdp_kwargs
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,123 @@
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
import numpy as np
|
||||
from torch.nn.utils.parametrizations import spectral_norm
|
||||
import os
|
||||
class DummyDiscriminator(nn.Module):
|
||||
def __init__(self, dim_in, num_layers):
|
||||
super().__init__()
|
||||
self.layers = nn.ModuleList()
|
||||
for _ in range(num_layers):
|
||||
self.layers.append(nn.Linear(dim_in, 1))
|
||||
|
||||
def forward(self, features):
|
||||
logits = []
|
||||
for layer, feature in zip(self.layers, features):
|
||||
mean = feature.mean(dim=1)
|
||||
logits.append(layer(mean))
|
||||
return torch.cat(logits, dim=1)
|
||||
|
||||
|
||||
class ResidualBlock(nn.Module):
|
||||
def __init__(self, fn):
|
||||
super().__init__()
|
||||
self.fn = fn
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return (self.fn(x) + x) / np.sqrt(2)
|
||||
|
||||
|
||||
class SpectralConv1d(nn.Module):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__()
|
||||
self.conv = spectral_norm(nn.Conv1d(*args, **kwargs))
|
||||
def forward(self, x):
|
||||
return self.conv(x)
|
||||
|
||||
class BatchNormLocal(nn.Module):
|
||||
def __init__(self, num_features: int, affine: bool = True, virtual_bs: int = 8, eps: float = 1e-5):
|
||||
super().__init__()
|
||||
self.virtual_bs = virtual_bs
|
||||
self.eps = eps
|
||||
self.affine = affine
|
||||
|
||||
if self.affine:
|
||||
self.weight = nn.Parameter(torch.ones(num_features))
|
||||
self.bias = nn.Parameter(torch.zeros(num_features))
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
shape = x.size()
|
||||
|
||||
# Calculate stats.
|
||||
mean = x.mean([0, 2], keepdim=True)
|
||||
var = x.var([0, 2], keepdim=True, unbiased=False)
|
||||
x = (x - mean) / (torch.sqrt(var + self.eps))
|
||||
|
||||
if self.affine:
|
||||
x = x * self.weight[None, :, None] + self.bias[None, :, None]
|
||||
|
||||
return x.view(shape)
|
||||
|
||||
def make_block(channels: int, kernel_size: int) -> nn.Module:
|
||||
return nn.Sequential(
|
||||
SpectralConv1d(
|
||||
channels,
|
||||
channels,
|
||||
kernel_size = kernel_size,
|
||||
padding = kernel_size//2,
|
||||
padding_mode = 'circular',
|
||||
),
|
||||
BatchNormLocal(channels),
|
||||
nn.LeakyReLU(0.2, True),
|
||||
)
|
||||
|
||||
class DiscHead(nn.Module):
|
||||
def __init__(self, feature_dim: int, text_c_dim: int, cmap_dim: int = 64, cnn_dim=512):
|
||||
super().__init__()
|
||||
self.channels = feature_dim
|
||||
self.text_c_dim = text_c_dim
|
||||
self.cmap_dim = cmap_dim
|
||||
self.down_proj = SpectralConv1d(feature_dim, cnn_dim, kernel_size=1, padding=0)
|
||||
self.main = nn.Sequential(
|
||||
make_block(cnn_dim, kernel_size=1),
|
||||
ResidualBlock(make_block(cnn_dim, kernel_size=9))
|
||||
)
|
||||
|
||||
self.cmapper = nn.Linear(self.text_c_dim, cmap_dim)
|
||||
self.cls = SpectralConv1d(cnn_dim, cmap_dim, kernel_size=1, padding=0)
|
||||
|
||||
def forward(self, x: torch.Tensor, c: torch.Tensor) -> torch.Tensor:
|
||||
h = self.down_proj(x)
|
||||
h = self.main(h)
|
||||
out = self.cls(h)
|
||||
|
||||
cmap = self.cmapper(c).unsqueeze(-1)
|
||||
out = (out * cmap).sum(1, keepdim=True) * (1 / np.sqrt(self.cmap_dim))
|
||||
|
||||
return out
|
||||
|
||||
|
||||
class LADDDiscriminator(nn.Module):
|
||||
def __init__(self, feature_dim, text_cond_dim, num_layers, layers_stride):
|
||||
super().__init__()
|
||||
heads = []
|
||||
for i in range(0, num_layers, layers_stride):
|
||||
heads.append(DiscHead(feature_dim, text_cond_dim))
|
||||
self.heads = nn.ModuleList(heads)
|
||||
self.layers_stride = layers_stride
|
||||
self.num_layers = num_layers
|
||||
|
||||
def forward(self, features, text_conditions) -> torch.Tensor:
|
||||
text_conditions = text_conditions.mean(1)
|
||||
# layer, B, L, C -> layer, B, C, L
|
||||
features = features.transpose(2, 3)
|
||||
logits = []
|
||||
for i in range(0, self.num_layers, self.layers_stride):
|
||||
head = self.heads[i//self.layers_stride]
|
||||
feat = features[i]
|
||||
logits.append(head(feat, text_conditions).view(feat.size(0), -1))
|
||||
logits = torch.cat(logits, dim=1)
|
||||
|
||||
|
||||
return logits
|
||||
@@ -0,0 +1,39 @@
|
||||
import torch
|
||||
mochi_latents_mean = torch.tensor([
|
||||
-0.06730895953510081,
|
||||
-0.038011381506090416,
|
||||
-0.07477820912866141,
|
||||
-0.05565264470995561,
|
||||
0.012767231469026969,
|
||||
-0.04703542746246419,
|
||||
0.043896967884726704,
|
||||
-0.09346305707025976,
|
||||
-0.09918314763016893,
|
||||
-0.008729793427399178,
|
||||
-0.011931556316503654,
|
||||
-0.0321993391887285
|
||||
]).view(1, 12, 1, 1, 1)
|
||||
mochi_latents_std = torch.tensor([
|
||||
0.9263795028493863,
|
||||
0.9248894543193766,
|
||||
0.9393059390890617,
|
||||
0.959253732819592,
|
||||
0.8244560132752793,
|
||||
0.917259975397747,
|
||||
0.9294154431013696,
|
||||
1.3720942357788521,
|
||||
0.881393668867029,
|
||||
0.9168315692124348,
|
||||
0.9185249279345552,
|
||||
0.9274757570805041
|
||||
]).view(1, 12, 1, 1, 1)
|
||||
mochi_scaling_factor = 1.0
|
||||
|
||||
|
||||
def normalize_mochi_dit_input(latents):
|
||||
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
|
||||
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
|
||||
latents = (latents - latents_mean) / latents_std
|
||||
return latents
|
||||
|
||||
|
||||
@@ -0,0 +1,654 @@
|
||||
# Copyright 2024 The Genmo team and The HuggingFace Team.
|
||||
# All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import diffusers
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.utils import is_torch_version, logging
|
||||
from diffusers.utils.torch_utils import maybe_allow_in_graph
|
||||
from diffusers.models.attention import FeedForward as HF_FeedForward
|
||||
from diffusers.models.attention_processor import Attention
|
||||
from diffusers.models.embeddings import MochiCombinedTimestepCaptionEmbedding, PatchEmbed
|
||||
from diffusers.models.modeling_outputs import Transformer2DModelOutput
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from fastvideo.model.norm import MochiLayerNormContinuous, MochiRMSNormZero, MochiModulatedRMSNorm, MochiRMSNorm
|
||||
from diffusers.models.normalization import AdaLayerNormContinuous
|
||||
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
from fastvideo.utils.communications import all_gather, all_to_all_4D
|
||||
import torch.nn.functional as F
|
||||
from diffusers.utils.torch_utils import is_torch_version, maybe_allow_in_graph
|
||||
from einops import rearrange
|
||||
|
||||
import numbers
|
||||
from flash_attn import flash_attn_varlen_qkvpacked_func
|
||||
from flash_attn.bert_padding import pad_input, unpad_input
|
||||
|
||||
from liger_kernel.ops.swiglu import LigerSiLUMulFunction
|
||||
class FeedForward(HF_FeedForward):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
dim_out: Optional[int] = None,
|
||||
mult: int = 4,
|
||||
dropout: float = 0.0,
|
||||
activation_fn: str = "geglu",
|
||||
final_dropout: bool = False,
|
||||
inner_dim=None,
|
||||
bias: bool = True,
|
||||
):
|
||||
super().__init__(dim, dim_out, mult, dropout, activation_fn, final_dropout, inner_dim, bias)
|
||||
assert activation_fn == "swiglu"
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.net[0].proj(hidden_states)
|
||||
hidden_states, gate = hidden_states.chunk(2, dim=-1)
|
||||
|
||||
return self.net[2](
|
||||
LigerSiLUMulFunction.apply(gate, hidden_states)
|
||||
)
|
||||
|
||||
def flash_attn_no_pad(qkv, key_padding_mask, causal=False, dropout_p=0.0, softmax_scale=None):
|
||||
# adapted from https://github.com/Dao-AILab/flash-attention/blob/13403e81157ba37ca525890f2f0f2137edf75311/flash_attn/flash_attention.py#L27
|
||||
batch_size = qkv.shape[0]
|
||||
seqlen = qkv.shape[1]
|
||||
nheads = qkv.shape[-2]
|
||||
x = rearrange(qkv, 'b s three h d -> b s (three h d)')
|
||||
x_unpad, indices, cu_seqlens, max_s, used_seqlens_in_batch = unpad_input(x, key_padding_mask)
|
||||
|
||||
|
||||
x_unpad = rearrange(x_unpad, 'nnz (three h d) -> nnz three h d', three=3, h=nheads)
|
||||
output_unpad = flash_attn_varlen_qkvpacked_func(
|
||||
x_unpad, cu_seqlens, max_s, dropout_p,
|
||||
softmax_scale=softmax_scale, causal=causal
|
||||
)
|
||||
output = rearrange(pad_input(rearrange(output_unpad, 'nnz h d -> nnz (h d)'),
|
||||
indices, batch_size, seqlen),
|
||||
'b s (h d) -> b s h d', h=nheads)
|
||||
return output
|
||||
|
||||
class MochiAttention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
query_dim: int,
|
||||
processor: "MochiAttnProcessor2_0",
|
||||
heads: int = 8,
|
||||
dim_head: int = 64,
|
||||
dropout: float = 0.0,
|
||||
bias: bool = False,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
added_proj_bias: Optional[bool] = True,
|
||||
out_dim: int = None,
|
||||
out_context_dim: int = None,
|
||||
out_bias: bool = True,
|
||||
context_pre_only: bool = False,
|
||||
eps: float = 1e-5,
|
||||
):
|
||||
super().__init__()
|
||||
self.inner_dim = out_dim if out_dim is not None else dim_head * heads
|
||||
self.out_dim = out_dim if out_dim is not None else query_dim
|
||||
self.out_context_dim = out_context_dim if out_context_dim else query_dim
|
||||
self.context_pre_only = context_pre_only
|
||||
|
||||
self.heads = out_dim // dim_head if out_dim is not None else heads
|
||||
|
||||
self.norm_q = MochiRMSNorm(dim_head, eps)
|
||||
self.norm_k = MochiRMSNorm(dim_head, eps)
|
||||
self.norm_added_q = MochiRMSNorm(dim_head, eps)
|
||||
self.norm_added_k = MochiRMSNorm(dim_head, eps)
|
||||
|
||||
self.to_q = nn.Linear(query_dim, self.inner_dim, bias=bias)
|
||||
self.to_k = nn.Linear(query_dim, self.inner_dim, bias=bias)
|
||||
self.to_v = nn.Linear(query_dim, self.inner_dim, bias=bias)
|
||||
|
||||
self.add_k_proj = nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias)
|
||||
self.add_v_proj = nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias)
|
||||
if self.context_pre_only is not None:
|
||||
self.add_q_proj = nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias)
|
||||
|
||||
self.to_out = nn.ModuleList([])
|
||||
self.to_out.append(nn.Linear(self.inner_dim, self.out_dim, bias=out_bias))
|
||||
self.to_out.append(nn.Dropout(dropout))
|
||||
|
||||
if not self.context_pre_only:
|
||||
self.to_add_out = nn.Linear(self.inner_dim, self.out_context_dim, bias=out_bias)
|
||||
|
||||
self.processor = processor
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
**kwargs,
|
||||
):
|
||||
return self.processor(
|
||||
self,
|
||||
hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
|
||||
class MochiAttnProcessor2_0:
|
||||
"""Attention processor used in Mochi."""
|
||||
|
||||
def __init__(self):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("MochiAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
# [b, s, h * d]
|
||||
query = attn.to_q(hidden_states)
|
||||
key = attn.to_k(hidden_states)
|
||||
value = attn.to_v(hidden_states)
|
||||
|
||||
# [b, s, h=24, d=128]
|
||||
query = query.unflatten(2, (attn.heads, -1))
|
||||
key = key.unflatten(2, (attn.heads, -1))
|
||||
value = value.unflatten(2, (attn.heads, -1))
|
||||
|
||||
|
||||
if attn.norm_q is not None:
|
||||
query = attn.norm_q(query)
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key)
|
||||
# [b, 256, h * d]
|
||||
encoder_query = attn.add_q_proj(encoder_hidden_states)
|
||||
encoder_key = attn.add_k_proj(encoder_hidden_states)
|
||||
encoder_value = attn.add_v_proj(encoder_hidden_states)
|
||||
|
||||
# [b, 256, h=24, d=128]
|
||||
encoder_query = encoder_query.unflatten(2, (attn.heads, -1))
|
||||
encoder_key = encoder_key.unflatten(2, (attn.heads, -1))
|
||||
encoder_value = encoder_value.unflatten(2, (attn.heads, -1))
|
||||
|
||||
|
||||
if attn.norm_added_q is not None:
|
||||
encoder_query = attn.norm_added_q(encoder_query)
|
||||
if attn.norm_added_k is not None:
|
||||
encoder_key = attn.norm_added_k(encoder_key)
|
||||
|
||||
if image_rotary_emb is not None:
|
||||
freqs_cos, freqs_sin = image_rotary_emb[0], image_rotary_emb[1]
|
||||
# shard the head dimension
|
||||
if get_sequence_parallel_state():
|
||||
# B, S, H, D to (S, B,) H, D
|
||||
# batch_size, seq_len, attn_heads, head_dim
|
||||
query = all_to_all_4D(query, scatter_dim=2, gather_dim=1)
|
||||
key = all_to_all_4D(key, scatter_dim=2, gather_dim=1)
|
||||
value = all_to_all_4D(value, scatter_dim=2, gather_dim=1)
|
||||
|
||||
|
||||
def shrink_head(encoder_state, dim):
|
||||
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
|
||||
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
|
||||
encoder_query = shrink_head(encoder_query, dim=2)
|
||||
encoder_key = shrink_head(encoder_key, dim=2)
|
||||
encoder_value = shrink_head(encoder_value, dim=2)
|
||||
if image_rotary_emb is not None:
|
||||
freqs_cos = shrink_head(freqs_cos, dim=1)
|
||||
freqs_sin = shrink_head(freqs_sin, dim=1)
|
||||
|
||||
|
||||
|
||||
if image_rotary_emb is not None:
|
||||
def apply_rotary_emb(x, freqs_cos, freqs_sin):
|
||||
x_even = x[..., 0::2].float()
|
||||
x_odd = x[..., 1::2].float()
|
||||
cos = (x_even * freqs_cos - x_odd * freqs_sin).to(x.dtype)
|
||||
sin = (x_even * freqs_sin + x_odd * freqs_cos).to(x.dtype)
|
||||
|
||||
return torch.stack([cos, sin], dim=-1).flatten(-2)
|
||||
query = apply_rotary_emb(query, freqs_cos, freqs_sin)
|
||||
key = apply_rotary_emb(key, freqs_cos, freqs_sin)
|
||||
|
||||
# query, key, value = query.transpose(1, 2), key.transpose(1, 2), value.transpose(1, 2)
|
||||
# encoder_query, encoder_key, encoder_value = (
|
||||
# encoder_query.transpose(1, 2),
|
||||
# encoder_key.transpose(1, 2),
|
||||
# encoder_value.transpose(1, 2),
|
||||
# )
|
||||
# [b, s, h, d]
|
||||
sequence_length = query.size(1)
|
||||
encoder_sequence_length = encoder_query.size(1)
|
||||
|
||||
# H
|
||||
query = torch.cat([query, encoder_query], dim=1).unsqueeze(2)
|
||||
key = torch.cat([key, encoder_key], dim=1).unsqueeze(2)
|
||||
value = torch.cat([value, encoder_value], dim=1).unsqueeze(2)
|
||||
# B, S, 3, H, D
|
||||
qkv = torch.cat([query, key, value], dim=2)
|
||||
|
||||
attn_mask = encoder_attention_mask[:, :].bool()
|
||||
attn_mask = F.pad(attn_mask, (sequence_length, 0), value=True)
|
||||
hidden_states = flash_attn_no_pad(qkv, attn_mask, causal=False, dropout_p=0.0, softmax_scale=None)
|
||||
|
||||
# hidden_states = F.scaled_dot_product_attention(query, key, value, attn_mask = None, dropout_p=0.0, is_causal=False)
|
||||
|
||||
# valid_lengths = encoder_attention_mask.sum(dim=1) + sequence_length
|
||||
# def no_padding_mask(score, b, h, q_idx, kv_idx):
|
||||
# return torch.where(kv_idx < valid_lengths[b],score, -float("inf"))
|
||||
|
||||
# hidden_states = flex_attention(query, key, value, score_mod=no_padding_mask)
|
||||
if get_sequence_parallel_state():
|
||||
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
|
||||
(sequence_length, encoder_sequence_length), dim=1
|
||||
)
|
||||
# B, S, H, D
|
||||
hidden_states = all_to_all_4D(hidden_states, scatter_dim=1, gather_dim=2)
|
||||
encoder_hidden_states = all_gather(encoder_hidden_states, dim=2).contiguous()
|
||||
hidden_states = hidden_states.flatten(2, 3)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
encoder_hidden_states = encoder_hidden_states.flatten(2, 3)
|
||||
encoder_hidden_states = encoder_hidden_states.to(query.dtype)
|
||||
else:
|
||||
hidden_states = hidden_states.flatten(2, 3)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
|
||||
(sequence_length, encoder_sequence_length), dim=1
|
||||
)
|
||||
|
||||
|
||||
|
||||
# linear proj
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
# dropout
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if hasattr(attn, "to_add_out"):
|
||||
encoder_hidden_states = attn.to_add_out(encoder_hidden_states)
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
@maybe_allow_in_graph
|
||||
class MochiTransformerBlock(nn.Module):
|
||||
r"""
|
||||
Transformer block used in [Mochi](https://huggingface.co/genmo/mochi-1-preview).
|
||||
|
||||
Args:
|
||||
dim (`int`):
|
||||
The number of channels in the input and output.
|
||||
num_attention_heads (`int`):
|
||||
The number of heads to use for multi-head attention.
|
||||
attention_head_dim (`int`):
|
||||
The number of channels in each head.
|
||||
qk_norm (`str`, defaults to `"rms_norm"`):
|
||||
The normalization layer to use.
|
||||
activation_fn (`str`, defaults to `"swiglu"`):
|
||||
Activation function to use in feed-forward.
|
||||
context_pre_only (`bool`, defaults to `False`):
|
||||
Whether or not to process context-related conditions with additional layers.
|
||||
eps (`float`, defaults to `1e-6`):
|
||||
Epsilon value for normalization layers.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
pooled_projection_dim: int,
|
||||
qk_norm: str = "rms_norm",
|
||||
activation_fn: str = "swiglu",
|
||||
context_pre_only: bool = False,
|
||||
eps: float = 1e-6,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.context_pre_only = context_pre_only
|
||||
self.ff_inner_dim = (4 * dim * 2) // 3
|
||||
self.ff_context_inner_dim = (4 * pooled_projection_dim * 2) // 3
|
||||
|
||||
self.norm1 = MochiRMSNormZero(dim, 4 * dim, eps=eps, elementwise_affine=False)
|
||||
|
||||
if not context_pre_only:
|
||||
self.norm1_context = MochiRMSNormZero(dim, 4 * pooled_projection_dim, eps=eps, elementwise_affine=False)
|
||||
else:
|
||||
self.norm1_context = MochiLayerNormContinuous(
|
||||
embedding_dim=pooled_projection_dim,
|
||||
conditioning_embedding_dim=dim,
|
||||
eps=eps,
|
||||
)
|
||||
|
||||
self.attn1 = MochiAttention(
|
||||
query_dim=dim,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
bias=False,
|
||||
added_kv_proj_dim=pooled_projection_dim,
|
||||
added_proj_bias=False,
|
||||
out_dim=dim,
|
||||
out_context_dim=pooled_projection_dim,
|
||||
context_pre_only=context_pre_only,
|
||||
processor=MochiAttnProcessor2_0(),
|
||||
eps=1e-5,
|
||||
)
|
||||
|
||||
# TODO(aryan): norm_context layers are not needed when `context_pre_only` is True
|
||||
self.norm2 = MochiModulatedRMSNorm(eps=eps)
|
||||
self.norm2_context = MochiModulatedRMSNorm(eps=eps) if not self.context_pre_only else None
|
||||
|
||||
self.norm3 = MochiModulatedRMSNorm(eps)
|
||||
self.norm3_context = MochiModulatedRMSNorm(eps=eps) if not self.context_pre_only else None
|
||||
|
||||
self.ff = FeedForward(dim, inner_dim=self.ff_inner_dim, activation_fn=activation_fn, bias=False)
|
||||
self.ff_context = None
|
||||
if not context_pre_only:
|
||||
self.ff_context = FeedForward(
|
||||
pooled_projection_dim,
|
||||
inner_dim=self.ff_context_inner_dim,
|
||||
activation_fn=activation_fn,
|
||||
bias=False,
|
||||
)
|
||||
|
||||
self.norm4 = MochiModulatedRMSNorm(eps=eps)
|
||||
self.norm4_context = MochiModulatedRMSNorm(eps=eps)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
output_attn = False,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
norm_hidden_states, gate_msa, scale_mlp, gate_mlp = self.norm1(hidden_states, temb)
|
||||
|
||||
if not self.context_pre_only:
|
||||
norm_encoder_hidden_states, enc_gate_msa, enc_scale_mlp, enc_gate_mlp = self.norm1_context(
|
||||
encoder_hidden_states, temb
|
||||
)
|
||||
else:
|
||||
norm_encoder_hidden_states = self.norm1_context(encoder_hidden_states, temb)
|
||||
|
||||
attn_hidden_states, context_attn_hidden_states = self.attn1(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=norm_encoder_hidden_states,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
encoder_attention_mask=encoder_attention_mask
|
||||
)
|
||||
|
||||
hidden_states = hidden_states + self.norm2(attn_hidden_states, torch.tanh(gate_msa).unsqueeze(1))
|
||||
norm_hidden_states = self.norm3(hidden_states, (1 + scale_mlp.unsqueeze(1).to(torch.float32)))
|
||||
ff_output = self.ff(norm_hidden_states)
|
||||
hidden_states = hidden_states + self.norm4(ff_output, torch.tanh(gate_mlp).unsqueeze(1))
|
||||
|
||||
if not self.context_pre_only:
|
||||
encoder_hidden_states = encoder_hidden_states + self.norm2_context(
|
||||
context_attn_hidden_states, torch.tanh(enc_gate_msa).unsqueeze(1)
|
||||
)
|
||||
norm_encoder_hidden_states = self.norm3_context(
|
||||
encoder_hidden_states, (1 + enc_scale_mlp.unsqueeze(1).to(torch.float32))
|
||||
)
|
||||
context_ff_output = self.ff_context(norm_encoder_hidden_states)
|
||||
encoder_hidden_states = encoder_hidden_states + self.norm4_context(
|
||||
context_ff_output, torch.tanh(enc_gate_mlp).unsqueeze(1)
|
||||
)
|
||||
|
||||
if not output_attn:
|
||||
attn_hidden_states = None
|
||||
return hidden_states, encoder_hidden_states, attn_hidden_states
|
||||
|
||||
|
||||
class MochiRoPE(nn.Module):
|
||||
r"""
|
||||
RoPE implementation used in [Mochi](https://huggingface.co/genmo/mochi-1-preview).
|
||||
|
||||
Args:
|
||||
base_height (`int`, defaults to `192`):
|
||||
Base height used to compute interpolation scale for rotary positional embeddings.
|
||||
base_width (`int`, defaults to `192`):
|
||||
Base width used to compute interpolation scale for rotary positional embeddings.
|
||||
"""
|
||||
|
||||
def __init__(self, base_height: int = 192, base_width: int = 192) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.target_area = base_height * base_width
|
||||
|
||||
def _centers(self, start, stop, num, device, dtype) -> torch.Tensor:
|
||||
edges = torch.linspace(start, stop, num + 1, device=device, dtype=dtype)
|
||||
return (edges[:-1] + edges[1:]) / 2
|
||||
|
||||
def _get_positions(
|
||||
self,
|
||||
num_frames: int,
|
||||
height: int,
|
||||
width: int,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
) -> torch.Tensor:
|
||||
scale = (self.target_area / (height * width)) ** 0.5
|
||||
t = torch.arange(num_frames * nccl_info.sp_size, device=device, dtype=dtype)
|
||||
h = self._centers(-height * scale / 2, height * scale / 2, height, device, dtype)
|
||||
w = self._centers(-width * scale / 2, width * scale / 2, width, device, dtype)
|
||||
|
||||
grid_t, grid_h, grid_w = torch.meshgrid(t, h, w, indexing="ij")
|
||||
|
||||
positions = torch.stack([grid_t, grid_h, grid_w], dim=-1).view(-1, 3)
|
||||
return positions
|
||||
|
||||
def _create_rope(self, freqs: torch.Tensor, pos: torch.Tensor) -> torch.Tensor:
|
||||
with torch.autocast(freqs.device.type, enabled=False):
|
||||
# Always run ROPE freqs computation in FP32
|
||||
freqs = torch.einsum("nd,dhf->nhf", pos.to(torch.float32), freqs.to(torch.float32))
|
||||
freqs_cos = torch.cos(freqs)
|
||||
freqs_sin = torch.sin(freqs)
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
def forward(
|
||||
self,
|
||||
pos_frequencies: torch.Tensor,
|
||||
num_frames: int,
|
||||
height: int,
|
||||
width: int,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
pos = self._get_positions(num_frames, height, width, device, dtype)
|
||||
rope_cos, rope_sin = self._create_rope(pos_frequencies, pos)
|
||||
return rope_cos, rope_sin
|
||||
|
||||
|
||||
@maybe_allow_in_graph
|
||||
class MochiTransformer3DModel(ModelMixin, ConfigMixin):
|
||||
r"""
|
||||
A Transformer model for video-like data introduced in [Mochi](https://huggingface.co/genmo/mochi-1-preview).
|
||||
|
||||
Args:
|
||||
patch_size (`int`, defaults to `2`):
|
||||
The size of the patches to use in the patch embedding layer.
|
||||
num_attention_heads (`int`, defaults to `24`):
|
||||
The number of heads to use for multi-head attention.
|
||||
attention_head_dim (`int`, defaults to `128`):
|
||||
The number of channels in each head.
|
||||
num_layers (`int`, defaults to `48`):
|
||||
The number of layers of Transformer blocks to use.
|
||||
in_channels (`int`, defaults to `12`):
|
||||
The number of channels in the input.
|
||||
out_channels (`int`, *optional*, defaults to `None`):
|
||||
The number of channels in the output.
|
||||
qk_norm (`str`, defaults to `"rms_norm"`):
|
||||
The normalization layer to use.
|
||||
text_embed_dim (`int`, defaults to `4096`):
|
||||
Input dimension of text embeddings from the text encoder.
|
||||
time_embed_dim (`int`, defaults to `256`):
|
||||
Output dimension of timestep embeddings.
|
||||
activation_fn (`str`, defaults to `"swiglu"`):
|
||||
Activation function to use in feed-forward.
|
||||
max_sequence_length (`int`, defaults to `256`):
|
||||
The maximum sequence length of text embeddings supported.
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: int = 2,
|
||||
num_attention_heads: int = 24,
|
||||
attention_head_dim: int = 128,
|
||||
num_layers: int = 48,
|
||||
pooled_projection_dim: int = 1536,
|
||||
in_channels: int = 12,
|
||||
out_channels: Optional[int] = None,
|
||||
qk_norm: str = "rms_norm",
|
||||
text_embed_dim: int = 4096,
|
||||
time_embed_dim: int = 256,
|
||||
activation_fn: str = "swiglu",
|
||||
max_sequence_length: int = 256,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
out_channels = out_channels or in_channels
|
||||
|
||||
self.patch_embed = PatchEmbed(
|
||||
patch_size=patch_size,
|
||||
in_channels=in_channels,
|
||||
embed_dim=inner_dim,
|
||||
pos_embed_type=None,
|
||||
)
|
||||
|
||||
self.time_embed = MochiCombinedTimestepCaptionEmbedding(
|
||||
embedding_dim=inner_dim,
|
||||
pooled_projection_dim=pooled_projection_dim,
|
||||
text_embed_dim=text_embed_dim,
|
||||
time_embed_dim=time_embed_dim,
|
||||
num_attention_heads=8,
|
||||
)
|
||||
|
||||
self.pos_frequencies = nn.Parameter(torch.full((3, num_attention_heads, attention_head_dim // 2), 0.0))
|
||||
self.rope = MochiRoPE()
|
||||
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
MochiTransformerBlock(
|
||||
dim=inner_dim,
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
pooled_projection_dim=pooled_projection_dim,
|
||||
qk_norm=qk_norm,
|
||||
activation_fn=activation_fn,
|
||||
context_pre_only=i == num_layers - 1,
|
||||
)
|
||||
for i in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
self.norm_out = AdaLayerNormContinuous(
|
||||
inner_dim, inner_dim, elementwise_affine=False, eps=1e-6, norm_type="layer_norm"
|
||||
)
|
||||
self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * out_channels)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def _set_gradient_checkpointing(self, module, value=False):
|
||||
if hasattr(module, "gradient_checkpointing"):
|
||||
module.gradient_checkpointing = value
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
output_attn = False,
|
||||
return_dict: bool = False,
|
||||
) -> torch.Tensor:
|
||||
assert return_dict is False, "return_dict is not supported in MochiTransformer3DModel"
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p = self.config.patch_size
|
||||
|
||||
post_patch_height = height // p
|
||||
post_patch_width = width // p
|
||||
# Peiyuan: This is hacked to force mochi to follow the behaviour of SD3 and Flux
|
||||
timestep = 1000 - timestep
|
||||
temb, encoder_hidden_states = self.time_embed(
|
||||
timestep, encoder_hidden_states, encoder_attention_mask, hidden_dtype=hidden_states.dtype
|
||||
)
|
||||
|
||||
hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1)
|
||||
hidden_states = self.patch_embed(hidden_states)
|
||||
hidden_states = hidden_states.unflatten(0, (batch_size, -1)).flatten(1, 2)
|
||||
|
||||
image_rotary_emb = self.rope(
|
||||
self.pos_frequencies,
|
||||
num_frames,
|
||||
post_patch_height,
|
||||
post_patch_width,
|
||||
device=hidden_states.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
attn_outputs_list = []
|
||||
for i, block in enumerate(self.transformer_blocks):
|
||||
if self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module):
|
||||
def custom_forward(*inputs):
|
||||
return module(*inputs)
|
||||
|
||||
return custom_forward
|
||||
|
||||
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
|
||||
hidden_states, encoder_hidden_states, attn_outputs = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(block),
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
temb,
|
||||
image_rotary_emb,
|
||||
output_attn,
|
||||
**ckpt_kwargs,
|
||||
)
|
||||
else:
|
||||
hidden_states, encoder_hidden_states, attn_outputs = block(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
temb=temb,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
output_attn = output_attn,
|
||||
)
|
||||
attn_outputs_list.append(attn_outputs)
|
||||
|
||||
hidden_states = self.norm_out(hidden_states, temb)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(batch_size, num_frames, post_patch_height, post_patch_width, p, p, -1)
|
||||
hidden_states = hidden_states.permute(0, 6, 1, 2, 4, 3, 5)
|
||||
output = hidden_states.reshape(batch_size, -1, num_frames, height, width)
|
||||
|
||||
if not output_attn :
|
||||
attn_outputs_list = None
|
||||
else:
|
||||
attn_outputs_list = torch.stack(attn_outputs_list, dim=0)
|
||||
# Peiyuan: This is hacked to force mochi to follow the behaviour of SD3 and Flux
|
||||
return (-output, attn_outputs_list)
|
||||
@@ -0,0 +1,124 @@
|
||||
# Copyright 2024 The Genmo team and The HuggingFace Team.
|
||||
# All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import numbers
|
||||
from typing import Dict, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class MochiModulatedRMSNorm(nn.Module):
|
||||
def __init__(self, eps: float):
|
||||
super().__init__()
|
||||
|
||||
self.eps = eps
|
||||
|
||||
def forward(self, hidden_states, scale=None):
|
||||
hidden_states_dtype = hidden_states.dtype
|
||||
hidden_states = hidden_states.to(torch.float32)
|
||||
variance = hidden_states.pow(2).mean(-1, keepdim=True)
|
||||
hidden_states = hidden_states * torch.rsqrt(variance + self.eps)
|
||||
if scale is not None:
|
||||
hidden_states = hidden_states * scale
|
||||
|
||||
hidden_states = hidden_states.to(hidden_states_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class MochiRMSNorm(nn.Module):
|
||||
def __init__(self, dim, eps: float, elementwise_affine=True):
|
||||
super().__init__()
|
||||
|
||||
self.eps = eps
|
||||
if elementwise_affine:
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
else:
|
||||
self.weight = None
|
||||
|
||||
def forward(self, hidden_states):
|
||||
hidden_states_dtype = hidden_states.dtype
|
||||
hidden_states = hidden_states.to(torch.float32)
|
||||
variance = hidden_states.pow(2).mean(-1, keepdim=True)
|
||||
hidden_states = hidden_states * torch.rsqrt(variance + self.eps)
|
||||
if self.weight is not None:
|
||||
# convert into half-precision if necessary
|
||||
if self.weight.dtype in [torch.float16, torch.bfloat16]:
|
||||
hidden_states = hidden_states.to(self.weight.dtype)
|
||||
hidden_states = hidden_states * self.weight
|
||||
hidden_states = hidden_states.to(hidden_states_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class MochiLayerNormContinuous(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
embedding_dim: int,
|
||||
conditioning_embedding_dim: int,
|
||||
eps=1e-5,
|
||||
bias=True,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# AdaLN
|
||||
self.silu = nn.SiLU()
|
||||
self.linear_1 = nn.Linear(conditioning_embedding_dim, embedding_dim, bias=bias)
|
||||
self.norm = MochiModulatedRMSNorm(eps=eps)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
conditioning_embedding: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
input_dtype = x.dtype
|
||||
|
||||
# convert back to the original dtype in case `conditioning_embedding`` is upcasted to float32 (needed for hunyuanDiT)
|
||||
scale = self.linear_1(self.silu(conditioning_embedding).to(x.dtype))
|
||||
x = self.norm(x, (1 + scale.unsqueeze(1).to(torch.float32)))
|
||||
|
||||
return x.to(input_dtype)
|
||||
|
||||
|
||||
class MochiRMSNormZero(nn.Module):
|
||||
r"""
|
||||
Adaptive RMS Norm used in Mochi.
|
||||
Parameters:
|
||||
embedding_dim (`int`): The size of each embedding vector.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, embedding_dim: int, hidden_dim: int, eps: float = 1e-5, elementwise_affine: bool = False
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.silu = nn.SiLU()
|
||||
self.linear = nn.Linear(embedding_dim, hidden_dim)
|
||||
self.norm = MochiModulatedRMSNorm(eps=eps)
|
||||
|
||||
def forward(
|
||||
self, hidden_states: torch.Tensor, emb: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
hidden_states_dtype = hidden_states.dtype
|
||||
|
||||
emb = self.linear(self.silu(emb))
|
||||
scale_msa, gate_msa, scale_mlp, gate_mlp = emb.chunk(4, dim=1)
|
||||
|
||||
hidden_states = self.norm(hidden_states, (1 + scale_msa[:, None].to(torch.float32)))
|
||||
hidden_states = hidden_states.to(hidden_states_dtype)
|
||||
|
||||
return hidden_states, gate_msa, scale_mlp, gate_mlp
|
||||
@@ -0,0 +1,756 @@
|
||||
# Copyright 2024 Black Forest Labs and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import inspect
|
||||
from typing import Callable, Dict, List, Optional, Union
|
||||
import copy
|
||||
import numpy as np
|
||||
import torch
|
||||
from transformers import T5EncoderModel, T5TokenizerFast
|
||||
|
||||
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
|
||||
from diffusers.models.autoencoders import AutoencoderKL
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
|
||||
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import (
|
||||
is_torch_xla_available,
|
||||
logging,
|
||||
replace_example_docstring,
|
||||
)
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.pipelines.mochi.pipeline_output import MochiPipelineOutput
|
||||
from einops import rearrange
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
from fastvideo.utils.communications import all_gather
|
||||
|
||||
if is_torch_xla_available():
|
||||
import torch_xla.core.xla_model as xm
|
||||
|
||||
XLA_AVAILABLE = True
|
||||
else:
|
||||
XLA_AVAILABLE = False
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
EXAMPLE_DOC_STRING = """
|
||||
Examples:
|
||||
```py
|
||||
>>> import torch
|
||||
>>> from diffusers import MochiPipeline
|
||||
>>> from diffusers.utils import export_to_video
|
||||
|
||||
>>> pipe = MochiPipeline.from_pretrained("genmo/mochi-1-preview", torch_dtype=torch.bfloat16)
|
||||
>>> pipe.to("cuda")
|
||||
>>> prompt = "Close-up of a chameleon's eye, with its scaly skin changing color. Ultra high resolution 4k."
|
||||
>>> frames = pipe(prompt, num_inference_steps=28, guidance_scale=3.5).frames[0]
|
||||
>>> export_to_video(frames, "mochi.mp4")
|
||||
```
|
||||
"""
|
||||
|
||||
|
||||
def calculate_shift(
|
||||
image_seq_len,
|
||||
base_seq_len: int = 256,
|
||||
max_seq_len: int = 4096,
|
||||
base_shift: float = 0.5,
|
||||
max_shift: float = 1.16,
|
||||
):
|
||||
m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
|
||||
b = base_shift - m * base_seq_len
|
||||
mu = image_seq_len * m + b
|
||||
return mu
|
||||
|
||||
|
||||
# from: https://github.com/genmoai/models/blob/075b6e36db58f1242921deff83a1066887b9c9e1/src/mochi_preview/infer.py#L77
|
||||
def linear_quadratic_schedule(num_steps, threshold_noise, linear_steps=None):
|
||||
if linear_steps is None:
|
||||
linear_steps = num_steps // 2
|
||||
linear_sigma_schedule = [i * threshold_noise / linear_steps for i in range(linear_steps)]
|
||||
threshold_noise_step_diff = linear_steps - threshold_noise * num_steps
|
||||
quadratic_steps = num_steps - linear_steps
|
||||
quadratic_coef = threshold_noise_step_diff / (linear_steps * quadratic_steps**2)
|
||||
linear_coef = threshold_noise / linear_steps - 2 * threshold_noise_step_diff / (quadratic_steps**2)
|
||||
const = quadratic_coef * (linear_steps**2)
|
||||
quadratic_sigma_schedule = [
|
||||
quadratic_coef * (i**2) + linear_coef * i + const for i in range(linear_steps, num_steps)
|
||||
]
|
||||
sigma_schedule = linear_sigma_schedule + quadratic_sigma_schedule
|
||||
sigma_schedule = [1.0 - x for x in sigma_schedule]
|
||||
return sigma_schedule
|
||||
|
||||
|
||||
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps
|
||||
def retrieve_timesteps(
|
||||
scheduler,
|
||||
num_inference_steps: Optional[int] = None,
|
||||
device: Optional[Union[str, torch.device]] = None,
|
||||
timesteps: Optional[List[int]] = None,
|
||||
sigmas: Optional[List[float]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
r"""
|
||||
Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles
|
||||
custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`.
|
||||
|
||||
Args:
|
||||
scheduler (`SchedulerMixin`):
|
||||
The scheduler to get timesteps from.
|
||||
num_inference_steps (`int`):
|
||||
The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps`
|
||||
must be `None`.
|
||||
device (`str` or `torch.device`, *optional*):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
timesteps (`List[int]`, *optional*):
|
||||
Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed,
|
||||
`num_inference_steps` and `sigmas` must be `None`.
|
||||
sigmas (`List[float]`, *optional*):
|
||||
Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed,
|
||||
`num_inference_steps` and `timesteps` must be `None`.
|
||||
|
||||
Returns:
|
||||
`Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the
|
||||
second element is the number of inference steps.
|
||||
"""
|
||||
if timesteps is not None and sigmas is not None:
|
||||
raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")
|
||||
if timesteps is not None:
|
||||
accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
|
||||
if not accepts_timesteps:
|
||||
raise ValueError(
|
||||
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
|
||||
f" timestep schedules. Please check whether you are using the correct scheduler."
|
||||
)
|
||||
scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
elif sigmas is not None:
|
||||
accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
|
||||
if not accept_sigmas:
|
||||
raise ValueError(
|
||||
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
|
||||
f" sigmas schedules. Please check whether you are using the correct scheduler."
|
||||
)
|
||||
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
else:
|
||||
scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
return timesteps, num_inference_steps
|
||||
|
||||
|
||||
class MochiPipeline(DiffusionPipeline):
|
||||
r"""
|
||||
The mochi pipeline for text-to-video generation.
|
||||
|
||||
Reference: https://github.com/genmoai/models
|
||||
|
||||
Args:
|
||||
transformer ([`MochiTransformer3DModel`]):
|
||||
Conditional Transformer architecture to denoise the encoded video latents.
|
||||
scheduler ([`FlowMatchEulerDiscreteScheduler`]):
|
||||
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
|
||||
vae ([`AutoencoderKL`]):
|
||||
Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations.
|
||||
text_encoder ([`T5EncoderModel`]):
|
||||
[T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5EncoderModel), specifically
|
||||
the [google/t5-v1_1-xxl](https://huggingface.co/google/t5-v1_1-xxl) variant.
|
||||
tokenizer (`CLIPTokenizer`):
|
||||
Tokenizer of class
|
||||
[CLIPTokenizer](https://huggingface.co/docs/transformers/en/model_doc/clip#transformers.CLIPTokenizer).
|
||||
tokenizer (`T5TokenizerFast`):
|
||||
Second Tokenizer of class
|
||||
[T5TokenizerFast](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5TokenizerFast).
|
||||
"""
|
||||
|
||||
model_cpu_offload_seq = "text_encoder->transformer->vae"
|
||||
_optional_components = []
|
||||
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
scheduler: FlowMatchEulerDiscreteScheduler,
|
||||
vae: AutoencoderKL,
|
||||
text_encoder: T5EncoderModel,
|
||||
tokenizer: T5TokenizerFast,
|
||||
transformer: MochiTransformer3DModel,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
tokenizer=tokenizer,
|
||||
transformer=transformer,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
# TODO: determine these scaling factors from model parameters
|
||||
self.vae_spatial_scale_factor = 8
|
||||
self.vae_temporal_scale_factor = 6
|
||||
self.patch_size = 2
|
||||
|
||||
self.video_processor = VideoProcessor(vae_scale_factor=self.vae_spatial_scale_factor)
|
||||
self.tokenizer_max_length = (
|
||||
self.tokenizer.model_max_length if hasattr(self, "tokenizer") and self.tokenizer is not None else 77
|
||||
)
|
||||
self.default_height = 480
|
||||
self.default_width = 848
|
||||
|
||||
# Adapted from diffusers.pipelines.cogvideo.pipeline_cogvideox.CogVideoXPipeline._get_t5_prompt_embeds
|
||||
def _get_t5_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
num_videos_per_prompt: int = 1,
|
||||
max_sequence_length: int = 256,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
device = device or self._execution_device
|
||||
dtype = dtype or self.text_encoder.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length,
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
prompt_attention_mask = text_inputs.attention_mask
|
||||
prompt_attention_mask = prompt_attention_mask.bool().to(device)
|
||||
|
||||
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
|
||||
|
||||
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
|
||||
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, max_sequence_length - 1 : -1])
|
||||
logger.warning(
|
||||
"The following part of your input was truncated because `max_sequence_length` is set to "
|
||||
f" {max_sequence_length} tokens: {removed_text}"
|
||||
)
|
||||
|
||||
prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=prompt_attention_mask)[0]
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
_, seq_len, _ = prompt_embeds.shape
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
|
||||
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1)
|
||||
|
||||
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
|
||||
prompt_attention_mask = prompt_attention_mask.repeat(num_videos_per_prompt, 1)
|
||||
|
||||
return prompt_embeds, prompt_attention_mask
|
||||
|
||||
# Adapted from diffusers.pipelines.cogvideo.pipeline_cogvideox.CogVideoXPipeline.encode_prompt
|
||||
def encode_prompt(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
do_classifier_free_guidance: bool = True,
|
||||
num_videos_per_prompt: int = 1,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.Tensor] = None,
|
||||
prompt_attention_mask: Optional[torch.Tensor] = None,
|
||||
negative_prompt_attention_mask: Optional[torch.Tensor] = None,
|
||||
max_sequence_length: int = 256,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
r"""
|
||||
Encodes the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
prompt to be encoded
|
||||
negative_prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts not to guide the image generation. If not defined, one has to pass
|
||||
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is
|
||||
less than `1`).
|
||||
do_classifier_free_guidance (`bool`, *optional*, defaults to `True`):
|
||||
Whether to use classifier free guidance or not.
|
||||
num_videos_per_prompt (`int`, *optional*, defaults to 1):
|
||||
Number of videos that should be generated per prompt. torch device to place the resulting embeddings on
|
||||
prompt_embeds (`torch.Tensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||
provided, text embeddings will be generated from `prompt` input argument.
|
||||
negative_prompt_embeds (`torch.Tensor`, *optional*):
|
||||
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
|
||||
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
|
||||
argument.
|
||||
device: (`torch.device`, *optional*):
|
||||
torch device
|
||||
dtype: (`torch.dtype`, *optional*):
|
||||
torch dtype
|
||||
"""
|
||||
device = device or self._execution_device
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
if prompt is not None:
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
if prompt_embeds is None:
|
||||
prompt_embeds, prompt_attention_mask = self._get_t5_prompt_embeds(
|
||||
prompt=prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
if do_classifier_free_guidance and negative_prompt_embeds is None:
|
||||
negative_prompt = negative_prompt or ""
|
||||
negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt
|
||||
|
||||
if prompt is not None and type(prompt) is not type(negative_prompt):
|
||||
raise TypeError(
|
||||
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
|
||||
f" {type(prompt)}."
|
||||
)
|
||||
elif batch_size != len(negative_prompt):
|
||||
raise ValueError(
|
||||
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
|
||||
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
|
||||
" the batch size of `prompt`."
|
||||
)
|
||||
|
||||
negative_prompt_embeds, negative_prompt_attention_mask = self._get_t5_prompt_embeds(
|
||||
prompt=negative_prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
return prompt_embeds, prompt_attention_mask, negative_prompt_embeds, negative_prompt_attention_mask
|
||||
|
||||
def check_inputs(
|
||||
self,
|
||||
prompt,
|
||||
height,
|
||||
width,
|
||||
callback_on_step_end_tensor_inputs=None,
|
||||
prompt_embeds=None,
|
||||
negative_prompt_embeds=None,
|
||||
prompt_attention_mask=None,
|
||||
negative_prompt_attention_mask=None,
|
||||
):
|
||||
if height % 8 != 0 or width % 8 != 0:
|
||||
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
|
||||
|
||||
if callback_on_step_end_tensor_inputs is not None and not all(
|
||||
k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs
|
||||
):
|
||||
raise ValueError(
|
||||
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
|
||||
)
|
||||
|
||||
if prompt is not None and prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two."
|
||||
)
|
||||
elif prompt is None and prompt_embeds is None:
|
||||
raise ValueError(
|
||||
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
|
||||
)
|
||||
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
|
||||
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
|
||||
|
||||
if prompt_embeds is not None and prompt_attention_mask is None:
|
||||
raise ValueError("Must provide `prompt_attention_mask` when specifying `prompt_embeds`.")
|
||||
|
||||
if negative_prompt_embeds is not None and negative_prompt_attention_mask is None:
|
||||
raise ValueError("Must provide `negative_prompt_attention_mask` when specifying `negative_prompt_embeds`.")
|
||||
|
||||
if prompt_embeds is not None and negative_prompt_embeds is not None:
|
||||
if prompt_embeds.shape != negative_prompt_embeds.shape:
|
||||
raise ValueError(
|
||||
"`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but"
|
||||
f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`"
|
||||
f" {negative_prompt_embeds.shape}."
|
||||
)
|
||||
if prompt_attention_mask.shape != negative_prompt_attention_mask.shape:
|
||||
raise ValueError(
|
||||
"`prompt_attention_mask` and `negative_prompt_attention_mask` must have the same shape when passed directly, but"
|
||||
f" got: `prompt_attention_mask` {prompt_attention_mask.shape} != `negative_prompt_attention_mask`"
|
||||
f" {negative_prompt_attention_mask.shape}."
|
||||
)
|
||||
|
||||
def enable_vae_slicing(self):
|
||||
r"""
|
||||
Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to
|
||||
compute decoding in several steps. This is useful to save some memory and allow larger batch sizes.
|
||||
"""
|
||||
self.vae.enable_slicing()
|
||||
|
||||
def disable_vae_slicing(self):
|
||||
r"""
|
||||
Disable sliced VAE decoding. If `enable_vae_slicing` was previously enabled, this method will go back to
|
||||
computing decoding in one step.
|
||||
"""
|
||||
self.vae.disable_slicing()
|
||||
|
||||
def enable_vae_tiling(self):
|
||||
r"""
|
||||
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
|
||||
compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow
|
||||
processing larger images.
|
||||
"""
|
||||
self.vae.enable_tiling()
|
||||
|
||||
def disable_vae_tiling(self):
|
||||
r"""
|
||||
Disable tiled VAE decoding. If `enable_vae_tiling` was previously enabled, this method will go back to
|
||||
computing decoding in one step.
|
||||
"""
|
||||
self.vae.disable_tiling()
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
num_frames,
|
||||
dtype,
|
||||
device,
|
||||
generator,
|
||||
latents=None,
|
||||
):
|
||||
height = height // self.vae_spatial_scale_factor
|
||||
width = width // self.vae_spatial_scale_factor
|
||||
num_frames = (num_frames - 1) // self.vae_temporal_scale_factor + 1
|
||||
|
||||
shape = (batch_size, num_channels_latents, num_frames, height, width)
|
||||
|
||||
if latents is not None:
|
||||
return latents.to(device=device, dtype=dtype)
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||
)
|
||||
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
return latents
|
||||
|
||||
@property
|
||||
def guidance_scale(self):
|
||||
return self._guidance_scale
|
||||
|
||||
@property
|
||||
def do_classifier_free_guidance(self):
|
||||
return self._guidance_scale > 1.0
|
||||
|
||||
@property
|
||||
def num_timesteps(self):
|
||||
return self._num_timesteps
|
||||
|
||||
@property
|
||||
def interrupt(self):
|
||||
return self._interrupt
|
||||
|
||||
@torch.no_grad()
|
||||
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_frames: int = 16,
|
||||
num_inference_steps: int = 28,
|
||||
timesteps: List[int] = None,
|
||||
guidance_scale: float = 4.5,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
prompt_attention_mask: Optional[torch.Tensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.Tensor] = None,
|
||||
negative_prompt_attention_mask: Optional[torch.Tensor] = None,
|
||||
output_type: Optional[str] = "pil",
|
||||
return_dict: bool = True,
|
||||
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
max_sequence_length: int = 256,
|
||||
return_all_states = False,
|
||||
):
|
||||
r"""
|
||||
Function invoked when calling the pipeline for generation.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
|
||||
instead.
|
||||
height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||
The height in pixels of the generated image. This is set to 1024 by default for the best results.
|
||||
width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):
|
||||
The width in pixels of the generated image. This is set to 1024 by default for the best results.
|
||||
num_frames (`int`, defaults to 16):
|
||||
The number of video frames to generate
|
||||
num_inference_steps (`int`, *optional*, defaults to 50):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference.
|
||||
timesteps (`List[int]`, *optional*):
|
||||
Custom timesteps to use for the denoising process with schedulers which support a `timesteps` argument
|
||||
in their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is
|
||||
passed will be used. Must be in descending order.
|
||||
guidance_scale (`float`, defaults to `4.5`):
|
||||
Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
|
||||
`guidance_scale` is defined as `w` of equation 2. of [Imagen
|
||||
Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >
|
||||
1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,
|
||||
usually at the expense of lower image quality.
|
||||
num_videos_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of videos to generate per prompt.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html)
|
||||
to make generation deterministic.
|
||||
latents (`torch.Tensor`, *optional*):
|
||||
Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor will ge generated by sampling using the supplied random `generator`.
|
||||
prompt_embeds (`torch.Tensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||
provided, text embeddings will be generated from `prompt` input argument.
|
||||
prompt_attention_mask (`torch.Tensor`, *optional*):
|
||||
Pre-generated attention mask for text embeddings.
|
||||
negative_prompt_embeds (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated negative text embeddings. For PixArt-Sigma this negative prompt should be "". If not
|
||||
provided, negative_prompt_embeds will be generated from `negative_prompt` input argument.
|
||||
negative_prompt_attention_mask (`torch.FloatTensor`, *optional*):
|
||||
Pre-generated attention mask for negative text embeddings.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generate image. Choose between
|
||||
[PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~pipelines.mochi.MochiPipelineOutput`] instead of a plain tuple.
|
||||
callback_on_step_end (`Callable`, *optional*):
|
||||
A function that calls at the end of each denoising steps during the inference. The function is called
|
||||
with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
|
||||
callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by
|
||||
`callback_on_step_end_tensor_inputs`.
|
||||
callback_on_step_end_tensor_inputs (`List`, *optional*):
|
||||
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
|
||||
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
|
||||
`._callback_tensor_inputs` attribute of your pipeline class.
|
||||
max_sequence_length (`int` defaults to `256`):
|
||||
Maximum sequence length to use with the `prompt`.
|
||||
|
||||
Examples:
|
||||
|
||||
Returns:
|
||||
[`~pipelines.mochi.MochiPipelineOutput`] or `tuple`:
|
||||
If `return_dict` is `True`, [`~pipelines.mochi.MochiPipelineOutput`] is returned, otherwise a `tuple`
|
||||
is returned where the first element is a list with the generated images.
|
||||
"""
|
||||
|
||||
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
|
||||
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
|
||||
|
||||
height = height or self.default_height
|
||||
width = width or self.default_width
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(
|
||||
prompt=prompt,
|
||||
height=height,
|
||||
width=width,
|
||||
callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
prompt_attention_mask=prompt_attention_mask,
|
||||
negative_prompt_attention_mask=negative_prompt_attention_mask,
|
||||
)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
self._interrupt = False
|
||||
|
||||
# 2. Define call parameters
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
device = self._execution_device
|
||||
|
||||
# 3. Prepare text embeddings
|
||||
(
|
||||
prompt_embeds,
|
||||
prompt_attention_mask,
|
||||
negative_prompt_embeds,
|
||||
negative_prompt_attention_mask,
|
||||
) = self.encode_prompt(
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
do_classifier_free_guidance=self.do_classifier_free_guidance,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
prompt_attention_mask=prompt_attention_mask,
|
||||
negative_prompt_attention_mask=negative_prompt_attention_mask,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
)
|
||||
if self.do_classifier_free_guidance:
|
||||
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
prompt_attention_mask = torch.cat([negative_prompt_attention_mask, prompt_attention_mask], dim=0)
|
||||
|
||||
# 4. Prepare latent variables
|
||||
num_channels_latents = self.transformer.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
num_frames,
|
||||
prompt_embeds.dtype,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
|
||||
if get_sequence_parallel_state():
|
||||
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
|
||||
latents = latents[:, :, rank, :, :, :]
|
||||
|
||||
|
||||
original_noise = copy.deepcopy(latents)
|
||||
# 5. Prepare timestep
|
||||
# from https://github.com/genmoai/models/blob/075b6e36db58f1242921deff83a1066887b9c9e1/src/mochi_preview/infer.py#L77
|
||||
threshold_noise = 0.025
|
||||
sigmas = linear_quadratic_schedule(num_inference_steps, threshold_noise)
|
||||
sigmas = np.array(sigmas)
|
||||
# check if of type FlowMatchEulerDiscreteScheduler
|
||||
if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler):
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
self.scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
timesteps,
|
||||
sigmas,
|
||||
)
|
||||
else:
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
self.scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
)
|
||||
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
# 6. Denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latent_model_input.shape[0]).to(latents.dtype)
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
# Mochi CFG + Sampling runs in FP32
|
||||
noise_pred = noise_pred.to(torch.float32)
|
||||
if self.do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents_dtype = latents.dtype
|
||||
latents = self.scheduler.step(noise_pred, t, latents.to(torch.float32), return_dict=False)[0]
|
||||
latents = latents.to(latents_dtype)
|
||||
|
||||
if latents.dtype != latents_dtype:
|
||||
if torch.backends.mps.is_available():
|
||||
# some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272
|
||||
latents = latents.to(latents_dtype)
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
for k in callback_on_step_end_tensor_inputs:
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
|
||||
if XLA_AVAILABLE:
|
||||
xm.mark_step()
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
latents = all_gather(latents, dim=2)
|
||||
#latents_shape = list(latents.shape)
|
||||
#full_shape = [latents_shape[0] * world_size] + latents_shape[1:]
|
||||
#all_latents = torch.zeros(full_shape, dtype=latents.dtype, device=latents.device)
|
||||
#torch.distributed.all_gather_into_tensor(all_latents, latents)
|
||||
#latents_list = list(all_latents.chunk(world_size, dim=0))
|
||||
#latents = torch.cat(latents_list, dim=2)
|
||||
|
||||
if output_type == "latent":
|
||||
video = latents
|
||||
else:
|
||||
# unscale/denormalize the latents
|
||||
# denormalize with the mean and std if available and not None
|
||||
has_latents_mean = hasattr(self.vae.config, "latents_mean") and self.vae.config.latents_mean is not None
|
||||
has_latents_std = hasattr(self.vae.config, "latents_std") and self.vae.config.latents_std is not None
|
||||
if has_latents_mean and has_latents_std:
|
||||
latents_mean = (
|
||||
torch.tensor(self.vae.config.latents_mean).view(1, 12, 1, 1, 1).to(latents.device, latents.dtype)
|
||||
)
|
||||
latents_std = (
|
||||
torch.tensor(self.vae.config.latents_std).view(1, 12, 1, 1, 1).to(latents.device, latents.dtype)
|
||||
)
|
||||
latents = latents * latents_std / self.vae.config.scaling_factor + latents_mean
|
||||
else:
|
||||
latents = latents / self.vae.config.scaling_factor
|
||||
|
||||
video = self.vae.decode(latents, return_dict=False)[0]
|
||||
video = self.video_processor.postprocess_video(video, output_type=output_type)
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
if return_all_states:
|
||||
# Pay extra attention here:
|
||||
# prompt_embeds with shape torch.Size([2, 256]), where prompt_embeds[1] is the prompt_embeds for the actual prompt
|
||||
# prompt_embeds[0] is for negative prompt
|
||||
return original_noise, video, latents, prompt_embeds, prompt_attention_mask
|
||||
|
||||
if not return_dict:
|
||||
return (video,)
|
||||
|
||||
return MochiPipelineOutput(frames=video)
|
||||
@@ -0,0 +1,155 @@
|
||||
import torch
|
||||
from fastvideo.model.pipeline_mochi import MochiPipeline
|
||||
import torch.distributed as dist
|
||||
|
||||
from diffusers.utils import export_to_video
|
||||
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, nccl_info
|
||||
import argparse
|
||||
import os
|
||||
from diffusers.models.transformers.transformer_mochi import MochiTransformerBlock
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
|
||||
|
||||
|
||||
import sys
|
||||
import pdb
|
||||
|
||||
class ForkedPdb(pdb.Pdb):
|
||||
"""
|
||||
PDB Subclass for debugging multi-processed code
|
||||
Suggested in: https://stackoverflow.com/questions/4716533/how-to-attach-debugger-to-a-python-subproccess
|
||||
"""
|
||||
def interaction(self, *args, **kwargs):
|
||||
_stdin = sys.stdin
|
||||
try:
|
||||
sys.stdin = open('/dev/stdin')
|
||||
pdb.Pdb.interaction(self, *args, **kwargs)
|
||||
finally:
|
||||
sys.stdin = _stdin
|
||||
def assert_all_close_list(input_list):
|
||||
for i in range(len(input_list) - 1):
|
||||
assert torch.allclose(input_list[i], input_list[i + 1]), f"input_list[{i}]: {input_list[i]}, input_list[{i+1}]: {input_list[i+1]}"
|
||||
|
||||
weight_dtype = torch.float32
|
||||
def initialize_distributed():
|
||||
local_rank = int(os.getenv('RANK', 0))
|
||||
world_size = int(os.getenv('WORLD_SIZE', 1))
|
||||
print('world_size', world_size)
|
||||
torch.cuda.set_device(local_rank)
|
||||
dist.init_process_group(backend='nccl', init_method='env://', world_size=world_size, rank=local_rank)
|
||||
return world_size
|
||||
|
||||
def main_print(content):
|
||||
if int(os.getenv('RANK', 0)) <= 0:
|
||||
print(content)
|
||||
|
||||
@torch.inference_mode
|
||||
def test_single_block(batch_size, device, seed):
|
||||
# set manual seed
|
||||
torch.manual_seed(seed)
|
||||
device = torch.cuda.current_device()
|
||||
block = MochiTransformerBlock(
|
||||
dim=768,
|
||||
num_attention_heads=12,
|
||||
attention_head_dim=64,
|
||||
pooled_projection_dim=256,
|
||||
qk_norm="rms_norm",
|
||||
activation_fn="swiglu",
|
||||
context_pre_only=False,
|
||||
).to(device)
|
||||
hidden_states = torch.randn(1, 16, 768).to(device).repeat(batch_size, 1, 1)
|
||||
encoder_hidden_states = torch.randn(1, 4, 256).to(device).repeat(batch_size, 1, 1)
|
||||
temb = torch.randn(1, 768).to(device).repeat(batch_size, 1)
|
||||
# shard hiddent_states according to world_size
|
||||
local_seq_length = hidden_states.shape[1] // nccl_info.sp_size
|
||||
hidden_states = hidden_states.narrow(1, nccl_info.global_rank * local_seq_length, local_seq_length)
|
||||
main_print(hidden_states.shape)
|
||||
hidden_states, encoder_hidden_states = block(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
temb=temb,
|
||||
)
|
||||
mean = hidden_states[0].mean()
|
||||
torch.distributed.all_reduce(mean, op=torch.distributed.ReduceOp.SUM)
|
||||
mean = mean / nccl_info.sp_size
|
||||
return mean
|
||||
|
||||
@torch.inference_mode
|
||||
def test_DiT(batch_size, transformer, seed):
|
||||
generator = torch.Generator(torch.cuda.current_device()).manual_seed(seed)
|
||||
device = torch.cuda.current_device()
|
||||
latent = torch.randn((1, 12, 8, 12, 8), device=device, dtype=weight_dtype, generator=generator).repeat(batch_size, 1, 1, 1, 1)
|
||||
prompt_embeds = torch.randn((1, 20, 4096), device=device, dtype=weight_dtype, generator=generator).repeat(batch_size, 1, 1)
|
||||
prompt_attention_mask = torch.ones((1, 20), device=device, dtype=weight_dtype).repeat(batch_size, 1)
|
||||
timestep = 0
|
||||
timestep = torch.tensor(timestep, device=device, dtype=weight_dtype).unsqueeze(0).repeat(batch_size)
|
||||
local_seq_length = latent.shape[2] // nccl_info.sp_size
|
||||
latent = latent.narrow(2, nccl_info.global_rank * local_seq_length, local_seq_length)
|
||||
# main_print(latent.shape)
|
||||
hidden_states = transformer(
|
||||
hidden_states=latent,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
timestep=timestep,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
def calculate_mean(states):
|
||||
mean = states.mean()
|
||||
torch.distributed.all_reduce(mean, op=torch.distributed.ReduceOp.SUM)
|
||||
mean = mean / int(os.getenv('WORLD_SIZE', 1))
|
||||
return mean
|
||||
mean1 = calculate_mean(hidden_states[0])
|
||||
main_print(hidden_states.shape)
|
||||
if hidden_states.shape[0] > 1:
|
||||
mean2 = calculate_mean(hidden_states[1])
|
||||
return mean1, mean2
|
||||
return mean1
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
world_size = initialize_distributed()
|
||||
device = torch.cuda.current_device()
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--seed", type=int, default=42)
|
||||
parser.add_argument("--test_single_block", action="store_true")
|
||||
args = parser.parse_args()
|
||||
seed = args.seed
|
||||
|
||||
if args.test_single_block:
|
||||
pass
|
||||
single_no_patch_bs_1 = test_single_block(1)
|
||||
single_no_patch_bs_2 = test_single_block(2)
|
||||
# check all close
|
||||
assert torch.allclose(single_no_patch_bs_1, single_no_patch_bs_2)
|
||||
single_patch_bs_1 = test_single_block(1)
|
||||
single_patch_bs_2 = test_single_block(2)
|
||||
assert torch.allclose(single_patch_bs_1, single_patch_bs_2)
|
||||
assert torch.allclose(single_no_patch_bs_1, single_patch_bs_2)
|
||||
initialize_sequence_parallel_state(world_size)
|
||||
|
||||
sp_patch_bs_1 = test_single_block(1)
|
||||
sp_patch_bs_2 = test_single_block(2)
|
||||
|
||||
assert torch.allclose(sp_patch_bs_1, sp_patch_bs_2)
|
||||
assert torch.allclose(single_no_patch_bs_1, sp_patch_bs_2)
|
||||
|
||||
else:
|
||||
transformer = MochiTransformer3DModel.from_pretrained("data/mochi/transformer", torch_dtype=weight_dtype).to(device)
|
||||
single_no_patch_bs_1 = test_DiT(1, transformer, seed)
|
||||
single_no_patch_bs_2_a, single_no_patch_bs_2_b = test_DiT(2, transformer, seed)
|
||||
|
||||
|
||||
single_patch_bs_1 = test_DiT(1, transformer, seed)
|
||||
single_patch_bs_2_a, single_patch_bs_2_b = test_DiT(2, transformer, seed)
|
||||
|
||||
initialize_sequence_parallel_state(world_size)
|
||||
|
||||
sp_patch_bs_1 = test_DiT(1, transformer, seed)
|
||||
sp_patch_bs_2_a, sp_patch_bs_2_b = test_DiT(2, transformer, seed)
|
||||
|
||||
assert_all_close_list([single_no_patch_bs_1, single_no_patch_bs_2_a, single_no_patch_bs_2_b, single_patch_bs_1, single_patch_bs_2_a, single_patch_bs_2_b, sp_patch_bs_1, sp_patch_bs_2_a, sp_patch_bs_2_b])
|
||||
|
||||
main_print(sp_patch_bs_1)
|
||||
|
||||
@@ -0,0 +1,107 @@
|
||||
import json
|
||||
|
||||
import torch.distributed as dist
|
||||
import torch
|
||||
from fastvideo.model.pipeline_mochi import MochiPipeline
|
||||
import os
|
||||
from diffusers.utils import export_to_video
|
||||
import argparse
|
||||
|
||||
def generate_video_and_latent(pipe, prompt, height, width, num_frames, num_inference_steps, guidance_scale):
|
||||
# Set the random seed for reproducibility
|
||||
generator = torch.Generator("cuda").manual_seed(12345)
|
||||
# Generate videos from the input prompt
|
||||
noise, video, latent, prompt_embed, prompt_attention_mask = pipe(
|
||||
prompt=prompt,
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
generator=generator,
|
||||
num_inference_steps=num_inference_steps,
|
||||
guidance_scale=guidance_scale,
|
||||
return_all_states=True,
|
||||
)
|
||||
# prompt_embed has negative prompt at index 0
|
||||
return noise[0], video[0], latent[0], prompt_embed[1], prompt_attention_mask[1]
|
||||
|
||||
# return dummy tensor to debug first
|
||||
# return torch.zeros(1, 3, 480, 848), torch.zeros(1, 256, 16, 16)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument("--height", type=int, default=480)
|
||||
parser.add_argument("--width", type=int, default=848)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=64)
|
||||
parser.add_argument("--guidance_scale", type=float, default=4.5)
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--prompt_path", type=str, default="data/dummyVid/videos2caption.json")
|
||||
parser.add_argument("--dataset_output_dir", type=str, default="data/dummySynthetic")
|
||||
args = parser.parse_args()
|
||||
|
||||
local_rank = int(os.getenv('RANK', 0))
|
||||
world_size = int(os.getenv('WORLD_SIZE', 1))
|
||||
print('world_size', world_size, 'local rank', local_rank)
|
||||
torch.cuda.set_device(local_rank)
|
||||
dist.init_process_group(backend='nccl', init_method='env://', world_size=world_size, rank=local_rank)
|
||||
|
||||
|
||||
if not isinstance(args.prompt_path, list):
|
||||
args.prompt_path = [args.prompt_path]
|
||||
if len(args.prompt_path) == 1 and args.prompt_path[0].endswith('txt'):
|
||||
text_prompt = open(args.prompt_path[0], 'r').readlines()
|
||||
text_prompt = [i.strip() for i in text_prompt]
|
||||
|
||||
|
||||
pipe = MochiPipeline.from_pretrained(args.model_path, torch_dtype=torch.bfloat16)
|
||||
pipe.enable_vae_tiling()
|
||||
pipe.enable_model_cpu_offload(gpu_id=local_rank)
|
||||
# make dir if not exist
|
||||
|
||||
os.makedirs(args.dataset_output_dir, exist_ok=True)
|
||||
os.makedirs(os.path.join(args.dataset_output_dir, "noise"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.dataset_output_dir, "video"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.dataset_output_dir, "latent"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.dataset_output_dir, "prompt_embed"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.dataset_output_dir, "prompt_attention_mask"), exist_ok=True)
|
||||
data = []
|
||||
for i, prompt in enumerate(text_prompt):
|
||||
if i % world_size != local_rank:
|
||||
continue
|
||||
noise, video, latent, prompt_embed, prompt_attention_mask = generate_video_and_latent(pipe, prompt, args.height, args.width, args.num_frames, args.num_inference_steps, args.guidance_scale)
|
||||
# save latent
|
||||
video_name = str(i)
|
||||
noise_path = os.path.join(args.dataset_output_dir, "noise", video_name + ".pt")
|
||||
latent_path = os.path.join(args.dataset_output_dir, "latent", video_name + ".pt")
|
||||
prompt_embed_path = os.path.join(args.dataset_output_dir, "prompt_embed", video_name + ".pt")
|
||||
video_path = os.path.join(args.dataset_output_dir, "video", video_name + ".mp4")
|
||||
prompt_attention_mask_path = os.path.join(args.dataset_output_dir, "prompt_attention_mask", video_name + ".pt")
|
||||
# save latent
|
||||
torch.save(noise, noise_path)
|
||||
torch.save(latent, latent_path)
|
||||
torch.save(prompt_embed, prompt_embed_path)
|
||||
torch.save(prompt_attention_mask, prompt_attention_mask_path)
|
||||
export_to_video(video, video_path, fps=30)
|
||||
item = {}
|
||||
|
||||
item["cap"] = prompt
|
||||
item["video"] = video_name + ".mp4"
|
||||
item["noise"] = video_name + ".pt"
|
||||
item["latent_path"] = video_name + ".pt"
|
||||
item["prompt_embed_path"] = video_name + ".pt"
|
||||
item["prompt_attention_mask"] = video_name + ".pt"
|
||||
data.append(item)
|
||||
dist.barrier()
|
||||
local_data = data
|
||||
gathered_data = [None] * world_size
|
||||
dist.all_gather_object(gathered_data, local_data)
|
||||
|
||||
# save json
|
||||
if local_rank == 0:
|
||||
all_data = [item for sublist in gathered_data for item in sublist]
|
||||
with open(os.path.join(args.dataset_output_dir, "videos2caption.json"), 'w') as f:
|
||||
json.dump(all_data, f, indent=4)
|
||||
|
||||
|
||||
@@ -0,0 +1,207 @@
|
||||
import torch
|
||||
from fastvideo.model.pipeline_mochi import MochiPipeline
|
||||
import torch.distributed as dist
|
||||
|
||||
from diffusers.utils import export_to_video
|
||||
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, nccl_info
|
||||
import argparse
|
||||
import os
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
import json
|
||||
from typing import Optional
|
||||
from safetensors.torch import save_file, load_file
|
||||
from peft import set_peft_model_state_dict, inject_adapter_in_model, load_peft_weights
|
||||
from peft import LoraConfig
|
||||
import sys
|
||||
import pdb
|
||||
import copy
|
||||
from typing import Dict
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from fastvideo.distill.solver import PCMFMScheduler
|
||||
def initialize_distributed():
|
||||
local_rank = int(os.getenv('RANK', 0))
|
||||
world_size = int(os.getenv('WORLD_SIZE', 1))
|
||||
print('world_size', world_size)
|
||||
torch.cuda.set_device(local_rank)
|
||||
dist.init_process_group(backend='nccl', init_method='env://', world_size=world_size, rank=local_rank)
|
||||
initialize_sequence_parallel_state(world_size)
|
||||
|
||||
def merge_lora_weights(
|
||||
base_model: torch.nn.Module,
|
||||
lora_weights: Dict[str, torch.Tensor],
|
||||
lora_config: LoraConfig,
|
||||
num_layers: Optional[int] = None
|
||||
) -> torch.nn.Module:
|
||||
merged_model = copy.deepcopy(base_model)
|
||||
if num_layers is None:
|
||||
num_layers = len(merged_model.transformer_blocks)
|
||||
scaling = lora_config.lora_alpha / lora_config.r
|
||||
|
||||
def merge_component(
|
||||
base_weight: torch.Tensor,
|
||||
lora_a: torch.Tensor,
|
||||
lora_b: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
device = base_weight.device
|
||||
lora_a = lora_a.to(device)
|
||||
lora_b = lora_b.to(device)
|
||||
lora_contribution = (lora_b @ lora_a) * scaling
|
||||
if lora_contribution.shape != base_weight.shape:
|
||||
raise ValueError(
|
||||
f"Shape mismatch: base={base_weight.shape}, "
|
||||
f"lora={lora_contribution.shape}"
|
||||
)
|
||||
return base_weight + lora_contribution
|
||||
|
||||
for layer_idx in range(num_layers):
|
||||
transformer_layer = merged_model.transformer_blocks[layer_idx].attn1
|
||||
for target_module in lora_config.target_modules:
|
||||
if target_module == "to_out.0":
|
||||
base_weight = transformer_layer.to_out[0].weight
|
||||
lora_a_key = f"transformer_blocks.{layer_idx}.attn1.to_out.0.lora_A.default.weight"
|
||||
lora_b_key = f"transformer_blocks.{layer_idx}.attn1.to_out.0.lora_B.default.weight"
|
||||
else:
|
||||
base_weight = getattr(transformer_layer, target_module).weight
|
||||
lora_a_key = f"transformer_blocks.{layer_idx}.attn1.{target_module}.lora_A.default.weight"
|
||||
lora_b_key = f"transformer_blocks.{layer_idx}.attn1.{target_module}.lora_B.default.weight"
|
||||
lora_a = lora_weights[lora_a_key]
|
||||
lora_b = lora_weights[lora_b_key]
|
||||
merged_weight = merge_component(base_weight, lora_a, lora_b)
|
||||
if target_module == "to_out.0":
|
||||
transformer_layer.to_out[0].weight.data.copy_(merged_weight)
|
||||
else:
|
||||
getattr(transformer_layer, target_module).weight.data.copy_(merged_weight)
|
||||
merged_model.transformer_blocks[layer_idx].attn1 = transformer_layer
|
||||
return merged_model
|
||||
|
||||
def load_lora_checkpoint(
|
||||
transformer: MochiTransformer3DModel,
|
||||
optimizer,
|
||||
lora_checkpoint_dir: str
|
||||
):
|
||||
config_path = os.path.join(lora_checkpoint_dir, "lora_config.json")
|
||||
with open(config_path, 'r') as f:
|
||||
lora_config_dict = json.load(f)
|
||||
|
||||
for key, value in lora_config['lora_params'].items():
|
||||
setattr(transformer.config, f"lora_{key}", value)
|
||||
|
||||
weight_path = os.path.join(lora_checkpoint_dir, "lora_weights.safetensors")
|
||||
lora_state_dict = load_file(weight_path)
|
||||
|
||||
lora_config = LoraConfig(
|
||||
r=lora_config_dict['lora_params']['lora_rank'],
|
||||
lora_alpha=lora_config_dict['lora_params']['lora_alpha'],
|
||||
target_modules=lora_config_dict['lora_params']['target_modules']
|
||||
)
|
||||
|
||||
transformer = merge_lora_weights(transformer, lora_state_dict, lora_config)
|
||||
step = lora_state_dict['step']
|
||||
print(f"--> Successfully loaded LoRA checkpoint from step {step}")
|
||||
return transformer
|
||||
|
||||
def main(args):
|
||||
initialize_distributed()
|
||||
print(nccl_info.sp_size)
|
||||
device = torch.cuda.current_device()
|
||||
generator = torch.Generator(device).manual_seed(args.seed)
|
||||
weight_dtype = torch.bfloat16
|
||||
if args.scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
else:
|
||||
linear_quadratic = True if "linear_quadratic" in args.scheduler_type else False
|
||||
scheduler = PCMFMScheduler(1000, args.shift, args.num_euler_timesteps, linear_quadratic,args.linear_threshold, args.linear_range)
|
||||
if args.transformer_path is not None:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
|
||||
else:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder = 'transformer/')
|
||||
if args.lora_checkpoint_dir is not None:
|
||||
# Load and merge LoRA weights
|
||||
transformer = load_lora_checkpoint(
|
||||
transformer=transformer,
|
||||
optimizer=None, # No optimizer needed for inference
|
||||
output_dir=args.lora_checkpoint_dir
|
||||
)
|
||||
print(f"Loaded and merged LoRA weights from {args.lora_checkpoint_dir}")
|
||||
pipe = MochiPipeline.from_pretrained(args.model_path, transformer = transformer,scheduler=scheduler)
|
||||
|
||||
pipe.enable_vae_tiling()
|
||||
# pipe.to(device)
|
||||
|
||||
pipe.enable_model_cpu_offload(device)
|
||||
# Generate videos from the input prompt
|
||||
|
||||
if args.prompt_embed_path is not None:
|
||||
prompt_embeds = torch.load(args.prompt_embed_path, map_location="cpu", weights_only=True).to(device).unsqueeze(0)
|
||||
encoder_attention_mask = torch.load(args.encoder_attention_mask_path, map_location="cpu", weights_only=True).to(device).unsqueeze(0)
|
||||
prompts = None
|
||||
elif args.prompt_path is not None:
|
||||
prompts = [line.strip() for line in open(args.prompt_path, "r")]
|
||||
prompt_embeds = None
|
||||
encoder_attention_mask = None
|
||||
else:
|
||||
prompts = args.prompts
|
||||
prompt_embeds = None
|
||||
encoder_attention_mask = None
|
||||
|
||||
if prompts is not None:
|
||||
videos = []
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
for prompt in prompts:
|
||||
video = pipe(
|
||||
prompt=[prompt],
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
generator=generator,
|
||||
).frames
|
||||
videos.append(video[0])
|
||||
else:
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
videos = pipe(
|
||||
prompt_embeds=prompt_embeds,
|
||||
prompt_attention_mask=encoder_attention_mask,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
generator=generator,
|
||||
).frames
|
||||
|
||||
if nccl_info.global_rank <= 0:
|
||||
if prompts is not None:
|
||||
# mkdir
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
for video, prompt in zip(videos, prompts):
|
||||
suffix = prompt.split(".")[0]
|
||||
export_to_video(video, os.path.join(args.output_path, f"{suffix}.mp4"), fps=30)
|
||||
else:
|
||||
export_to_video(videos[0], args.output_path + ".mp4", fps=30)
|
||||
|
||||
if __name__ == "__main__":
|
||||
# arg parse
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--prompts", nargs='+', default=[])
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument("--height", type=int, default=480)
|
||||
parser.add_argument("--width", type=int, default=848)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=64)
|
||||
parser.add_argument("--guidance_scale", type=float, default=4.5)
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--seed", type=int, default=42)
|
||||
parser.add_argument("--output_path", type=str, default="./outputs.mp4")
|
||||
parser.add_argument("--transformer_path", type=str, default=None)
|
||||
parser.add_argument("--prompt_embed_path", type=str, default=None)
|
||||
parser.add_argument("--prompt_path", type=str, default=None)
|
||||
parser.add_argument("--scheduler_type", type=str, default="euler")
|
||||
parser.add_argument("--encoder_attention_mask_path", type=str, default=None)
|
||||
parser.add_argument('--lora_checkpoint_dir', type=str, default=None, help='Path to the directory containing LoRA checkpoints')
|
||||
parser.add_argument("--shift", type=float, default=8.0)
|
||||
parser.add_argument("--num_euler_timesteps", type=int, default=100)
|
||||
parser.add_argument("--linear_threshold", type=float, default=0.025)
|
||||
parser.add_argument("--linear_range", type=float, default=0.5)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,50 @@
|
||||
import torch
|
||||
from fastvideo.model.pipeline_mochi import MochiPipeline
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
from diffusers.utils import export_to_video, load_image, load_video
|
||||
import argparse
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
def main(args):
|
||||
# Set the random seed for reproducibility
|
||||
generator = torch.Generator("cuda").manual_seed(args.seed)
|
||||
# do not invert
|
||||
scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
if args.transformer_path is not None:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
|
||||
else:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder = 'transformer/')
|
||||
pipe = MochiPipeline.from_pretrained(args.model_path, transformer = transformer, scheduler = scheduler)
|
||||
pipe.enable_vae_tiling()
|
||||
# pipe.to("cuda:1")
|
||||
pipe.enable_model_cpu_offload()
|
||||
|
||||
# Generate videos from the input prompt
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
videos = pipe(
|
||||
prompt=args.prompts,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
generator=generator,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
).frames
|
||||
|
||||
for prompt,video in zip(args.prompts, videos):
|
||||
export_to_video(video, args.output_path + f"_{prompt}.mp4", fps=30)
|
||||
|
||||
if __name__ == "__main__":
|
||||
# arg parse
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--prompts", nargs='+', default=[])
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument("--height", type=int, default=480)
|
||||
parser.add_argument("--width", type=int, default=848)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=64)
|
||||
parser.add_argument("--guidance_scale", type=float, default=4.5)
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--seed", type=int, default=12345)
|
||||
parser.add_argument("--transformer_path", type=str, default=None)
|
||||
parser.add_argument("--output_path", type=str, default="./outputs.mp4")
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,482 @@
|
||||
import argparse
|
||||
from email.policy import strict
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, \
|
||||
destroy_sequence_parallel_group, get_sequence_parallel_state, nccl_info
|
||||
from fastvideo.utils.communications import sp_parallel_dataloader_wrapper, broadcast
|
||||
from fastvideo.model.mochi_latents_utils import normalize_mochi_dit_input
|
||||
from fastvideo.utils.validation import log_validation
|
||||
import time
|
||||
from torch.utils.data import DataLoader
|
||||
import torch
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
StateDictType,
|
||||
FullStateDictConfig,
|
||||
)
|
||||
import json
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from fastvideo.utils.dataset_utils import LengthGroupedSampler
|
||||
import wandb
|
||||
from accelerate.utils import set_seed
|
||||
from tqdm.auto import tqdm
|
||||
from fastvideo.fsdp_util import get_dit_fsdp_kwargs, apply_fsdp_checkpointing
|
||||
import diffusers
|
||||
from diffusers import (
|
||||
FlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
from diffusers.optimization import get_scheduler
|
||||
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
|
||||
from diffusers.utils import check_min_version
|
||||
from fastvideo.dataset.latent_datasets import LatentDataset, latent_collate_function
|
||||
import torch.distributed as dist
|
||||
from safetensors.torch import save_file, load_file
|
||||
from peft import LoraConfig, inject_adapter_in_model
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
)
|
||||
from fastvideo.utils.checkpoint import save_checkpoint, save_lora_checkpoint, resume_lora_training
|
||||
from fastvideo.utils.logging import main_print
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
check_min_version("0.31.0")
|
||||
import time
|
||||
from collections import deque
|
||||
|
||||
|
||||
|
||||
|
||||
def compute_density_for_timestep_sampling(
|
||||
weighting_scheme: str, batch_size: int, generator, logit_mean: float = None, logit_std: float = None, mode_scale: float = None
|
||||
):
|
||||
"""
|
||||
Compute the density for sampling the timesteps when doing SD3 training.
|
||||
|
||||
Courtesy: This was contributed by Rafie Walker in https://github.com/huggingface/diffusers/pull/8528.
|
||||
|
||||
SD3 paper reference: https://arxiv.org/abs/2403.03206v1.
|
||||
"""
|
||||
if weighting_scheme == "logit_normal":
|
||||
# See 3.1 in the SD3 paper ($rf/lognorm(0.00,1.00)$).
|
||||
u = torch.normal(mean=logit_mean, std=logit_std, size=(batch_size,), device="cpu", generator=generator)
|
||||
u = torch.nn.functional.sigmoid(u)
|
||||
elif weighting_scheme == "mode":
|
||||
u = torch.rand(size=(batch_size,), device="cpu", generator=generator)
|
||||
u = 1 - u - mode_scale * (torch.cos(math.pi * u / 2) ** 2 - 1 + u)
|
||||
else:
|
||||
u = torch.rand(size=(batch_size,), device="cpu", generator=generator)
|
||||
return u
|
||||
|
||||
def get_sigmas(noise_scheduler, device, timesteps, n_dim=4, dtype=torch.float32):
|
||||
sigmas = noise_scheduler.sigmas.to(device=device, dtype=dtype)
|
||||
schedule_timesteps = noise_scheduler.timesteps.to(device)
|
||||
timesteps = timesteps.to(device)
|
||||
step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps]
|
||||
|
||||
sigma = sigmas[step_indices].flatten()
|
||||
while len(sigma.shape) < n_dim:
|
||||
sigma = sigma.unsqueeze(-1)
|
||||
return sigma
|
||||
|
||||
|
||||
def train_one_step_mochi(transformer, optimizer, lr_scheduler,loader, noise_scheduler, noise_random_generator, gradient_accumulation_steps, sp_size, precondition_outputs, max_grad_norm, weighting_scheme, logit_mean, logit_std, mode_scale):
|
||||
total_loss = 0.0
|
||||
optimizer.zero_grad()
|
||||
for _ in range(gradient_accumulation_steps):
|
||||
latents, encoder_hidden_states, latents_attention_mask, encoder_attention_mask = next(loader)
|
||||
latents = normalize_mochi_dit_input(latents)
|
||||
|
||||
batch_size = latents.shape[0]
|
||||
noise = torch.randn_like(latents)
|
||||
u = compute_density_for_timestep_sampling(
|
||||
weighting_scheme=weighting_scheme,
|
||||
batch_size=batch_size,
|
||||
generator=noise_random_generator,
|
||||
logit_mean=logit_mean,
|
||||
logit_std=logit_std,
|
||||
mode_scale=mode_scale,
|
||||
)
|
||||
indices = (u * noise_scheduler.config.num_train_timesteps).long()
|
||||
timesteps = noise_scheduler.timesteps[indices].to(device=latents.device)
|
||||
if sp_size > 1:
|
||||
# Make sure that the timesteps are the same across all sp processes.
|
||||
broadcast(timesteps)
|
||||
|
||||
sigmas = get_sigmas(noise_scheduler, latents.device, timesteps, n_dim=latents.ndim, dtype=latents.dtype)
|
||||
noisy_model_input = (1.0 - sigmas) * latents + sigmas * noise
|
||||
with torch.autocast("cuda", torch.bfloat16):
|
||||
model_pred = transformer(
|
||||
noisy_model_input,
|
||||
encoder_hidden_states,
|
||||
timesteps,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict= False
|
||||
)[0]
|
||||
|
||||
if precondition_outputs:
|
||||
model_pred = noisy_model_input - model_pred * sigmas
|
||||
if precondition_outputs:
|
||||
target = latents
|
||||
else:
|
||||
target = noise - latents
|
||||
|
||||
loss = torch.mean((model_pred.float() - target.float()) ** 2) / gradient_accumulation_steps
|
||||
|
||||
loss.backward()
|
||||
|
||||
avg_loss = loss.detach().clone()
|
||||
dist.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
|
||||
total_loss += avg_loss.item()
|
||||
|
||||
|
||||
grad_norm = transformer.clip_grad_norm_(max_grad_norm)
|
||||
optimizer.step()
|
||||
lr_scheduler.step()
|
||||
return total_loss, grad_norm.item()
|
||||
|
||||
def get_lora_model(transformer, lora_config):
|
||||
transformer.requires_grad_(False)
|
||||
transformer = inject_adapter_in_model(lora_config, transformer)
|
||||
return transformer
|
||||
|
||||
|
||||
|
||||
def main(args):
|
||||
# use LayerNorm, GeLu, SiLu always as fp32 mode
|
||||
# TODO:
|
||||
if args.enable_stable_fp32:
|
||||
raise NotImplementedError("enable_stable_fp32 is not supported now.")
|
||||
torch.backends.cuda.matmul.allow_tf32 = True
|
||||
|
||||
local_rank = int(os.environ['LOCAL_RANK'])
|
||||
rank = int(os.environ['RANK'])
|
||||
world_size = int(os.environ['WORLD_SIZE'])
|
||||
dist.init_process_group("nccl")
|
||||
torch.cuda.set_device(local_rank)
|
||||
device = torch.cuda.current_device()
|
||||
initialize_sequence_parallel_state(args.sp_size)
|
||||
|
||||
# If passed along, set the training seed now. On GPU...
|
||||
if args.seed is not None:
|
||||
# TODO: t within the same seq parallel group should be the same. Noise should be different.
|
||||
set_seed(args.seed + rank)
|
||||
# We use different seeds for the noise generation in each process to ensure that the noise is different in a batch.
|
||||
noise_random_generator = None
|
||||
|
||||
# Handle the repository creation
|
||||
if rank <=0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weigths to half-precision
|
||||
# as these weights are only used for inference, keeping weights in full precision is not required.
|
||||
|
||||
# Create model:
|
||||
|
||||
main_print(f"--> loading model from {args.pretrained_model_name_or_path}")
|
||||
# keep the master weight to float32
|
||||
transformer = MochiTransformer3DModel.from_pretrained(
|
||||
args.pretrained_model_name_or_path,
|
||||
subfolder="transformer",
|
||||
torch_dtype = torch.float32,
|
||||
#torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
|
||||
if args.use_lora:
|
||||
lora_config = LoraConfig(
|
||||
r=args.lora_rank,
|
||||
lora_alpha=args.lora_alpha,
|
||||
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
|
||||
init_lora_weights=True,
|
||||
)
|
||||
transformer = get_lora_model(transformer, lora_config)
|
||||
|
||||
main_print(f" Total training parameters = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e6} M")
|
||||
main_print(f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}")
|
||||
fsdp_kwargs = get_dit_fsdp_kwargs(args.fsdp_sharding_startegy, args.use_lora, args.use_cpu_offload)
|
||||
|
||||
|
||||
if args.use_lora:
|
||||
transformer.config.lora_rank = args.lora_rank
|
||||
transformer.config.lora_alpha = args.lora_alpha
|
||||
transformer.config.lora_target_modules = ["to_k", "to_q", "to_v", "to_out.0"]
|
||||
transformer._no_split_modules = ["MochiTransformerBlock"]
|
||||
fsdp_kwargs['auto_wrap_policy'] = fsdp_kwargs['auto_wrap_policy'](transformer)
|
||||
|
||||
|
||||
transformer = FSDP(
|
||||
transformer,
|
||||
**fsdp_kwargs,
|
||||
)
|
||||
main_print(f"--> model loaded")
|
||||
|
||||
if args.gradient_checkpointing:
|
||||
apply_fsdp_checkpointing(transformer, args.selective_checkpointing)
|
||||
|
||||
# Set model as trainable.
|
||||
transformer.train()
|
||||
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
|
||||
params_to_optimize = transformer.parameters()
|
||||
params_to_optimize = list(filter(lambda p: p.requires_grad, params_to_optimize))
|
||||
|
||||
optimizer = torch.optim.AdamW(
|
||||
params_to_optimize,
|
||||
lr=args.learning_rate,
|
||||
betas=(0.9,0.999),
|
||||
weight_decay=args.weight_decay,
|
||||
eps=1e-8,
|
||||
)
|
||||
|
||||
init_steps = 0
|
||||
if args.resume_from_lora_checkpoint:
|
||||
transformer, optimizer, init_steps = resume_lora_training(
|
||||
transformer, args.resume_from_lora_checkpoint, optimizer
|
||||
)
|
||||
main_print(f"optimizer: {optimizer}")
|
||||
|
||||
#todo add lr scheduler
|
||||
lr_scheduler = get_scheduler(
|
||||
args.lr_scheduler,
|
||||
optimizer=optimizer,
|
||||
num_warmup_steps=args.lr_warmup_steps * world_size,
|
||||
num_training_steps=args.max_train_steps * world_size,
|
||||
num_cycles=args.lr_num_cycles,
|
||||
power=args.lr_power,
|
||||
last_epoch=init_steps - 1,
|
||||
)
|
||||
|
||||
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t, args.cfg)
|
||||
sampler = LengthGroupedSampler(
|
||||
args.train_batch_size,
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
lengths=train_dataset.lengths,
|
||||
group_frame=args.group_frame,
|
||||
group_resolution=args.group_resolution,
|
||||
) if (args.group_frame or args.group_resolution) else DistributedSampler(train_dataset, rank=rank, num_replicas=world_size, shuffle=False)
|
||||
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
collate_fn=latent_collate_function,
|
||||
pin_memory=True,
|
||||
batch_size=args.train_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
drop_last=True,
|
||||
)
|
||||
|
||||
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
|
||||
|
||||
if rank <= 0:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
# Train!
|
||||
total_batch_size = args.train_batch_size * world_size * args.gradient_accumulation_steps / args.sp_size * args.train_sp_batch_size
|
||||
main_print("***** Running training *****")
|
||||
main_print(f" Num examples = {len(train_dataset)}")
|
||||
main_print(f" Dataloader size = {len(train_dataloader)}")
|
||||
main_print(f" Num Epochs = {args.num_train_epochs}")
|
||||
main_print(f" Resume training from step {init_steps}")
|
||||
main_print(f" Instantaneous batch size per device = {args.train_batch_size}")
|
||||
main_print(f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}")
|
||||
main_print(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
|
||||
main_print(f" Total optimization steps = {args.max_train_steps}")
|
||||
main_print(f" Total training parameters per FSDP shard = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e9} B")
|
||||
# print dtype
|
||||
main_print(f" Master weight dtype: {transformer.parameters().__next__().dtype}")
|
||||
|
||||
|
||||
# Potentially load in the weights and states from a previous save
|
||||
if args.resume_from_checkpoint:
|
||||
assert NotImplementedError("resume_from_checkpoint is not supported now.")
|
||||
# TODO
|
||||
|
||||
|
||||
progress_bar = tqdm(
|
||||
range(0, args.max_train_steps),
|
||||
initial=init_steps,
|
||||
desc="Steps",
|
||||
# Only show the progress bar once on each machine.
|
||||
disable= local_rank > 0,
|
||||
)
|
||||
|
||||
loader = sp_parallel_dataloader_wrapper(train_dataloader, device, args.train_batch_size, args.sp_size, args.train_sp_batch_size)
|
||||
|
||||
step_times = deque(maxlen=100)
|
||||
|
||||
#todo future
|
||||
for i in range(init_steps):
|
||||
next(loader)
|
||||
for step in range(init_steps + 1, args.max_train_steps+1):
|
||||
start_time = time.time()
|
||||
loss, grad_norm= train_one_step_mochi(transformer, optimizer, lr_scheduler, loader, noise_scheduler, noise_random_generator, args.gradient_accumulation_steps, args.sp_size, args.precondition_outputs, args.max_grad_norm, args.weighting_scheme, args.logit_mean, args.logit_std, args.mode_scale)
|
||||
|
||||
step_time = time.time() - start_time
|
||||
step_times.append(step_time)
|
||||
avg_step_time = sum(step_times) / len(step_times)
|
||||
|
||||
progress_bar.set_postfix({
|
||||
"loss": f"{loss:.4f}",
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
"grad_norm": grad_norm
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if rank <= 0:
|
||||
wandb.log({
|
||||
"train_loss": loss,
|
||||
"learning_rate": lr_scheduler.get_last_lr()[0],
|
||||
"step_time": step_time,
|
||||
"avg_step_time": avg_step_time,
|
||||
"grad_norm": grad_norm
|
||||
}, step=step)
|
||||
if step % args.checkpointing_steps == 0:
|
||||
if args.use_lora:
|
||||
# Save LoRA weights
|
||||
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, step)
|
||||
else:
|
||||
# Your existing checkpoint saving code
|
||||
save_checkpoint(transformer, optimizer, rank, args.output_dir, step)
|
||||
dist.barrier()
|
||||
if args.log_validation and step % args.validation_steps == 0:
|
||||
log_validation(args, transformer, device,
|
||||
torch.bfloat16, step)
|
||||
|
||||
if args.use_lora:
|
||||
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, args.max_train_steps)
|
||||
else:
|
||||
save_checkpoint(transformer, optimizer, rank, args.output_dir, args.max_train_steps)
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
destroy_sequence_parallel_group()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--data_json_path", type=str, required=True)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument("--dataloader_num_workers", type=int, default=10, help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.")
|
||||
parser.add_argument("--train_batch_size", type=int, default=16, help="Batch size (per device) for the training dataloader.")
|
||||
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--pretrained_model_name_or_path", type=str)
|
||||
parser.add_argument("--cache_dir", type=str, default='./cache_dir')
|
||||
parser.add_argument('--enable_stable_fp32', action='store_true') # TODO
|
||||
|
||||
# diffusion setting
|
||||
parser.add_argument("--ema_decay", type=float, default=0.999)
|
||||
parser.add_argument("--ema_start_step", type=int, default=0)
|
||||
parser.add_argument('--cfg', type=float, default=0.1)
|
||||
parser.add_argument("--precondition_outputs", action="store_true", help="Whether to precondition the outputs of the model.")
|
||||
|
||||
# validation & logs
|
||||
parser.add_argument("--validation_prompt_dir", type=str)
|
||||
parser.add_argument("--uncond_prompt_dir", type=str)
|
||||
parser.add_argument("--validation_sampling_steps", type=int, default=64)
|
||||
parser.add_argument('--validation_guidance_scale', type=float, default=4.5)
|
||||
parser.add_argument('--validation_steps', type=float, default=4.5)
|
||||
parser.add_argument("--log_validation", action="store_true")
|
||||
parser.add_argument("--tracker_project_name", type=str, default=None)
|
||||
parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.")
|
||||
parser.add_argument("--output_dir", type=str, default=None, help="The output directory where the model predictions and checkpoints will be written.")
|
||||
parser.add_argument("--checkpoints_total_limit", type=int, default=None, help=("Max number of checkpoints to store."))
|
||||
parser.add_argument("--checkpointing_steps", type=int, default=500,
|
||||
help=(
|
||||
"Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
|
||||
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
|
||||
" training using `--resume_from_checkpoint`."
|
||||
),
|
||||
)
|
||||
parser.add_argument("--resume_from_checkpoint", type=str, default=None,
|
||||
help=(
|
||||
"Whether training should be resumed from a previous checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
|
||||
),
|
||||
)
|
||||
parser.add_argument("--resume_from_lora_checkpoint", type=str, default=None,
|
||||
help=(
|
||||
"Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
|
||||
),
|
||||
)
|
||||
parser.add_argument("--logging_dir", type=str, default="logs",
|
||||
help=(
|
||||
"[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
|
||||
),
|
||||
)
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
parser.add_argument("--num_train_epochs", type=int, default=100)
|
||||
parser.add_argument("--max_train_steps", type=int, default=None, help="Total number of training steps to perform. If provided, overrides num_train_epochs.")
|
||||
parser.add_argument("--gradient_accumulation_steps", type=int, default=1, help="Number of updates steps to accumulate before performing a backward/update pass.")
|
||||
parser.add_argument("--learning_rate", type=float, default=1e-4, help="Initial learning rate (after the potential warmup period) to use.")
|
||||
parser.add_argument("--scale_lr", action="store_true", default=False, help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.")
|
||||
parser.add_argument("--lr_warmup_steps", type=int, default=10, help="Number of steps for the warmup in the lr scheduler.")
|
||||
parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.")
|
||||
parser.add_argument("--gradient_checkpointing", action="store_true", help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.")
|
||||
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
|
||||
parser.add_argument("--allow_tf32", action="store_true",
|
||||
help=(
|
||||
"Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
|
||||
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
|
||||
),
|
||||
)
|
||||
parser.add_argument("--mixed_precision", type=str, default=None, choices=["no", "fp16", "bf16"],
|
||||
help=(
|
||||
"Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
|
||||
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
|
||||
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
|
||||
),
|
||||
)
|
||||
parser.add_argument("--use_cpu_offload", action="store_true", help="Whether to use CPU offload for param & gradient & optimizer states.")
|
||||
|
||||
parser.add_argument("--sp_size", type=int, default=1, help="For sequence parallel")
|
||||
parser.add_argument("--train_sp_batch_size", type=int, default=1, help="Batch size for sequence parallel training")
|
||||
|
||||
parser.add_argument("--use_lora", action="store_true", default=False, help="Whether to use LoRA for finetuning.")
|
||||
parser.add_argument("--lora_alpha", type=int, default=256, help="Alpha parameter for LoRA.")
|
||||
parser.add_argument("--lora_rank", type=int, default=128, help="LoRA rank parameter. ")
|
||||
parser.add_argument("--fsdp_sharding_startegy", default="full")
|
||||
|
||||
parser.add_argument(
|
||||
"--weighting_scheme",
|
||||
type=str,
|
||||
default="uniform",
|
||||
choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "uniform"],
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mode_scale",
|
||||
type=float,
|
||||
default=1.29,
|
||||
help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.",
|
||||
)
|
||||
# lr_scheduler
|
||||
parser.add_argument("--lr_scheduler", type=str, default="constant",
|
||||
help=(
|
||||
'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
|
||||
' "constant", "constant_with_warmup"]'
|
||||
),
|
||||
)
|
||||
parser.add_argument("--lr_num_cycles", type=int, default=1, help="Number of cycles in the learning rate scheduler.")
|
||||
parser.add_argument("--lr_power", type=float, default=1.0, help="Power factor of the polynomial scheduler.",)
|
||||
parser.add_argument("--weight_decay", type=float, default=0.01, help="Weight decay to apply.")
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,244 @@
|
||||
# import
|
||||
import os
|
||||
import json
|
||||
import torch
|
||||
from fastvideo.utils.logging import main_print
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP, StateDictType, FullStateDictConfig
|
||||
from safetensors.torch import save_file, load_file
|
||||
import torch.distributed.checkpoint as dist_cp
|
||||
from torch.distributed.checkpoint.default_planner import DefaultSavePlanner, DefaultLoadPlanner
|
||||
from torch.distributed.checkpoint.optimizer import load_sharded_optimizer_state_dict
|
||||
from torch.distributed.fsdp import FullOptimStateDictConfig
|
||||
def save_checkpoint(model, optimizer, rank, output_dir, step, discriminator=False):
|
||||
|
||||
with FSDP.state_dict_type(
|
||||
model, StateDictType.FULL_STATE_DICT, FullStateDictConfig(offload_to_cpu=True, rank0_only=True), FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
):
|
||||
cpu_state = model.state_dict()
|
||||
optim_state = FSDP.optim_state_dict(
|
||||
model,
|
||||
optimizer,
|
||||
)
|
||||
|
||||
#todo move to get_state_dict
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
if rank <= 0 and not discriminator:
|
||||
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
|
||||
save_file(cpu_state, weight_path)
|
||||
config_dict = dict(model.config)
|
||||
config_path = os.path.join(save_dir, "config.json")
|
||||
# save dict as json
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
optimizer_path = os.path.join(save_dir, "optimizer.pt")
|
||||
torch.save(optim_state, optimizer_path)
|
||||
else:
|
||||
weight_path = os.path.join(save_dir, "discriminator_pytorch_model.safetensors")
|
||||
save_file(cpu_state, weight_path)
|
||||
optimizer_path = os.path.join(save_dir, "discriminator_optimizer.pt")
|
||||
torch.save(optim_state, optimizer_path)
|
||||
|
||||
|
||||
|
||||
def save_checkpoint_generator_discriminator(model, optimizer, discriminator, discriminator_optimizer, rank, output_dir, step,):
|
||||
with FSDP.state_dict_type(
|
||||
model, StateDictType.FULL_STATE_DICT, FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
):
|
||||
cpu_state = model.state_dict()
|
||||
|
||||
#todo move to get_state_dict
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
hf_weight_dir = os.path.join(save_dir, "hf_weights")
|
||||
os.makedirs(hf_weight_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
if rank <= 0:
|
||||
config_dict = dict(model.config)
|
||||
config_path = os.path.join(hf_weight_dir, "config.json")
|
||||
# save dict as json
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
weight_path = os.path.join(hf_weight_dir, "diffusion_pytorch_model.safetensors")
|
||||
save_file(cpu_state, weight_path)
|
||||
|
||||
|
||||
main_print(f"--> saved HF weight checkpoint at path {hf_weight_dir}")
|
||||
model_weight_dir = os.path.join(save_dir, "model_weights_state")
|
||||
os.makedirs(model_weight_dir, exist_ok=True)
|
||||
model_optimizer_dir = os.path.join(save_dir, "model_optimizer_state")
|
||||
os.makedirs(model_optimizer_dir, exist_ok=True)
|
||||
with FSDP.state_dict_type(model, StateDictType.SHARDED_STATE_DICT):
|
||||
optim_state = FSDP.optim_state_dict(model, optimizer)
|
||||
model_state = model.state_dict()
|
||||
weight_state_dict = {"model": model_state}
|
||||
dist_cp.save_state_dict(
|
||||
state_dict=weight_state_dict,
|
||||
storage_writer=dist_cp.FileSystemWriter(model_weight_dir),
|
||||
planner=DefaultSavePlanner(),
|
||||
)
|
||||
optimizer_state_dict = {"optimizer": optim_state}
|
||||
dist_cp.save_state_dict(
|
||||
state_dict=optimizer_state_dict,
|
||||
storage_writer=dist_cp.FileSystemWriter(model_optimizer_dir),
|
||||
planner=DefaultSavePlanner(),
|
||||
)
|
||||
|
||||
|
||||
discriminator_fsdp_state_dir = os.path.join(save_dir, "discriminator_fsdp_state")
|
||||
os.makedirs(discriminator_fsdp_state_dir, exist_ok=True)
|
||||
with FSDP.state_dict_type(discriminator, StateDictType.FULL_STATE_DICT, FullStateDictConfig(offload_to_cpu=True, rank0_only=True), FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True)):
|
||||
optim_state = FSDP.optim_state_dict(discriminator, discriminator_optimizer)
|
||||
model_state = discriminator.state_dict()
|
||||
state_dict = {"optimizer": optim_state, "model": model_state}
|
||||
if rank <=0:
|
||||
discriminator_fsdp_state_fil = os.path.join(discriminator_fsdp_state_dir, "discriminator_state.pt")
|
||||
torch.save(state_dict, discriminator_fsdp_state_fil)
|
||||
|
||||
main_print("--> saved FSDP state checkpoint")
|
||||
|
||||
def load_sharded_model(model, optimizer, model_dir, optimizer_dir):
|
||||
with FSDP.state_dict_type(model, StateDictType.SHARDED_STATE_DICT):
|
||||
weight_state_dict = {"model": model.state_dict()}
|
||||
|
||||
optim_state = load_sharded_optimizer_state_dict(
|
||||
model_state_dict=weight_state_dict["model"],
|
||||
optimizer_key="optimizer",
|
||||
storage_reader=dist_cp.FileSystemReader(optimizer_dir),
|
||||
)
|
||||
optim_state = optim_state["optimizer"]
|
||||
flattened_osd = FSDP.optim_state_dict_to_load(model=model, optim=optimizer, optim_state_dict=optim_state)
|
||||
optimizer.load_state_dict(flattened_osd)
|
||||
dist_cp.load_state_dict(
|
||||
state_dict = weight_state_dict,
|
||||
storage_reader=dist_cp.FileSystemReader(model_dir),
|
||||
planner=DefaultLoadPlanner(),
|
||||
)
|
||||
model_state = weight_state_dict["model"]
|
||||
model.load_state_dict(model_state)
|
||||
main_print(f"--> loaded model and optimizer from path {model_dir}")
|
||||
return model, optimizer
|
||||
|
||||
def load_full_state_model(model, optimizer, checkpoint_file, rank):
|
||||
with FSDP.state_dict_type(model, StateDictType.FULL_STATE_DICT, FullStateDictConfig(offload_to_cpu=True, rank0_only=True), FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True)):
|
||||
discriminator_state = torch.load(checkpoint_file)
|
||||
model_state = discriminator_state["model"]
|
||||
if rank <= 0:
|
||||
optim_state = discriminator_state["optimizer"]
|
||||
else:
|
||||
optim_state = None
|
||||
model.load_state_dict(model_state)
|
||||
discriminator_optim_state = FSDP.optim_state_dict_to_load(model=model, optim=optimizer, optim_state_dict=optim_state)
|
||||
optimizer.load_state_dict(discriminator_optim_state)
|
||||
main_print(f"--> loaded discriminator and discriminator optimizer from path {checkpoint_file}")
|
||||
return model, optimizer
|
||||
|
||||
def resume_training_generator_discriminator(model, optimizer, discriminator, discriminator_optimizer, checkpoint_dir, rank):
|
||||
step = int(checkpoint_dir.split("-")[-1])
|
||||
model_weight_dir = os.path.join(checkpoint_dir, "model_weights_state")
|
||||
model_optimizer_dir = os.path.join(checkpoint_dir, "model_optimizer_state")
|
||||
model, optimizer = load_sharded_model(model, optimizer, model_weight_dir, model_optimizer_dir)
|
||||
discriminator_ckpt_file = os.path.join(checkpoint_dir, "discriminator_fsdp_state", "discriminator_state.pt")
|
||||
discriminator, discriminator_optimizer = load_full_state_model(discriminator, discriminator_optimizer, discriminator_ckpt_file, rank)
|
||||
return model, optimizer, discriminator, discriminator_optimizer, step
|
||||
|
||||
|
||||
def resume_training(model, optimizer, checkpoint_dir, discriminator=False):
|
||||
weight_path = os.path.join(checkpoint_dir, "diffusion_pytorch_model.safetensors")
|
||||
if discriminator:
|
||||
weight_path = os.path.join(checkpoint_dir, "discriminator_pytorch_model.safetensors")
|
||||
model_weights = load_file(weight_path)
|
||||
|
||||
with FSDP.state_dict_type(
|
||||
model, StateDictType.FULL_STATE_DICT, FullStateDictConfig(offload_to_cpu=True, rank0_only=True), FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
):
|
||||
current_state = model.state_dict()
|
||||
current_state.update(model_weights)
|
||||
model.load_state_dict(current_state, strict=False)
|
||||
if discriminator:
|
||||
optim_path = os.path.join(checkpoint_dir, "discriminator_optimizer.pt")
|
||||
else:
|
||||
optim_path = os.path.join(checkpoint_dir, "optimizer.pt")
|
||||
optimizer_state_dict = torch.load(optim_path, weights_only=False)
|
||||
optim_state = FSDP.optim_state_dict_to_load(
|
||||
model=model,
|
||||
optim=optimizer,
|
||||
optim_state_dict=optimizer_state_dict
|
||||
)
|
||||
optimizer.load_state_dict(optim_state)
|
||||
step = int(checkpoint_dir.split("-")[-1])
|
||||
return model, optimizer, step
|
||||
|
||||
|
||||
def save_lora_checkpoint(
|
||||
transformer,
|
||||
optimizer,
|
||||
rank,
|
||||
output_dir,
|
||||
step
|
||||
):
|
||||
main_print(f"--> saving LoRA checkpoint at step {step}")
|
||||
with FSDP.state_dict_type(
|
||||
transformer,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
):
|
||||
full_state_dict = transformer.state_dict()
|
||||
lora_state_dict = {
|
||||
k: v for k, v in full_state_dict.items()
|
||||
if 'lora' in k.lower()
|
||||
}
|
||||
lora_optim_state = FSDP.optim_state_dict(
|
||||
transformer,
|
||||
optimizer,
|
||||
)
|
||||
if rank <= 0:
|
||||
save_dir = os.path.join(output_dir, f"lora-checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
weight_path = os.path.join(save_dir, "lora_weights.safetensors")
|
||||
save_file(lora_state_dict, weight_path)
|
||||
optim_path = os.path.join(save_dir, "lora_optimizer.pt")
|
||||
torch.save(lora_optim_state, optim_path)
|
||||
lora_config = {
|
||||
'step': step,
|
||||
'lora_params': {
|
||||
'lora_rank': transformer.config.lora_rank,
|
||||
'lora_alpha': transformer.config.lora_alpha,
|
||||
'target_modules': transformer.config.lora_target_modules
|
||||
}
|
||||
}
|
||||
config_path = os.path.join(save_dir, "lora_config.json")
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(lora_config, f, indent=4)
|
||||
main_print(f"--> LoRA checkpoint saved at step {step}")
|
||||
|
||||
def resume_lora_training(
|
||||
transformer,
|
||||
checkpoint_dir,
|
||||
optimizer
|
||||
):
|
||||
weight_path = os.path.join(checkpoint_dir, "lora_weights.safetensors")
|
||||
lora_weights = load_file(weight_path)
|
||||
config_path = os.path.join(checkpoint_dir, "lora_config.json")
|
||||
with open(config_path, "r") as f:
|
||||
config_dict = json.load(f)
|
||||
with FSDP.state_dict_type(
|
||||
transformer,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
|
||||
):
|
||||
current_state = transformer.state_dict()
|
||||
current_state.update(lora_weights)
|
||||
transformer.load_state_dict(current_state, strict=False)
|
||||
optim_path = os.path.join(checkpoint_dir, "lora_optimizer.pt")
|
||||
optimizer_state_dict = torch.load(optim_path, weights_only=False)
|
||||
optim_state = FSDP.optim_state_dict_to_load(
|
||||
model=transformer,
|
||||
optim=optimizer,
|
||||
optim_state_dict=optimizer_state_dict
|
||||
)
|
||||
optimizer.load_state_dict(optim_state)
|
||||
step = config_dict['step']
|
||||
main_print(f"--> Successfully resuming LoRA training from step {step}")
|
||||
return transformer, optimizer, step
|
||||
@@ -0,0 +1,298 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from fastvideo.utils.parallel_states import nccl_info
|
||||
from typing import Any, Tuple
|
||||
from torch import Tensor
|
||||
from torch.nn import Module
|
||||
|
||||
def broadcast(input_: torch.Tensor):
|
||||
src = nccl_info.group_id * nccl_info.sp_size
|
||||
dist.broadcast(input_, src=src, group=nccl_info.group)
|
||||
|
||||
|
||||
def _all_to_all_4D(
|
||||
input: torch.tensor, scatter_idx: int = 2, gather_idx: int = 1, group=None
|
||||
) -> torch.tensor:
|
||||
"""
|
||||
all-to-all for QKV
|
||||
|
||||
Args:
|
||||
input (torch.tensor): a tensor sharded along dim scatter dim
|
||||
scatter_idx (int): default 1
|
||||
gather_idx (int): default 2
|
||||
group : torch process group
|
||||
|
||||
Returns:
|
||||
torch.tensor: resharded tensor (bs, seqlen/P, hc, hs)
|
||||
"""
|
||||
assert (
|
||||
input.dim() == 4
|
||||
), f"input must be 4D tensor, got {input.dim()} and shape {input.shape}"
|
||||
|
||||
seq_world_size = dist.get_world_size(group)
|
||||
|
||||
if scatter_idx == 2 and gather_idx == 1:
|
||||
# input (torch.tensor): a tensor sharded along dim 1 (bs, seqlen/P, hc, hs) output: (bs, seqlen, hc/P, hs)
|
||||
bs, shard_seqlen, hc, hs = input.shape
|
||||
seqlen = shard_seqlen * seq_world_size
|
||||
shard_hc = hc // seq_world_size
|
||||
|
||||
# transpose groups of heads with the seq-len parallel dimension, so that we can scatter them!
|
||||
# (bs, seqlen/P, hc, hs) -reshape-> (bs, seq_len/P, P, hc/P, hs) -transpose(0,2)-> (P, seq_len/P, bs, hc/P, hs)
|
||||
input_t = (
|
||||
input.reshape(bs, shard_seqlen, seq_world_size, shard_hc, hs)
|
||||
.transpose(0, 2)
|
||||
.contiguous()
|
||||
)
|
||||
|
||||
output = torch.empty_like(input_t)
|
||||
# https://pytorch.org/docs/stable/distributed.html#torch.distributed.all_to_all_single
|
||||
# (P, seq_len/P, bs, hc/P, hs) scatter seqlen -all2all-> (P, seq_len/P, bs, hc/P, hs) scatter head
|
||||
if seq_world_size > 1:
|
||||
dist.all_to_all_single(output, input_t, group=group)
|
||||
torch.cuda.synchronize()
|
||||
else:
|
||||
output = input_t
|
||||
# if scattering the seq-dim, transpose the heads back to the original dimension
|
||||
output = output.reshape(seqlen, bs, shard_hc, hs)
|
||||
|
||||
# (seq_len, bs, hc/P, hs) -reshape-> (bs, seq_len, hc/P, hs)
|
||||
output = output.transpose(0, 1).contiguous().reshape(bs, seqlen, shard_hc, hs)
|
||||
|
||||
return output
|
||||
|
||||
elif scatter_idx == 1 and gather_idx == 2:
|
||||
# input (torch.tensor): a tensor sharded along dim 1 (bs, seqlen, hc/P, hs) output: (bs, seqlen/P, hc, hs)
|
||||
bs, seqlen, shard_hc, hs = input.shape
|
||||
hc = shard_hc * seq_world_size
|
||||
shard_seqlen = seqlen // seq_world_size
|
||||
seq_world_size = dist.get_world_size(group)
|
||||
|
||||
# transpose groups of heads with the seq-len parallel dimension, so that we can scatter them!
|
||||
# (bs, seqlen, hc/P, hs) -reshape-> (bs, P, seq_len/P, hc/P, hs) -transpose(0, 3)-> (hc/P, P, seqlen/P, bs, hs) -transpose(0, 1) -> (P, hc/P, seqlen/P, bs, hs)
|
||||
input_t = (
|
||||
input.reshape(bs, seq_world_size, shard_seqlen, shard_hc, hs)
|
||||
.transpose(0, 3)
|
||||
.transpose(0, 1)
|
||||
.contiguous()
|
||||
.reshape(seq_world_size, shard_hc, shard_seqlen, bs, hs)
|
||||
)
|
||||
|
||||
output = torch.empty_like(input_t)
|
||||
# https://pytorch.org/docs/stable/distributed.html#torch.distributed.all_to_all_single
|
||||
# (P, bs x hc/P, seqlen/P, hs) scatter seqlen -all2all-> (P, bs x seq_len/P, hc/P, hs) scatter head
|
||||
if seq_world_size > 1:
|
||||
dist.all_to_all_single(output, input_t, group=group)
|
||||
torch.cuda.synchronize()
|
||||
else:
|
||||
output = input_t
|
||||
|
||||
# if scattering the seq-dim, transpose the heads back to the original dimension
|
||||
output = output.reshape(hc, shard_seqlen, bs, hs)
|
||||
|
||||
# (hc, seqlen/N, bs, hs) -tranpose(0,2)-> (bs, seqlen/N, hc, hs)
|
||||
output = output.transpose(0, 2).contiguous().reshape(bs, shard_seqlen, hc, hs)
|
||||
|
||||
return output
|
||||
else:
|
||||
raise RuntimeError("scatter_idx must be 1 or 2 and gather_idx must be 1 or 2")
|
||||
|
||||
|
||||
class SeqAllToAll4D(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx: Any,
|
||||
group: dist.ProcessGroup,
|
||||
input: Tensor,
|
||||
scatter_idx: int,
|
||||
gather_idx: int,
|
||||
) -> Tensor:
|
||||
|
||||
ctx.group = group
|
||||
ctx.scatter_idx = scatter_idx
|
||||
ctx.gather_idx = gather_idx
|
||||
|
||||
return _all_to_all_4D(input, scatter_idx, gather_idx, group=group)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx: Any, *grad_output: Tensor) -> Tuple[None, Tensor, None, None]:
|
||||
return (
|
||||
None,
|
||||
SeqAllToAll4D.apply(
|
||||
ctx.group, *grad_output, ctx.gather_idx, ctx.scatter_idx
|
||||
),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def all_to_all_4D(
|
||||
input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1,
|
||||
):
|
||||
return SeqAllToAll4D.apply( nccl_info.group,input_, scatter_dim, gather_dim)
|
||||
|
||||
|
||||
|
||||
|
||||
def _all_to_all(
|
||||
input_: torch.Tensor,
|
||||
world_size: int,
|
||||
group: dist.ProcessGroup,
|
||||
scatter_dim: int,
|
||||
gather_dim: int,
|
||||
):
|
||||
input_list = [t.contiguous() for t in torch.tensor_split(input_, world_size, scatter_dim)]
|
||||
output_list = [torch.empty_like(input_list[0]) for _ in range(world_size)]
|
||||
dist.all_to_all(output_list, input_list, group=group)
|
||||
return torch.cat(output_list, dim=gather_dim).contiguous()
|
||||
|
||||
|
||||
class _AllToAll(torch.autograd.Function):
|
||||
"""All-to-all communication.
|
||||
|
||||
Args:
|
||||
input_: input matrix
|
||||
process_group: communication group
|
||||
scatter_dim: scatter dimension
|
||||
gather_dim: gather dimension
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input_, process_group, scatter_dim, gather_dim):
|
||||
ctx.process_group = process_group
|
||||
ctx.scatter_dim = scatter_dim
|
||||
ctx.gather_dim = gather_dim
|
||||
ctx.world_size = dist.get_world_size(process_group)
|
||||
output = _all_to_all(input_, ctx.world_size, process_group, scatter_dim, gather_dim)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
grad_output = _all_to_all(
|
||||
grad_output,
|
||||
ctx.world_size,
|
||||
ctx.process_group,
|
||||
ctx.gather_dim,
|
||||
ctx.scatter_dim,
|
||||
)
|
||||
return (
|
||||
grad_output,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def all_to_all(
|
||||
input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1,
|
||||
):
|
||||
return _AllToAll.apply(input_, nccl_info.group, scatter_dim, gather_dim)
|
||||
|
||||
|
||||
|
||||
class _AllGather(torch.autograd.Function):
|
||||
"""All-gather communication with autograd support.
|
||||
|
||||
Args:
|
||||
input_: input tensor
|
||||
dim: dimension along which to concatenate
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input_, dim):
|
||||
ctx.dim = dim
|
||||
world_size = nccl_info.sp_size
|
||||
group = nccl_info.group
|
||||
input_size = list(input_.size())
|
||||
|
||||
ctx.input_size = input_size[dim]
|
||||
|
||||
tensor_list = [torch.empty_like(input_) for _ in range(world_size)]
|
||||
input_ = input_.contiguous()
|
||||
dist.all_gather(tensor_list, input_, group=group)
|
||||
|
||||
output = torch.cat(tensor_list, dim=dim)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
world_size = nccl_info.sp_size
|
||||
rank = nccl_info.rank_within_group
|
||||
dim = ctx.dim
|
||||
input_size = ctx.input_size
|
||||
|
||||
sizes = [input_size] * world_size
|
||||
|
||||
grad_input_list = torch.split(grad_output, sizes, dim=dim)
|
||||
grad_input = grad_input_list[rank]
|
||||
|
||||
return grad_input, None
|
||||
|
||||
def all_gather(input_: torch.Tensor, dim: int = 1):
|
||||
"""Performs an all-gather operation on the input tensor along the specified dimension.
|
||||
|
||||
Args:
|
||||
input_ (torch.Tensor): Input tensor of shape [B, H, S, D].
|
||||
dim (int, optional): Dimension along which to concatenate. Defaults to 1.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor after all-gather operation, concatenated along 'dim'.
|
||||
"""
|
||||
return _AllGather.apply(input_, dim)
|
||||
|
||||
|
||||
def prepare_sequence_parallel_data(hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask):
|
||||
if nccl_info.sp_size == 1:
|
||||
return hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask
|
||||
def prepare(hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask):
|
||||
|
||||
hidden_states = all_to_all(hidden_states, scatter_dim=2, gather_dim=0)
|
||||
encoder_hidden_states = all_to_all(encoder_hidden_states, scatter_dim=1, gather_dim=0)
|
||||
attention_mask = all_to_all(attention_mask, scatter_dim=1, gather_dim=0)
|
||||
encoder_attention_mask = all_to_all(encoder_attention_mask, scatter_dim=1, gather_dim=0)
|
||||
return hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask
|
||||
|
||||
sp_size = nccl_info.sp_size
|
||||
frame = hidden_states.shape[2]
|
||||
assert frame % sp_size == 0, "frame should be a multiple of sp_size"
|
||||
|
||||
hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask = prepare(hidden_states,
|
||||
encoder_hidden_states.repeat(1, sp_size, 1),
|
||||
attention_mask.repeat(1, sp_size, 1, 1),
|
||||
encoder_attention_mask.repeat(1, sp_size))
|
||||
|
||||
|
||||
return hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask
|
||||
|
||||
|
||||
|
||||
def sp_parallel_dataloader_wrapper(dataloader, device, train_batch_size, sp_size, train_sp_batch_size):
|
||||
while True:
|
||||
for data_item in dataloader:
|
||||
latents, cond,attn_mask, cond_mask = data_item
|
||||
latents = latents.to(device)
|
||||
cond = cond.to(device)
|
||||
attn_mask = attn_mask.to(device)
|
||||
cond_mask = cond_mask.to(device)
|
||||
frame = latents.shape[2]
|
||||
if frame == 1:
|
||||
yield latents, cond, attn_mask, cond_mask
|
||||
else:
|
||||
latents, cond, attn_mask, cond_mask = prepare_sequence_parallel_data(latents, cond, attn_mask, cond_mask)
|
||||
assert train_batch_size * sp_size >= train_sp_batch_size, "train_batch_size * sp_size should be greater than train_sp_batch_size"
|
||||
for iter in range(train_batch_size * sp_size // train_sp_batch_size):
|
||||
st_idx = iter * train_sp_batch_size
|
||||
ed_idx = (iter + 1) * train_sp_batch_size
|
||||
encoder_hidden_states=cond[st_idx: ed_idx]
|
||||
attention_mask=attn_mask[st_idx: ed_idx]
|
||||
encoder_attention_mask=cond_mask[st_idx: ed_idx]
|
||||
yield latents[st_idx: ed_idx], encoder_hidden_states, attention_mask, encoder_attention_mask
|
||||
@@ -0,0 +1,117 @@
|
||||
import argparse
|
||||
import torch
|
||||
from accelerate.logging import get_logger
|
||||
from fastvideo.model.pipeline_mochi import MochiPipeline
|
||||
from diffusers.utils import export_to_video
|
||||
import json
|
||||
import os
|
||||
import torch.distributed as dist
|
||||
logger = get_logger(__name__)
|
||||
from torch.utils.data import Dataset
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
class T5dataset(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
json_path,
|
||||
vae_debug,
|
||||
):
|
||||
self.json_path = json_path
|
||||
self.vae_debug = vae_debug
|
||||
with open(self.json_path, "r") as f:
|
||||
train_dataset = json.load(f)
|
||||
self.train_dataset = sorted(train_dataset, key=lambda x: x['latent_path'])
|
||||
def __getitem__(self, idx):
|
||||
caption = self.train_dataset[idx]['caption']
|
||||
filename = self.train_dataset[idx]['latent_path'].split('.')[0]
|
||||
length = self.train_dataset[idx]['length']
|
||||
if self.vae_debug:
|
||||
latents = torch.load(os.path.join(args.output_dir, 'latent', self.train_dataset[idx]['latent_path']), map_location="cpu")
|
||||
else:
|
||||
latents = []
|
||||
|
||||
return dict(caption=caption, latents=latents, filename=filename, length=length)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.train_dataset)
|
||||
|
||||
def main(args):
|
||||
local_rank = int(os.getenv('RANK', 0))
|
||||
world_size = int(os.getenv('WORLD_SIZE', 1))
|
||||
print('world_size', world_size, 'local rank', local_rank)
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
torch.cuda.set_device(local_rank)
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(backend='nccl', init_method='env://', world_size=world_size, rank=local_rank)
|
||||
|
||||
pipe = MochiPipeline.from_pretrained(args.model_path).to(device)
|
||||
pipe.vae.enable_tiling()
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "video"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "prompt_embed"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "prompt_attention_mask"), exist_ok=True)
|
||||
|
||||
latents_json_path = os.path.join(args.output_dir, "videos2caption_temp.json")
|
||||
train_dataset = T5dataset(latents_json_path, args.vae_debug)
|
||||
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
batch_size=args.train_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
json_data = []
|
||||
for _, data in enumerate(train_dataloader):
|
||||
with torch.inference_mode():
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
prompt_embeds, prompt_attention_mask, _, _ = pipe.encode_prompt(
|
||||
prompt=data['caption'],
|
||||
)
|
||||
if args.vae_debug:
|
||||
latents = data['latents']
|
||||
video = pipe.vae.decode(latents.to(device), return_dict=False)[0]
|
||||
video = pipe.video_processor.postprocess_video(video)
|
||||
for idx, video_name in enumerate(data['filename']):
|
||||
prompt_embed_path = os.path.join(args.output_dir, "prompt_embed", video_name + ".pt")
|
||||
video_path = os.path.join(args.output_dir, "video", video_name + ".mp4")
|
||||
prompt_attention_mask_path = os.path.join(args.output_dir, "prompt_attention_mask", video_name + ".pt")
|
||||
# save latent
|
||||
torch.save(prompt_embeds[idx], prompt_embed_path)
|
||||
torch.save(prompt_attention_mask[idx], prompt_attention_mask_path)
|
||||
print(f"sample {video_name} saved")
|
||||
if args.vae_debug:
|
||||
export_to_video(video[idx], video_path, fps=30)
|
||||
item = {}
|
||||
item['length'] = int(data['length'][idx])
|
||||
item["latent_path"] = video_name + ".pt"
|
||||
item["prompt_embed_path"] = video_name + ".pt"
|
||||
item["prompt_attention_mask"] = video_name + ".pt"
|
||||
item["caption"] = data['caption'][idx]
|
||||
json_data.append(item)
|
||||
dist.barrier()
|
||||
local_data = json_data
|
||||
gathered_data = [None] * world_size
|
||||
dist.all_gather_object(gathered_data, local_data)
|
||||
if local_rank == 0:
|
||||
# os.remove(latents_json_path)
|
||||
all_json_data = [item for sublist in gathered_data for item in sublist]
|
||||
with open(os.path.join(args.output_dir, "videos2caption.json"), 'w') as f:
|
||||
json.dump(all_json_data, f, indent=4)
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--dataloader_num_workers", type=int, default=1, help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.")
|
||||
parser.add_argument("--train_batch_size", type=int, default=1, help="Batch size (per device) for the training dataloader.")
|
||||
parser.add_argument("--text_encoder_name", type=str, default='google/t5-v1_1-xxl')
|
||||
parser.add_argument("--cache_dir", type=str, default='./cache_dir')
|
||||
parser.add_argument("--output_dir", type=str, default=None, help="The output directory where the model predictions and checkpoints will be written.")
|
||||
parser.add_argument("--vae_debug",action="store_true")
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,106 @@
|
||||
from fastvideo.dataset import getdataset
|
||||
from torch.utils.data import DataLoader
|
||||
from fastvideo.utils.dataset_utils import Collate
|
||||
import argparse
|
||||
import torch
|
||||
from accelerate import Accelerator
|
||||
from accelerate.logging import get_logger
|
||||
from accelerate.utils import ProjectConfiguration
|
||||
import json
|
||||
import os
|
||||
from diffusers import AutoencoderKLMochi
|
||||
import torch.distributed as dist
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
def main(args):
|
||||
local_rank = int(os.getenv('RANK', 0))
|
||||
world_size = int(os.getenv('WORLD_SIZE', 1))
|
||||
print('world_size', world_size, 'local rank', local_rank)
|
||||
args.ae_stride_t, args.ae_stride_h, args.ae_stride_w = 4, 8, 8
|
||||
args.ae_stride = args.ae_stride_h
|
||||
patch_size_t, patch_size_h, patch_size_w = 1, 2, 2
|
||||
args.patch_size = patch_size_h
|
||||
args.patch_size_t, args.patch_size_h, args.patch_size_w = patch_size_t, patch_size_h, patch_size_w
|
||||
accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=args.logging_dir)
|
||||
accelerator = Accelerator(
|
||||
project_config=accelerator_project_config,
|
||||
)
|
||||
train_dataset = getdataset(args)
|
||||
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
batch_size=args.train_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
|
||||
encoder_device = torch.device(f"cuda" if torch.cuda.is_available() else "cpu")
|
||||
torch.cuda.set_device(local_rank)
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(backend='nccl', init_method='env://', world_size=world_size, rank=local_rank)
|
||||
vae = AutoencoderKLMochi.from_pretrained(args.model_path, subfolder="vae").to("cuda")
|
||||
vae.enable_tiling()
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
|
||||
|
||||
json_data = []
|
||||
for _, data in enumerate(train_dataloader):
|
||||
with torch.inference_mode():
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
latents = vae.encode(data['pixel_values'].to(encoder_device))['latent_dist'].sample()
|
||||
for idx, video_path in enumerate(data['path']):
|
||||
video_name = os.path.basename(video_path).split(".")[0]
|
||||
latent_path = os.path.join(args.output_dir, "latent", video_name + ".pt")
|
||||
torch.save(latents[idx].to(torch.bfloat16), latent_path)
|
||||
item = {}
|
||||
item["length"] = latents[idx].shape[1]
|
||||
item["latent_path"] = video_name + ".pt"
|
||||
item["caption"] = data['text'][idx]
|
||||
json_data.append(item)
|
||||
print(f"{video_name} processed")
|
||||
dist.barrier()
|
||||
local_data = json_data
|
||||
gathered_data = [None] * world_size
|
||||
dist.all_gather_object(gathered_data, local_data)
|
||||
if local_rank == 0:
|
||||
all_json_data = [item for sublist in gathered_data for item in sublist]
|
||||
with open(os.path.join(args.output_dir, "videos2caption_temp.json"), 'w') as f:
|
||||
json.dump(all_json_data, f, indent=4)
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--data_merge_path", type=str, required=True)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument("--dataloader_num_workers", type=int, default=1, help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.")
|
||||
parser.add_argument("--train_batch_size", type=int, default=16, help="Batch size (per device) for the training dataloader.")
|
||||
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
|
||||
parser.add_argument("--max_height", type=int, default=480)
|
||||
parser.add_argument("--max_width", type=int, default=848)
|
||||
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument("--dataset", default='t2v')
|
||||
parser.add_argument("--train_fps", type=int, default=30)
|
||||
parser.add_argument("--use_image_num", type=int, default=0)
|
||||
parser.add_argument("--text_max_length", type=int, default=256)
|
||||
parser.add_argument("--speed_factor", type=float, default=1.0)
|
||||
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--text_encoder_name", type=str, default='google/t5-v1_1-xxl')
|
||||
parser.add_argument("--cache_dir", type=str, default='./cache_dir')
|
||||
parser.add_argument('--cfg', type=float, default=0.0)
|
||||
parser.add_argument("--output_dir", type=str, default=None, help="The output directory where the model predictions and checkpoints will be written.")
|
||||
parser.add_argument("--logging_dir", type=str, default="logs",
|
||||
help=(
|
||||
"[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
|
||||
),
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,293 @@
|
||||
import math
|
||||
from einops import rearrange
|
||||
import decord
|
||||
from torch.nn import functional as F
|
||||
import torch
|
||||
from typing import Optional
|
||||
import torch.utils
|
||||
import torch.utils.data
|
||||
import torch
|
||||
from torch.utils.data import Sampler
|
||||
from typing import List
|
||||
from collections import Counter
|
||||
import random
|
||||
|
||||
|
||||
IMG_EXTENSIONS = ['.jpg', '.JPG', '.jpeg', '.JPEG', '.png', '.PNG']
|
||||
|
||||
def is_image_file(filename):
|
||||
return any(filename.endswith(extension) for extension in IMG_EXTENSIONS)
|
||||
|
||||
class DecordInit(object):
|
||||
"""Using Decord(https://github.com/dmlc/decord) to initialize the video_reader."""
|
||||
|
||||
def __init__(self, num_threads=1):
|
||||
self.num_threads = num_threads
|
||||
self.ctx = decord.cpu(0)
|
||||
|
||||
def __call__(self, filename):
|
||||
"""Perform the Decord initialization.
|
||||
Args:
|
||||
results (dict): The resulting dict to be modified and passed
|
||||
to the next transform in pipeline.
|
||||
"""
|
||||
reader = decord.VideoReader(filename,
|
||||
ctx=self.ctx,
|
||||
num_threads=self.num_threads)
|
||||
return reader
|
||||
|
||||
def __repr__(self):
|
||||
repr_str = (f'{self.__class__.__name__}('
|
||||
f'sr={self.sr},'
|
||||
f'num_threads={self.num_threads})')
|
||||
return repr_str
|
||||
|
||||
def pad_to_multiple(number, ds_stride):
|
||||
remainder = number % ds_stride
|
||||
if remainder == 0:
|
||||
return number
|
||||
else:
|
||||
padding = ds_stride - remainder
|
||||
return number + padding
|
||||
|
||||
class Collate:
|
||||
def __init__(self, args):
|
||||
self.batch_size = args.train_batch_size
|
||||
self.group_frame = args.group_frame
|
||||
self.group_resolution = args.group_resolution
|
||||
|
||||
self.max_height = args.max_height
|
||||
self.max_width = args.max_width
|
||||
self.ae_stride = args.ae_stride
|
||||
|
||||
self.ae_stride_t = args.ae_stride_t
|
||||
self.ae_stride_thw = (self.ae_stride_t, self.ae_stride, self.ae_stride)
|
||||
|
||||
self.patch_size = args.patch_size
|
||||
self.patch_size_t = args.patch_size_t
|
||||
|
||||
self.num_frames = args.num_frames
|
||||
self.use_image_num = args.use_image_num
|
||||
self.max_thw = (self.num_frames, self.max_height, self.max_width)
|
||||
|
||||
def package(self, batch):
|
||||
batch_tubes = [i['pixel_values'] for i in batch] # b [c t h w]
|
||||
input_ids = [i['input_ids'] for i in batch] # b [1 l]
|
||||
cond_mask = [i['cond_mask'] for i in batch] # b [1 l]
|
||||
return batch_tubes, input_ids, cond_mask
|
||||
|
||||
def __call__(self, batch):
|
||||
batch_tubes, input_ids, cond_mask = self.package(batch)
|
||||
|
||||
ds_stride = self.ae_stride * self.patch_size
|
||||
t_ds_stride = self.ae_stride_t * self.patch_size_t
|
||||
|
||||
pad_batch_tubes, attention_mask, input_ids, cond_mask = self.process(batch_tubes, input_ids, cond_mask, t_ds_stride, ds_stride, self.max_thw, self.ae_stride_thw)
|
||||
assert not torch.any(torch.isnan(pad_batch_tubes)), 'after pad_batch_tubes'
|
||||
return pad_batch_tubes, attention_mask, input_ids, cond_mask
|
||||
|
||||
|
||||
def process(self, batch_tubes, input_ids, cond_mask, t_ds_stride, ds_stride, max_thw, ae_stride_thw):
|
||||
# pad to max multiple of ds_stride
|
||||
batch_input_size = [i.shape for i in batch_tubes] # [(c t h w), (c t h w)]
|
||||
assert len(batch_input_size) == self.batch_size
|
||||
if self.group_frame or self.group_resolution or self.batch_size == 1: #
|
||||
len_each_batch = batch_input_size
|
||||
idx_length_dict = dict([*zip(list(range(self.batch_size)), len_each_batch)])
|
||||
count_dict = Counter(len_each_batch)
|
||||
if len(count_dict) != 1:
|
||||
sorted_by_value = sorted(count_dict.items(), key=lambda item: item[1])
|
||||
pick_length = sorted_by_value[-1][0] # the highest frequency
|
||||
candidate_batch = [idx for idx, length in idx_length_dict.items() if length == pick_length]
|
||||
random_select_batch = [random.choice(candidate_batch) for _ in range(len(len_each_batch) - len(candidate_batch))]
|
||||
print(batch_input_size, idx_length_dict, count_dict, sorted_by_value, pick_length, candidate_batch, random_select_batch)
|
||||
pick_idx = candidate_batch + random_select_batch
|
||||
|
||||
batch_tubes = [batch_tubes[i] for i in pick_idx]
|
||||
batch_input_size = [i.shape for i in batch_tubes] # [(c t h w), (c t h w)]
|
||||
input_ids = [input_ids[i] for i in pick_idx] # b [1, l]
|
||||
cond_mask = [cond_mask[i] for i in pick_idx] # b [1, l]
|
||||
|
||||
for i in range(1, self.batch_size):
|
||||
assert batch_input_size[0] == batch_input_size[i]
|
||||
max_t = max([i[1] for i in batch_input_size])
|
||||
max_h = max([i[2] for i in batch_input_size])
|
||||
max_w = max([i[3] for i in batch_input_size])
|
||||
else:
|
||||
max_t, max_h, max_w = max_thw
|
||||
pad_max_t, pad_max_h, pad_max_w = pad_to_multiple(max_t-1+self.ae_stride_t, t_ds_stride), \
|
||||
pad_to_multiple(max_h, ds_stride), \
|
||||
pad_to_multiple(max_w, ds_stride)
|
||||
pad_max_t = pad_max_t + 1 - self.ae_stride_t
|
||||
each_pad_t_h_w = [
|
||||
[
|
||||
pad_max_t - i.shape[1],
|
||||
pad_max_h - i.shape[2],
|
||||
pad_max_w - i.shape[3]
|
||||
] for i in batch_tubes
|
||||
]
|
||||
pad_batch_tubes = [
|
||||
F.pad(im, (0, pad_w, 0, pad_h, 0, pad_t), value=0)
|
||||
for (pad_t, pad_h, pad_w), im in zip(each_pad_t_h_w, batch_tubes)
|
||||
]
|
||||
pad_batch_tubes = torch.stack(pad_batch_tubes, dim=0)
|
||||
|
||||
|
||||
max_tube_size = [pad_max_t, pad_max_h, pad_max_w]
|
||||
max_latent_size = [
|
||||
((max_tube_size[0]-1) // ae_stride_thw[0] + 1),
|
||||
max_tube_size[1] // ae_stride_thw[1],
|
||||
max_tube_size[2] // ae_stride_thw[2]
|
||||
]
|
||||
valid_latent_size = [
|
||||
[
|
||||
int(math.ceil((i[1]-1) / ae_stride_thw[0])) + 1,
|
||||
int(math.ceil(i[2] / ae_stride_thw[1])),
|
||||
int(math.ceil(i[3] / ae_stride_thw[2]))
|
||||
] for i in batch_input_size]
|
||||
attention_mask = [
|
||||
F.pad(torch.ones(i, dtype=pad_batch_tubes.dtype), (0, max_latent_size[2] - i[2],
|
||||
0, max_latent_size[1] - i[1],
|
||||
0, max_latent_size[0] - i[0]), value=0) for i in valid_latent_size]
|
||||
attention_mask = torch.stack(attention_mask) # b t h w
|
||||
if self.batch_size == 1 or self.group_frame or self.group_resolution:
|
||||
assert torch.all(attention_mask.bool())
|
||||
|
||||
input_ids = torch.stack(input_ids) # b 1 l
|
||||
cond_mask = torch.stack(cond_mask) # b 1 l
|
||||
|
||||
return pad_batch_tubes, attention_mask, input_ids, cond_mask
|
||||
|
||||
|
||||
def split_to_even_chunks(indices, lengths, num_chunks, batch_size):
|
||||
"""
|
||||
Split a list of indices into `chunks` chunks of roughly equal lengths.
|
||||
"""
|
||||
|
||||
if len(indices) % num_chunks != 0:
|
||||
chunks = [indices[i::num_chunks] for i in range(num_chunks)]
|
||||
else:
|
||||
num_indices_per_chunk = len(indices) // num_chunks
|
||||
|
||||
chunks = [[] for _ in range(num_chunks)]
|
||||
chunks_lengths = [0 for _ in range(num_chunks)]
|
||||
for index in indices:
|
||||
shortest_chunk = chunks_lengths.index(min(chunks_lengths))
|
||||
chunks[shortest_chunk].append(index)
|
||||
chunks_lengths[shortest_chunk] += lengths[index]
|
||||
if len(chunks[shortest_chunk]) == num_indices_per_chunk:
|
||||
chunks_lengths[shortest_chunk] = float("inf")
|
||||
# return chunks
|
||||
|
||||
pad_chunks = []
|
||||
for idx, chunk in enumerate(chunks):
|
||||
if batch_size != len(chunk):
|
||||
assert batch_size > len(chunk)
|
||||
if len(chunk) != 0:
|
||||
chunk = chunk + [random.choice(chunk) for _ in range(batch_size - len(chunk))]
|
||||
else:
|
||||
chunk = random.choice(pad_chunks)
|
||||
print(chunks[idx], '->', chunk)
|
||||
pad_chunks.append(chunk)
|
||||
return pad_chunks
|
||||
|
||||
def group_frame_fun(indices, lengths):
|
||||
# sort by num_frames
|
||||
indices.sort(key=lambda i: lengths[i], reverse=True)
|
||||
return indices
|
||||
|
||||
|
||||
def megabatch_frame_alignment(megabatches, lengths):
|
||||
aligned_magabatches = []
|
||||
for _, megabatch in enumerate(megabatches):
|
||||
assert len(megabatch) != 0
|
||||
len_each_megabatch = [lengths[i] for i in megabatch]
|
||||
idx_length_dict = dict([*zip(megabatch, len_each_megabatch)])
|
||||
count_dict = Counter(len_each_megabatch)
|
||||
|
||||
# mixed frame length, align megabatch inside
|
||||
if len(count_dict) != 1:
|
||||
sorted_by_value = sorted(count_dict.items(), key=lambda item: item[1])
|
||||
pick_length = sorted_by_value[-1][0] # the highest frequency
|
||||
candidate_batch = [idx for idx, length in idx_length_dict.items() if length == pick_length]
|
||||
random_select_batch = [random.choice(candidate_batch) for i in range(len(idx_length_dict) - len(candidate_batch))]
|
||||
aligned_magabatch = candidate_batch + random_select_batch
|
||||
aligned_magabatches.append(aligned_magabatch)
|
||||
# already aligned megabatches
|
||||
else:
|
||||
aligned_magabatches.append(megabatch)
|
||||
|
||||
return aligned_magabatches
|
||||
|
||||
|
||||
def get_length_grouped_indices(lengths, batch_size, world_size, generator=None, group_frame=False, group_resolution=False, seed=42):
|
||||
# We need to use torch for the random part as a distributed sampler will set the random seed for torch.
|
||||
if generator is None:
|
||||
generator = torch.Generator().manual_seed(seed) # every rank will generate a fixed order but random index
|
||||
|
||||
indices = torch.randperm(len(lengths), generator=generator).tolist()
|
||||
|
||||
# sort dataset according to frame
|
||||
indices = group_frame_fun(indices, lengths)
|
||||
|
||||
# chunk dataset to megabatches
|
||||
megabatch_size = world_size * batch_size
|
||||
megabatches = [indices[i: i + megabatch_size] for i in range(0, len(lengths), megabatch_size)]
|
||||
|
||||
# make sure the length in each magabatch is align with each other
|
||||
megabatches = megabatch_frame_alignment(megabatches, lengths)
|
||||
|
||||
# aplit aligned megabatch into batches
|
||||
megabatches = [split_to_even_chunks(megabatch, lengths, world_size, batch_size) for megabatch in megabatches]
|
||||
|
||||
# random megabatches to do video-image mix training
|
||||
indices = torch.randperm(len(megabatches), generator=generator).tolist()
|
||||
shuffled_megabatches = [megabatches[i] for i in indices]
|
||||
|
||||
# expand indices and return
|
||||
return [i for megabatch in shuffled_megabatches for batch in megabatch for i in batch]
|
||||
|
||||
|
||||
class LengthGroupedSampler(Sampler):
|
||||
r"""
|
||||
Sampler that samples indices in a way that groups together features of the dataset of roughly the same length while
|
||||
keeping a bit of randomness.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
batch_size: int,
|
||||
rank: int,
|
||||
world_size: int,
|
||||
lengths: Optional[List[int]] = None,
|
||||
group_frame=False,
|
||||
group_resolution=False,
|
||||
generator=None,
|
||||
):
|
||||
if lengths is None:
|
||||
raise ValueError("Lengths must be provided.")
|
||||
|
||||
self.batch_size = batch_size
|
||||
self.rank = rank
|
||||
self.world_size = world_size
|
||||
self.lengths = lengths
|
||||
self.group_frame = group_frame
|
||||
self.group_resolution = group_resolution
|
||||
self.generator = generator
|
||||
|
||||
def __len__(self):
|
||||
return len(self.lengths)
|
||||
|
||||
def __iter__(self):
|
||||
indices = get_length_grouped_indices(self.lengths, self.batch_size, self.world_size, group_frame=self.group_frame,
|
||||
group_resolution=self.group_resolution, generator=self.generator)
|
||||
def distributed_sampler(lst, rank, batch_size, world_size):
|
||||
result = []
|
||||
index = rank * batch_size
|
||||
while index < len(lst):
|
||||
result.extend(lst[index:index + batch_size])
|
||||
index += batch_size * world_size
|
||||
return result
|
||||
|
||||
indices = distributed_sampler(indices, self.rank, self.batch_size, self.world_size)
|
||||
return iter(indices)
|
||||
@@ -0,0 +1,328 @@
|
||||
import contextlib
|
||||
import copy
|
||||
import random
|
||||
from typing import Any, Dict, Iterable, List, Optional, Union
|
||||
|
||||
from diffusers.utils import (
|
||||
deprecate,
|
||||
is_torchvision_available,
|
||||
is_transformers_available,
|
||||
)
|
||||
|
||||
if is_transformers_available():
|
||||
import transformers
|
||||
|
||||
if is_torchvision_available():
|
||||
from torchvision import transforms
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
# Adapted from diffusers-style ema https://github.com/huggingface/diffusers/blob/main/src/diffusers/training_utils.py#L263
|
||||
class EMAModel:
|
||||
"""
|
||||
Exponential Moving Average of models weights
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
parameters: Iterable[torch.nn.Parameter],
|
||||
decay: float = 0.9999,
|
||||
min_decay: float = 0.0,
|
||||
update_after_step: int = 0,
|
||||
use_ema_warmup: bool = False,
|
||||
inv_gamma: Union[float, int] = 1.0,
|
||||
power: Union[float, int] = 2 / 3,
|
||||
model_cls: Optional[Any] = None,
|
||||
model_config: Dict[str, Any] = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
parameters (Iterable[torch.nn.Parameter]): The parameters to track.
|
||||
decay (float): The decay factor for the exponential moving average.
|
||||
min_decay (float): The minimum decay factor for the exponential moving average.
|
||||
update_after_step (int): The number of steps to wait before starting to update the EMA weights.
|
||||
use_ema_warmup (bool): Whether to use EMA warmup.
|
||||
inv_gamma (float):
|
||||
Inverse multiplicative factor of EMA warmup. Default: 1. Only used if `use_ema_warmup` is True.
|
||||
power (float): Exponential factor of EMA warmup. Default: 2/3. Only used if `use_ema_warmup` is True.
|
||||
device (Optional[Union[str, torch.device]]): The device to store the EMA weights on. If None, the EMA
|
||||
weights will be stored on CPU.
|
||||
|
||||
@crowsonkb's notes on EMA Warmup:
|
||||
If gamma=1 and power=1, implements a simple average. gamma=1, power=2/3 are good values for models you plan
|
||||
to train for a million or more steps (reaches decay factor 0.999 at 31.6K steps, 0.9999 at 1M steps),
|
||||
gamma=1, power=3/4 for models you plan to train for less (reaches decay factor 0.999 at 10K steps, 0.9999
|
||||
at 215.4k steps).
|
||||
"""
|
||||
|
||||
if isinstance(parameters, torch.nn.Module):
|
||||
deprecation_message = (
|
||||
"Passing a `torch.nn.Module` to `ExponentialMovingAverage` is deprecated. "
|
||||
"Please pass the parameters of the module instead."
|
||||
)
|
||||
deprecate(
|
||||
"passing a `torch.nn.Module` to `ExponentialMovingAverage`",
|
||||
"1.0.0",
|
||||
deprecation_message,
|
||||
standard_warn=False,
|
||||
)
|
||||
parameters = parameters.parameters()
|
||||
|
||||
# set use_ema_warmup to True if a torch.nn.Module is passed for backwards compatibility
|
||||
use_ema_warmup = True
|
||||
|
||||
if kwargs.get("max_value", None) is not None:
|
||||
deprecation_message = "The `max_value` argument is deprecated. Please use `decay` instead."
|
||||
deprecate("max_value", "1.0.0", deprecation_message, standard_warn=False)
|
||||
decay = kwargs["max_value"]
|
||||
|
||||
if kwargs.get("min_value", None) is not None:
|
||||
deprecation_message = "The `min_value` argument is deprecated. Please use `min_decay` instead."
|
||||
deprecate("min_value", "1.0.0", deprecation_message, standard_warn=False)
|
||||
min_decay = kwargs["min_value"]
|
||||
|
||||
parameters = list(parameters)
|
||||
self.shadow_params = [p.clone().detach() for p in parameters]
|
||||
|
||||
if kwargs.get("device", None) is not None:
|
||||
deprecation_message = "The `device` argument is deprecated. Please use `to` instead."
|
||||
deprecate("device", "1.0.0", deprecation_message, standard_warn=False)
|
||||
self.to(device=kwargs["device"])
|
||||
|
||||
self.temp_stored_params = None
|
||||
|
||||
self.decay = decay
|
||||
self.min_decay = min_decay
|
||||
self.update_after_step = update_after_step
|
||||
self.use_ema_warmup = use_ema_warmup
|
||||
self.inv_gamma = inv_gamma
|
||||
self.power = power
|
||||
self.optimization_step = 0
|
||||
self.cur_decay_value = None # set in `step()`
|
||||
|
||||
self.model_cls = model_cls
|
||||
self.model_config = model_config
|
||||
|
||||
@classmethod
|
||||
def extract_ema_kwargs(cls, kwargs):
|
||||
"""
|
||||
Extracts the EMA kwargs from the kwargs of a class method.
|
||||
"""
|
||||
ema_kwargs = {}
|
||||
for key in [
|
||||
"decay",
|
||||
"min_decay",
|
||||
"optimization_step",
|
||||
"update_after_step",
|
||||
"use_ema_warmup",
|
||||
"inv_gamma",
|
||||
"power",
|
||||
]:
|
||||
if kwargs.get(key, None) is not None:
|
||||
ema_kwargs[key] = kwargs.pop(key)
|
||||
return ema_kwargs
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, path, model_cls) -> "EMAModel":
|
||||
config = model_cls.load_config(path)
|
||||
ema_kwargs = cls.extract_ema_kwargs(config)
|
||||
model = model_cls.from_pretrained(path)
|
||||
|
||||
ema_model = cls(model.parameters(), model_cls=model_cls, model_config=config)
|
||||
|
||||
ema_model.load_state_dict(ema_kwargs)
|
||||
return ema_model
|
||||
|
||||
def save_pretrained(self, path):
|
||||
if self.model_cls is None:
|
||||
raise ValueError("`save_pretrained` can only be used if `model_cls` was defined at __init__.")
|
||||
|
||||
if self.model_config is None:
|
||||
raise ValueError("`save_pretrained` can only be used if `model_config` was defined at __init__.")
|
||||
|
||||
model = self.model_cls.from_config(self.model_config)
|
||||
state_dict = self.state_dict()
|
||||
state_dict.pop("shadow_params", None)
|
||||
|
||||
model.register_to_config(**state_dict)
|
||||
self.copy_to(model.parameters())
|
||||
model.save_pretrained(path)
|
||||
|
||||
def get_decay(self, optimization_step: int) -> float:
|
||||
"""
|
||||
Compute the decay factor for the exponential moving average.
|
||||
"""
|
||||
step = max(0, optimization_step - self.update_after_step - 1)
|
||||
|
||||
if step <= 0:
|
||||
return 0.0
|
||||
|
||||
if self.use_ema_warmup:
|
||||
cur_decay_value = 1 - (1 + step / self.inv_gamma) ** -self.power
|
||||
else:
|
||||
cur_decay_value = (1 + step) / (10 + step)
|
||||
|
||||
cur_decay_value = min(cur_decay_value, self.decay)
|
||||
# make sure decay is not smaller than min_decay
|
||||
cur_decay_value = max(cur_decay_value, self.min_decay)
|
||||
return cur_decay_value
|
||||
|
||||
@torch.no_grad()
|
||||
def step(self, parameters: Iterable[torch.nn.Parameter]):
|
||||
if isinstance(parameters, torch.nn.Module):
|
||||
deprecation_message = (
|
||||
"Passing a `torch.nn.Module` to `ExponentialMovingAverage.step` is deprecated. "
|
||||
"Please pass the parameters of the module instead."
|
||||
)
|
||||
deprecate(
|
||||
"passing a `torch.nn.Module` to `ExponentialMovingAverage.step`",
|
||||
"1.0.0",
|
||||
deprecation_message,
|
||||
standard_warn=False,
|
||||
)
|
||||
parameters = parameters.parameters()
|
||||
|
||||
parameters = list(parameters)
|
||||
|
||||
self.optimization_step += 1
|
||||
|
||||
# Compute the decay factor for the exponential moving average.
|
||||
decay = self.get_decay(self.optimization_step)
|
||||
self.cur_decay_value = decay
|
||||
one_minus_decay = 1 - decay
|
||||
|
||||
context_manager = contextlib.nullcontext
|
||||
if is_transformers_available() and transformers.integrations.is_deepspeed_zero3_enabled():
|
||||
import deepspeed
|
||||
|
||||
for s_param, param in zip(self.shadow_params, parameters):
|
||||
if is_transformers_available() and transformers.integrations.is_deepspeed_zero3_enabled():
|
||||
context_manager = deepspeed.zero.GatheredParameters(param, modifier_rank=None)
|
||||
|
||||
with context_manager():
|
||||
if param.requires_grad:
|
||||
s_param.sub_(one_minus_decay * (s_param - param))
|
||||
else:
|
||||
s_param.copy_(param)
|
||||
|
||||
def copy_to(self, parameters: Iterable[torch.nn.Parameter]) -> None:
|
||||
"""
|
||||
Copy current averaged parameters into given collection of parameters.
|
||||
|
||||
Args:
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
updated with the stored moving averages. If `None`, the parameters with which this
|
||||
`ExponentialMovingAverage` was initialized will be used.
|
||||
"""
|
||||
parameters = list(parameters)
|
||||
for s_param, param in zip(self.shadow_params, parameters):
|
||||
param.data.copy_(s_param.to(param.device).data)
|
||||
|
||||
|
||||
def to(self, device=None, dtype=None) -> None:
|
||||
r"""Move internal buffers of the ExponentialMovingAverage to `device`.
|
||||
|
||||
Args:
|
||||
device: like `device` argument to `torch.Tensor.to`
|
||||
"""
|
||||
# .to() on the tensors handles None correctly
|
||||
self.shadow_params = [
|
||||
p.to(device=device, dtype=dtype) if p.is_floating_point() else p.to(device=device)
|
||||
for p in self.shadow_params
|
||||
]
|
||||
|
||||
def state_dict(self) -> dict:
|
||||
r"""
|
||||
Returns the state of the ExponentialMovingAverage as a dict. This method is used by accelerate during
|
||||
checkpointing to save the ema state dict.
|
||||
"""
|
||||
# Following PyTorch conventions, references to tensors are returned:
|
||||
# "returns a reference to the state and not its copy!" -
|
||||
# https://pytorch.org/tutorials/beginner/saving_loading_models.html#what-is-a-state-dict
|
||||
return {
|
||||
"decay": self.decay,
|
||||
"min_decay": self.min_decay,
|
||||
"optimization_step": self.optimization_step,
|
||||
"update_after_step": self.update_after_step,
|
||||
"use_ema_warmup": self.use_ema_warmup,
|
||||
"inv_gamma": self.inv_gamma,
|
||||
"power": self.power,
|
||||
"shadow_params": self.shadow_params,
|
||||
}
|
||||
|
||||
def store(self, parameters: Iterable[torch.nn.Parameter]) -> None:
|
||||
r"""
|
||||
Args:
|
||||
Save the current parameters for restoring later.
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
temporarily stored.
|
||||
"""
|
||||
self.temp_stored_params = [param.detach().cpu().clone() for param in parameters]
|
||||
|
||||
def restore(self, parameters: Iterable[torch.nn.Parameter]) -> None:
|
||||
r"""
|
||||
Args:
|
||||
Restore the parameters stored with the `store` method. Useful to validate the model with EMA parameters without:
|
||||
affecting the original optimization process. Store the parameters before the `copy_to()` method. After
|
||||
validation (or model saving), use this to restore the former parameters.
|
||||
parameters: Iterable of `torch.nn.Parameter`; the parameters to be
|
||||
updated with the stored parameters. If `None`, the parameters with which this
|
||||
`ExponentialMovingAverage` was initialized will be used.
|
||||
"""
|
||||
if self.temp_stored_params is None:
|
||||
raise RuntimeError("This ExponentialMovingAverage has no `store()`ed weights " "to `restore()`")
|
||||
for c_param, param in zip(self.temp_stored_params, parameters):
|
||||
param.data.copy_(c_param.data)
|
||||
|
||||
# Better memory-wise.
|
||||
self.temp_stored_params = None
|
||||
|
||||
def load_state_dict(self, state_dict: dict) -> None:
|
||||
r"""
|
||||
Args:
|
||||
Loads the ExponentialMovingAverage state. This method is used by accelerate during checkpointing to save the
|
||||
ema state dict.
|
||||
state_dict (dict): EMA state. Should be an object returned
|
||||
from a call to :meth:`state_dict`.
|
||||
"""
|
||||
# deepcopy, to be consistent with module API
|
||||
state_dict = copy.deepcopy(state_dict)
|
||||
|
||||
self.decay = state_dict.get("decay", self.decay)
|
||||
if self.decay < 0.0 or self.decay > 1.0:
|
||||
raise ValueError("Decay must be between 0 and 1")
|
||||
|
||||
self.min_decay = state_dict.get("min_decay", self.min_decay)
|
||||
if not isinstance(self.min_decay, float):
|
||||
raise ValueError("Invalid min_decay")
|
||||
|
||||
self.optimization_step = state_dict.get("optimization_step", self.optimization_step)
|
||||
if not isinstance(self.optimization_step, int):
|
||||
raise ValueError("Invalid optimization_step")
|
||||
|
||||
self.update_after_step = state_dict.get("update_after_step", self.update_after_step)
|
||||
if not isinstance(self.update_after_step, int):
|
||||
raise ValueError("Invalid update_after_step")
|
||||
|
||||
self.use_ema_warmup = state_dict.get("use_ema_warmup", self.use_ema_warmup)
|
||||
if not isinstance(self.use_ema_warmup, bool):
|
||||
raise ValueError("Invalid use_ema_warmup")
|
||||
|
||||
self.inv_gamma = state_dict.get("inv_gamma", self.inv_gamma)
|
||||
if not isinstance(self.inv_gamma, (float, int)):
|
||||
raise ValueError("Invalid inv_gamma")
|
||||
|
||||
self.power = state_dict.get("power", self.power)
|
||||
if not isinstance(self.power, (float, int)):
|
||||
raise ValueError("Invalid power")
|
||||
|
||||
shadow_params = state_dict.get("shadow_params", None)
|
||||
if shadow_params is not None:
|
||||
self.shadow_params = shadow_params
|
||||
if not isinstance(self.shadow_params, list):
|
||||
raise ValueError("shadow_params must be a list")
|
||||
if not all(isinstance(p, torch.Tensor) for p in self.shadow_params):
|
||||
raise ValueError("shadow_params must all be Tensors")
|
||||
@@ -0,0 +1,23 @@
|
||||
|
||||
import sys
|
||||
import pdb
|
||||
import os
|
||||
|
||||
def main_print(content):
|
||||
if int(os.environ['LOCAL_RANK']) <= 0:
|
||||
print(content)
|
||||
|
||||
#ForkedPdb().set_trace()
|
||||
class ForkedPdb(pdb.Pdb):
|
||||
"""A Pdb subclass that may be used
|
||||
from a forked multiprocessing child
|
||||
|
||||
"""
|
||||
def interaction(self, *args, **kwargs):
|
||||
_stdin = sys.stdin
|
||||
try:
|
||||
sys.stdin = open('/dev/stdin')
|
||||
pdb.Pdb.interaction(self, *args, **kwargs)
|
||||
finally:
|
||||
sys.stdin = _stdin
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import os
|
||||
|
||||
class COMM_INFO:
|
||||
def __init__(self):
|
||||
self.group = None
|
||||
self.sp_size = 1
|
||||
self.global_rank = 0
|
||||
self.rank_within_group = 0
|
||||
self.group_id = 0
|
||||
|
||||
nccl_info = COMM_INFO()
|
||||
_SEQUENCE_PARALLEL_STATE = False
|
||||
def initialize_sequence_parallel_state(sequence_parallel_size):
|
||||
global _SEQUENCE_PARALLEL_STATE
|
||||
if sequence_parallel_size > 1:
|
||||
_SEQUENCE_PARALLEL_STATE = True
|
||||
initialize_sequence_parallel_group(sequence_parallel_size)
|
||||
else:
|
||||
nccl_info.sp_size = 1
|
||||
nccl_info.global_rank = int(os.getenv('RANK', '0'))
|
||||
nccl_info.rank_within_group = 0
|
||||
nccl_info.group_id = int(os.getenv('RANK', '0'))
|
||||
|
||||
def set_sequence_parallel_state(state):
|
||||
global _SEQUENCE_PARALLEL_STATE
|
||||
_SEQUENCE_PARALLEL_STATE = state
|
||||
|
||||
def get_sequence_parallel_state():
|
||||
return _SEQUENCE_PARALLEL_STATE
|
||||
|
||||
def initialize_sequence_parallel_group(sequence_parallel_size):
|
||||
"""Initialize the sequence parallel group."""
|
||||
rank = int(os.getenv('RANK', '0'))
|
||||
world_size = int(os.getenv("WORLD_SIZE", '1'))
|
||||
assert world_size % sequence_parallel_size == 0, "world_size must be divisible by sequence_parallel_size, but got world_size: {}, sequence_parallel_size: {}".format(world_size, sequence_parallel_size)
|
||||
nccl_info.sp_size = sequence_parallel_size
|
||||
nccl_info.global_rank = rank
|
||||
num_sequence_parallel_groups: int = world_size // sequence_parallel_size
|
||||
for i in range(num_sequence_parallel_groups):
|
||||
ranks = range(i * sequence_parallel_size, (i + 1) * sequence_parallel_size)
|
||||
group = dist.new_group(ranks)
|
||||
if rank in ranks:
|
||||
nccl_info.group = group
|
||||
nccl_info.rank_within_group = rank - i * sequence_parallel_size
|
||||
nccl_info.group_id = i
|
||||
|
||||
|
||||
def destroy_sequence_parallel_group():
|
||||
"""Destroy the sequence parallel group."""
|
||||
dist.destroy_process_group()
|
||||
@@ -0,0 +1,471 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
import os
|
||||
import math
|
||||
import torch
|
||||
import logging
|
||||
import random
|
||||
import subprocess
|
||||
import numpy as np
|
||||
import torch.distributed as dist
|
||||
|
||||
# from torch._six import inf
|
||||
from torch import inf
|
||||
from PIL import Image
|
||||
from typing import Union, Iterable
|
||||
import collections
|
||||
from collections import OrderedDict
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
|
||||
from diffusers.utils import is_bs4_available, is_ftfy_available
|
||||
|
||||
import html
|
||||
import re
|
||||
import urllib.parse as ul
|
||||
|
||||
if is_bs4_available():
|
||||
from bs4 import BeautifulSoup
|
||||
|
||||
if is_ftfy_available():
|
||||
import ftfy
|
||||
|
||||
_tensor_or_tensors = Union[torch.Tensor, Iterable[torch.Tensor]]
|
||||
|
||||
def to_2tuple(x):
|
||||
if isinstance(x, collections.abc.Iterable):
|
||||
return x
|
||||
return (x, x)
|
||||
|
||||
def find_model(model_name):
|
||||
"""
|
||||
Finds a pre-trained Latte model, downloading it if necessary. Alternatively, loads a model from a local path.
|
||||
"""
|
||||
assert os.path.isfile(model_name), f'Could not find Latte checkpoint at {model_name}'
|
||||
checkpoint = torch.load(model_name, map_location=lambda storage, loc: storage)
|
||||
|
||||
# if "ema" in checkpoint: # supports checkpoints from train.py
|
||||
# print('Using Ema!')
|
||||
# checkpoint = checkpoint["ema"]
|
||||
# else:
|
||||
print('Using model!')
|
||||
checkpoint = checkpoint['model']
|
||||
return checkpoint
|
||||
|
||||
#################################################################################
|
||||
# Training Clip Gradients #
|
||||
#################################################################################
|
||||
|
||||
def get_grad_norm(
|
||||
parameters: _tensor_or_tensors, norm_type: float = 2.0) -> torch.Tensor:
|
||||
r"""
|
||||
Copy from torch.nn.utils.clip_grad_norm_
|
||||
|
||||
Clips gradient norm of an iterable of parameters.
|
||||
|
||||
The norm is computed over all gradients together, as if they were
|
||||
concatenated into a single vector. Gradients are modified in-place.
|
||||
|
||||
Args:
|
||||
parameters (Iterable[Tensor] or Tensor): an iterable of Tensors or a
|
||||
single Tensor that will have gradients normalized
|
||||
max_norm (float or int): max norm of the gradients
|
||||
norm_type (float or int): type of the used p-norm. Can be ``'inf'`` for
|
||||
infinity norm.
|
||||
error_if_nonfinite (bool): if True, an error is thrown if the total
|
||||
norm of the gradients from :attr:`parameters` is ``nan``,
|
||||
``inf``, or ``-inf``. Default: False (will switch to True in the future)
|
||||
|
||||
Returns:
|
||||
Total norm of the parameter gradients (viewed as a single vector).
|
||||
"""
|
||||
if isinstance(parameters, torch.Tensor):
|
||||
parameters = [parameters]
|
||||
grads = [p.grad for p in parameters if p.grad is not None]
|
||||
norm_type = float(norm_type)
|
||||
if len(grads) == 0:
|
||||
return torch.tensor(0.)
|
||||
device = grads[0].device
|
||||
if norm_type == inf:
|
||||
norms = [g.detach().abs().max().to(device) for g in grads]
|
||||
total_norm = norms[0] if len(norms) == 1 else torch.max(torch.stack(norms))
|
||||
else:
|
||||
total_norm = torch.norm(torch.stack([torch.norm(g.detach(), norm_type).to(device) for g in grads]), norm_type)
|
||||
return total_norm
|
||||
|
||||
|
||||
def clip_grad_norm_(
|
||||
parameters: _tensor_or_tensors, max_norm: float, norm_type: float = 2.0,
|
||||
error_if_nonfinite: bool = False, clip_grad=True) -> torch.Tensor:
|
||||
r"""
|
||||
Copy from torch.nn.utils.clip_grad_norm_
|
||||
|
||||
Clips gradient norm of an iterable of parameters.
|
||||
|
||||
The norm is computed over all gradients together, as if they were
|
||||
concatenated into a single vector. Gradients are modified in-place.
|
||||
|
||||
Args:
|
||||
parameters (Iterable[Tensor] or Tensor): an iterable of Tensors or a
|
||||
single Tensor that will have gradients normalized
|
||||
max_norm (float or int): max norm of the gradients
|
||||
norm_type (float or int): type of the used p-norm. Can be ``'inf'`` for
|
||||
infinity norm.
|
||||
error_if_nonfinite (bool): if True, an error is thrown if the total
|
||||
norm of the gradients from :attr:`parameters` is ``nan``,
|
||||
``inf``, or ``-inf``. Default: False (will switch to True in the future)
|
||||
|
||||
Returns:
|
||||
Total norm of the parameter gradients (viewed as a single vector).
|
||||
"""
|
||||
if isinstance(parameters, torch.Tensor):
|
||||
parameters = [parameters]
|
||||
grads = [p.grad for p in parameters if p.grad is not None]
|
||||
max_norm = float(max_norm)
|
||||
norm_type = float(norm_type)
|
||||
if len(grads) == 0:
|
||||
return torch.tensor(0.)
|
||||
device = grads[0].device
|
||||
if norm_type == inf:
|
||||
norms = [g.detach().abs().max().to(device) for g in grads]
|
||||
total_norm = norms[0] if len(norms) == 1 else torch.max(torch.stack(norms))
|
||||
else:
|
||||
total_norm = torch.norm(torch.stack([torch.norm(g.detach(), norm_type).to(device) for g in grads]), norm_type)
|
||||
|
||||
if clip_grad:
|
||||
if error_if_nonfinite and torch.logical_or(total_norm.isnan(), total_norm.isinf()):
|
||||
raise RuntimeError(
|
||||
f'The total norm of order {norm_type} for gradients from '
|
||||
'`parameters` is non-finite, so it cannot be clipped. To disable '
|
||||
'this error and scale the gradients by the non-finite norm anyway, '
|
||||
'set `error_if_nonfinite=False`')
|
||||
clip_coef = max_norm / (total_norm + 1e-6)
|
||||
# Note: multiplying by the clamped coef is redundant when the coef is clamped to 1, but doing so
|
||||
# avoids a `if clip_coef < 1:` conditional which can require a CPU <=> device synchronization
|
||||
# when the gradients do not reside in CPU memory.
|
||||
clip_coef_clamped = torch.clamp(clip_coef, max=1.0)
|
||||
for g in grads:
|
||||
g.detach().mul_(clip_coef_clamped.to(g.device))
|
||||
# gradient_cliped = torch.norm(torch.stack([torch.norm(g.detach(), norm_type).to(device) for g in grads]), norm_type)
|
||||
# print(gradient_cliped)
|
||||
return total_norm
|
||||
|
||||
|
||||
def get_experiment_dir(root_dir, args):
|
||||
# if args.pretrained is not None and 'Latte-XL-2-256x256.pt' not in args.pretrained:
|
||||
# root_dir += '-WOPRE'
|
||||
if args.use_compile:
|
||||
root_dir += '-Compile' # speedup by torch compile
|
||||
if args.attention_mode:
|
||||
root_dir += f'-{args.attention_mode.upper()}'
|
||||
# if args.enable_xformers_memory_efficient_attention:
|
||||
# root_dir += '-Xfor'
|
||||
if args.gradient_checkpointing:
|
||||
root_dir += '-Gc'
|
||||
if args.mixed_precision:
|
||||
root_dir += f'-{args.mixed_precision.upper()}'
|
||||
root_dir += f'-{args.max_image_size}'
|
||||
return root_dir
|
||||
|
||||
def get_precision(args):
|
||||
if args.mixed_precision == "bf16":
|
||||
dtype = torch.bfloat16
|
||||
elif args.mixed_precision == "fp16":
|
||||
dtype = torch.float16
|
||||
else:
|
||||
dtype = torch.float32
|
||||
return dtype
|
||||
|
||||
#################################################################################
|
||||
# Training Logger #
|
||||
#################################################################################
|
||||
|
||||
def create_logger(logging_dir):
|
||||
"""
|
||||
Create a logger that writes to a log file and stdout.
|
||||
"""
|
||||
if dist.get_rank() == 0: # real logger
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
# format='[\033[34m%(asctime)s\033[0m] %(message)s',
|
||||
format='[%(asctime)s] %(message)s',
|
||||
datefmt='%Y-%m-%d %H:%M:%S',
|
||||
handlers=[logging.StreamHandler(), logging.FileHandler(f"{logging_dir}/log.txt")]
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
else: # dummy logger (does nothing)
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.addHandler(logging.NullHandler())
|
||||
return logger
|
||||
|
||||
|
||||
def create_tensorboard(tensorboard_dir):
|
||||
"""
|
||||
Create a tensorboard that saves losses.
|
||||
"""
|
||||
if dist.get_rank() == 0: # real tensorboard
|
||||
# tensorboard
|
||||
writer = SummaryWriter(tensorboard_dir)
|
||||
|
||||
return writer
|
||||
|
||||
|
||||
def write_tensorboard(writer, *args):
|
||||
'''
|
||||
write the loss information to a tensorboard file.
|
||||
Only for pytorch DDP mode.
|
||||
'''
|
||||
if dist.get_rank() == 0: # real tensorboard
|
||||
writer.add_scalar(args[0], args[1], args[2])
|
||||
|
||||
|
||||
#################################################################################
|
||||
# EMA Update/ DDP Training Utils #
|
||||
#################################################################################
|
||||
|
||||
@torch.no_grad()
|
||||
def update_ema(ema_model, model, decay=0.9999):
|
||||
"""
|
||||
Step the EMA model towards the current model.
|
||||
"""
|
||||
ema_params = OrderedDict(ema_model.named_parameters())
|
||||
model_params = OrderedDict(model.named_parameters())
|
||||
|
||||
for name, param in model_params.items():
|
||||
# TODO: Consider applying only to params that require_grad to avoid small numerical changes of pos_embed
|
||||
ema_params[name].mul_(decay).add_(param.data, alpha=1 - decay)
|
||||
|
||||
|
||||
def requires_grad(model, flag=True):
|
||||
"""
|
||||
Set requires_grad flag for all parameters in a model.
|
||||
"""
|
||||
for p in model.parameters():
|
||||
p.requires_grad = flag
|
||||
|
||||
|
||||
def cleanup():
|
||||
"""
|
||||
End DDP training.
|
||||
"""
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
def setup_distributed(backend="nccl", port=None):
|
||||
"""Initialize distributed training environment.
|
||||
support both slurm and torch.distributed.launch
|
||||
see torch.distributed.init_process_group() for more details
|
||||
"""
|
||||
num_gpus = torch.cuda.device_count()
|
||||
|
||||
if "SLURM_JOB_ID" in os.environ:
|
||||
rank = int(os.environ["SLURM_PROCID"])
|
||||
world_size = int(os.environ["SLURM_NTASKS"])
|
||||
node_list = os.environ["SLURM_NODELIST"]
|
||||
addr = subprocess.getoutput(f"scontrol show hostname {node_list} | head -n1")
|
||||
# specify master port
|
||||
if port is not None:
|
||||
os.environ["MASTER_PORT"] = str(port)
|
||||
elif "MASTER_PORT" not in os.environ:
|
||||
# os.environ["MASTER_PORT"] = "29566"
|
||||
os.environ["MASTER_PORT"] = str(29567 + num_gpus)
|
||||
if "MASTER_ADDR" not in os.environ:
|
||||
os.environ["MASTER_ADDR"] = addr
|
||||
os.environ["WORLD_SIZE"] = str(world_size)
|
||||
os.environ["LOCAL_RANK"] = str(rank % num_gpus)
|
||||
os.environ["RANK"] = str(rank)
|
||||
else:
|
||||
rank = int(os.environ["RANK"])
|
||||
world_size = int(os.environ["WORLD_SIZE"])
|
||||
|
||||
# torch.cuda.set_device(rank % num_gpus)
|
||||
|
||||
dist.init_process_group(
|
||||
backend=backend,
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
)
|
||||
|
||||
|
||||
#################################################################################
|
||||
# Testing Utils #
|
||||
#################################################################################
|
||||
|
||||
def save_video_grid(video, nrow=None):
|
||||
b, t, h, w, c = video.shape
|
||||
|
||||
if nrow is None:
|
||||
nrow = math.ceil(math.sqrt(b))
|
||||
ncol = math.ceil(b / nrow)
|
||||
padding = 1
|
||||
video_grid = torch.zeros((t, (padding + h) * nrow + padding,
|
||||
(padding + w) * ncol + padding, c), dtype=torch.uint8)
|
||||
|
||||
print(video_grid.shape)
|
||||
for i in range(b):
|
||||
r = i // ncol
|
||||
c = i % ncol
|
||||
start_r = (padding + h) * r
|
||||
start_c = (padding + w) * c
|
||||
video_grid[:, start_r:start_r + h, start_c:start_c + w] = video[i]
|
||||
|
||||
return video_grid
|
||||
|
||||
|
||||
#################################################################################
|
||||
# MMCV Utils #
|
||||
#################################################################################
|
||||
|
||||
|
||||
def collect_env():
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from mmcv.utils import collect_env as collect_base_env
|
||||
from mmcv.utils import get_git_hash
|
||||
"""Collect the information of the running environments."""
|
||||
|
||||
env_info = collect_base_env()
|
||||
env_info['MMClassification'] = get_git_hash()[:7]
|
||||
|
||||
for name, val in env_info.items():
|
||||
print(f'{name}: {val}')
|
||||
|
||||
print(torch.cuda.get_arch_list())
|
||||
print(torch.version.cuda)
|
||||
|
||||
|
||||
#################################################################################
|
||||
# Pixart-alpha Utils #
|
||||
#################################################################################
|
||||
|
||||
|
||||
bad_punct_regex = re.compile(r'['+'#®•©™&@·º½¾¿¡§~'+'\)'+'\('+'\]'+'\['+'\}'+'\{'+'\|'+'\\'+'\/'+'\*' + r']{1,}') # noqa
|
||||
|
||||
def text_preprocessing(text, support_Chinese=True):
|
||||
# The exact text cleaning as was in the training stage:
|
||||
text = clean_caption(text, support_Chinese=support_Chinese)
|
||||
text = clean_caption(text, support_Chinese=support_Chinese)
|
||||
return text
|
||||
|
||||
def basic_clean(text):
|
||||
text = ftfy.fix_text(text)
|
||||
text = html.unescape(html.unescape(text))
|
||||
return text.strip()
|
||||
|
||||
def clean_caption(caption, support_Chinese=True):
|
||||
caption = str(caption)
|
||||
caption = ul.unquote_plus(caption)
|
||||
caption = caption.strip().lower()
|
||||
caption = re.sub('<person>', 'person', caption)
|
||||
# urls:
|
||||
caption = re.sub(
|
||||
r'\b((?:https?:(?:\/{1,3}|[a-zA-Z0-9%])|[a-zA-Z0-9.\-]+[.](?:com|co|ru|net|org|edu|gov|it)[\w/-]*\b\/?(?!@)))', # noqa
|
||||
'', caption) # regex for urls
|
||||
caption = re.sub(
|
||||
r'\b((?:www:(?:\/{1,3}|[a-zA-Z0-9%])|[a-zA-Z0-9.\-]+[.](?:com|co|ru|net|org|edu|gov|it)[\w/-]*\b\/?(?!@)))', # noqa
|
||||
'', caption) # regex for urls
|
||||
# html:
|
||||
caption = BeautifulSoup(caption, features='html.parser').text
|
||||
|
||||
# @<nickname>
|
||||
caption = re.sub(r'@[\w\d]+\b', '', caption)
|
||||
|
||||
# 31C0—31EF CJK Strokes
|
||||
# 31F0—31FF Katakana Phonetic Extensions
|
||||
# 3200—32FF Enclosed CJK Letters and Months
|
||||
# 3300—33FF CJK Compatibility
|
||||
# 3400—4DBF CJK Unified Ideographs Extension A
|
||||
# 4DC0—4DFF Yijing Hexagram Symbols
|
||||
# 4E00—9FFF CJK Unified Ideographs
|
||||
caption = re.sub(r'[\u31c0-\u31ef]+', '', caption)
|
||||
caption = re.sub(r'[\u31f0-\u31ff]+', '', caption)
|
||||
caption = re.sub(r'[\u3200-\u32ff]+', '', caption)
|
||||
caption = re.sub(r'[\u3300-\u33ff]+', '', caption)
|
||||
caption = re.sub(r'[\u3400-\u4dbf]+', '', caption)
|
||||
caption = re.sub(r'[\u4dc0-\u4dff]+', '', caption)
|
||||
if not support_Chinese:
|
||||
caption = re.sub(r'[\u4e00-\u9fff]+', '', caption) # Chinese
|
||||
#######################################################
|
||||
|
||||
# все виды тире / all types of dash --> "-"
|
||||
caption = re.sub(
|
||||
r'[\u002D\u058A\u05BE\u1400\u1806\u2010-\u2015\u2E17\u2E1A\u2E3A\u2E3B\u2E40\u301C\u3030\u30A0\uFE31\uFE32\uFE58\uFE63\uFF0D]+', # noqa
|
||||
'-', caption)
|
||||
|
||||
# кавычки к одному стандарту
|
||||
caption = re.sub(r'[`´«»“”¨]', '"', caption)
|
||||
caption = re.sub(r'[‘’]', "'", caption)
|
||||
|
||||
# "
|
||||
caption = re.sub(r'"?', '', caption)
|
||||
# &
|
||||
caption = re.sub(r'&', '', caption)
|
||||
|
||||
# ip adresses:
|
||||
caption = re.sub(r'\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3}', ' ', caption)
|
||||
|
||||
# article ids:
|
||||
caption = re.sub(r'\d:\d\d\s+$', '', caption)
|
||||
|
||||
# \n
|
||||
caption = re.sub(r'\\n', ' ', caption)
|
||||
|
||||
# "#123"
|
||||
caption = re.sub(r'#\d{1,3}\b', '', caption)
|
||||
# "#12345.."
|
||||
caption = re.sub(r'#\d{5,}\b', '', caption)
|
||||
# "123456.."
|
||||
caption = re.sub(r'\b\d{6,}\b', '', caption)
|
||||
# filenames:
|
||||
caption = re.sub(r'[\S]+\.(?:png|jpg|jpeg|bmp|webp|eps|pdf|apk|mp4)', '', caption)
|
||||
|
||||
#
|
||||
caption = re.sub(r'[\"\']{2,}', r'"', caption) # """AUSVERKAUFT"""
|
||||
caption = re.sub(r'[\.]{2,}', r' ', caption) # """AUSVERKAUFT"""
|
||||
|
||||
caption = re.sub(bad_punct_regex, r' ', caption) # ***AUSVERKAUFT***, #AUSVERKAUFT
|
||||
caption = re.sub(r'\s+\.\s+', r' ', caption) # " . "
|
||||
|
||||
# this-is-my-cute-cat / this_is_my_cute_cat
|
||||
regex2 = re.compile(r'(?:\-|\_)')
|
||||
if len(re.findall(regex2, caption)) > 3:
|
||||
caption = re.sub(regex2, ' ', caption)
|
||||
|
||||
caption = basic_clean(caption)
|
||||
|
||||
caption = re.sub(r'\b[a-zA-Z]{1,3}\d{3,15}\b', '', caption) # jc6640
|
||||
caption = re.sub(r'\b[a-zA-Z]+\d+[a-zA-Z]+\b', '', caption) # jc6640vc
|
||||
caption = re.sub(r'\b\d+[a-zA-Z]+\d+\b', '', caption) # 6640vc231
|
||||
|
||||
caption = re.sub(r'(worldwide\s+)?(free\s+)?shipping', '', caption)
|
||||
caption = re.sub(r'(free\s)?download(\sfree)?', '', caption)
|
||||
caption = re.sub(r'\bclick\b\s(?:for|on)\s\w+', '', caption)
|
||||
caption = re.sub(r'\b(?:png|jpg|jpeg|bmp|webp|eps|pdf|apk|mp4)(\simage[s]?)?', '', caption)
|
||||
caption = re.sub(r'\bpage\s+\d+\b', '', caption)
|
||||
|
||||
caption = re.sub(r'\b\d*[a-zA-Z]+\d+[a-zA-Z]+\d+[a-zA-Z\d]*\b', r' ', caption) # j2d1a2a...
|
||||
|
||||
caption = re.sub(r'\b\d+\.?\d*[xх×]\d+\.?\d*\b', '', caption)
|
||||
|
||||
caption = re.sub(r'\b\s+\:\s+', r': ', caption)
|
||||
caption = re.sub(r'(\D[,\./])\b', r'\1 ', caption)
|
||||
caption = re.sub(r'\s+', ' ', caption)
|
||||
|
||||
caption.strip()
|
||||
|
||||
caption = re.sub(r'^[\"\']([\w\W]+)[\"\']$', r'\1', caption)
|
||||
caption = re.sub(r'^[\'\_,\-\:;]', r'', caption)
|
||||
caption = re.sub(r'[\'\_,\-\:\-\+]$', r'', caption)
|
||||
caption = re.sub(r'^\.\S+$', '', caption)
|
||||
|
||||
return caption.strip()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
# caption = re.sub(r'[\u4e00-\u9fff]+', '', caption)
|
||||
a = "امرأة مسنة بشعر أبيض ووجه مليء بالتجاعيد تجلس داخل سيارة قديمة الطراز، تنظر من خلال النافذة الجانبية بتعبير تأملي أو حزين قليلاً."
|
||||
print(a)
|
||||
print(text_preprocessing(a))
|
||||
|
||||
@@ -0,0 +1,274 @@
|
||||
|
||||
|
||||
from typing import Optional, Union, List
|
||||
import numpy as np
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
from fastvideo.utils.communications import all_gather
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from fastvideo.model.pipeline_mochi import linear_quadratic_schedule, retrieve_timesteps
|
||||
from tqdm import tqdm
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from diffusers import (
|
||||
FlowMatchEulerDiscreteScheduler,
|
||||
AutoencoderKLMochi,
|
||||
)
|
||||
from fastvideo.utils.logging import main_print
|
||||
from fastvideo.distill.solver import PCMFMScheduler
|
||||
from diffusers.utils import export_to_video
|
||||
import os
|
||||
import wandb
|
||||
import gc
|
||||
def prepare_latents(
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
num_frames,
|
||||
dtype,
|
||||
device,
|
||||
generator,
|
||||
vae_spatial_scale_factor,
|
||||
vae_temporal_scale_factor,
|
||||
):
|
||||
height = height // vae_spatial_scale_factor
|
||||
width = width // vae_spatial_scale_factor
|
||||
num_frames = (num_frames - 1) // vae_temporal_scale_factor + 1
|
||||
|
||||
shape = (batch_size, num_channels_latents, num_frames, height, width)
|
||||
|
||||
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
return latents
|
||||
|
||||
def sample_validation_video(
|
||||
transformer,
|
||||
vae,
|
||||
scheduler,
|
||||
scheduler_type="euler",
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
num_frames: int = 16,
|
||||
num_inference_steps: int = 28,
|
||||
timesteps: List[int] = None,
|
||||
guidance_scale: float = 4.5,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
prompt_attention_mask: Optional[torch.Tensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.Tensor] = None,
|
||||
negative_prompt_attention_mask: Optional[torch.Tensor] = None,
|
||||
output_type: Optional[str] = "pil",
|
||||
vae_spatial_scale_factor = 8,
|
||||
vae_temporal_scale_factor = 6,
|
||||
):
|
||||
device = vae.device
|
||||
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
do_classifier_free_guidance = guidance_scale > 1.0
|
||||
if do_classifier_free_guidance:
|
||||
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
prompt_attention_mask = torch.cat([negative_prompt_attention_mask, prompt_attention_mask], dim=0)
|
||||
|
||||
# 4. Prepare latent variables
|
||||
# TODO: Remove hardcore
|
||||
num_channels_latents = 12
|
||||
latents = prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
num_frames,
|
||||
prompt_embeds.dtype,
|
||||
device,
|
||||
generator,
|
||||
vae_spatial_scale_factor,
|
||||
vae_temporal_scale_factor
|
||||
)
|
||||
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
|
||||
if get_sequence_parallel_state():
|
||||
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
|
||||
latents = latents[:, :, rank, :, :, :]
|
||||
|
||||
|
||||
# 5. Prepare timestep
|
||||
# from https://github.com/genmoai/models/blob/075b6e36db58f1242921deff83a1066887b9c9e1/src/mochi_preview/infer.py#L77
|
||||
threshold_noise = 0.025
|
||||
sigmas = linear_quadratic_schedule(num_inference_steps, threshold_noise)
|
||||
sigmas = np.array(sigmas)
|
||||
if scheduler_type == "euler":
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
timesteps,
|
||||
sigmas,
|
||||
)
|
||||
else:
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
)
|
||||
num_warmup_steps = max(len(timesteps) - num_inference_steps * scheduler.order, 0)
|
||||
|
||||
# 6. Denoising loop
|
||||
# with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
# write with tqdm instead
|
||||
# only enable if nccl_info.global_rank == 0
|
||||
|
||||
with tqdm(total=num_inference_steps, disable= nccl_info.rank_within_group != 0, desc="Validation sampling...") as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
|
||||
|
||||
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latent_model_input.shape[0]).to(latents.dtype)
|
||||
noise_pred = transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
# Mochi CFG + Sampling runs in FP32
|
||||
noise_pred = noise_pred.to(torch.float32)
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents_dtype = latents.dtype
|
||||
latents = scheduler.step(noise_pred, t, latents.to(torch.float32), return_dict=False)[0]
|
||||
latents = latents.to(latents_dtype)
|
||||
|
||||
if latents.dtype != latents_dtype:
|
||||
if torch.backends.mps.is_available():
|
||||
# some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272
|
||||
latents = latents.to(latents_dtype)
|
||||
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
|
||||
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
latents = all_gather(latents, dim=2)
|
||||
|
||||
|
||||
if output_type == "latent":
|
||||
video = latents
|
||||
else:
|
||||
# unscale/denormalize the latents
|
||||
# denormalize with the mean and std if available and not None
|
||||
has_latents_mean = hasattr(vae.config, "latents_mean") and vae.config.latents_mean is not None
|
||||
has_latents_std = hasattr(vae.config, "latents_std") and vae.config.latents_std is not None
|
||||
if has_latents_mean and has_latents_std:
|
||||
latents_mean = (
|
||||
torch.tensor(vae.config.latents_mean).view(1, 12, 1, 1, 1).to(latents.device, latents.dtype)
|
||||
)
|
||||
latents_std = (
|
||||
torch.tensor(vae.config.latents_std).view(1, 12, 1, 1, 1).to(latents.device, latents.dtype)
|
||||
)
|
||||
latents = latents * latents_std / vae.config.scaling_factor + latents_mean
|
||||
else:
|
||||
latents = latents / vae.config.scaling_factor
|
||||
|
||||
video = vae.decode(latents, return_dict=False)[0]
|
||||
video_processor = VideoProcessor(vae_scale_factor=vae_spatial_scale_factor)
|
||||
video = video_processor.postprocess_video(video, output_type=output_type)
|
||||
|
||||
|
||||
|
||||
return (video,)
|
||||
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.autocast("cuda", dtype=torch.bfloat16)
|
||||
def log_validation(args, transformer, device, weight_dtype, global_step, scheduler_type="euler",shift=1.0, num_euler_timesteps=100, linear_quadratic_threshold=0.025, linear_range=0.5, ema=False):
|
||||
#TODO
|
||||
print(f"Running validation....\n")
|
||||
vae = AutoencoderKLMochi.from_pretrained(args.pretrained_model_name_or_path, subfolder="vae", torch_dtype=weight_dtype).to("cuda")
|
||||
vae.enable_tiling()
|
||||
if scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
else:
|
||||
linear_quadraic = True if scheduler_type == "pcm_linear_quadratic" else False
|
||||
scheduler = PCMFMScheduler(1000, shift, num_euler_timesteps, linear_quadraic, linear_quadratic_threshold, linear_range)
|
||||
# args.validation_prompt_dir
|
||||
|
||||
validation_guidance_scale_ls = args.validation_guidance_scale.split(",")
|
||||
validation_guidance_scale_ls = [float(scale) for scale in validation_guidance_scale_ls]
|
||||
for validation_sampling_step in args.validation_sampling_steps.split(","):
|
||||
validation_sampling_step = int(validation_sampling_step)
|
||||
for validation_guidance_scale in validation_guidance_scale_ls:
|
||||
|
||||
videos = []
|
||||
# prompt_embed are named embed0 to embedN
|
||||
# check how many embeds are there
|
||||
num_embeds = len([f for f in os.listdir(args.validation_prompt_dir) if "embed" in f])
|
||||
validation_prompt_ids = list(range(num_embeds))
|
||||
num_sp_groups = int(os.getenv("WORLD_SIZE", '1')) // nccl_info.sp_size
|
||||
# pad to multiple of groups
|
||||
validation_prompt_ids += [0] * (num_sp_groups - num_embeds % num_sp_groups)
|
||||
num_embeds_per_group = len(validation_prompt_ids) // num_sp_groups
|
||||
local_prompt_ids = validation_prompt_ids[nccl_info.group_id * num_embeds_per_group: (nccl_info.group_id + 1) * num_embeds_per_group]
|
||||
|
||||
for i in local_prompt_ids:
|
||||
prompt_embed_path = os.path.join(args.validation_prompt_dir, f"embed{i}.pt")
|
||||
prompt_mask_path = os.path.join(args.validation_prompt_dir, f"mask{i}.pt")
|
||||
prompt_embeds = torch.load(prompt_embed_path, map_location="cpu", weights_only=True).to(device).to(weight_dtype).unsqueeze(0)
|
||||
prompt_attention_mask = torch.load(prompt_mask_path, map_location="cpu", weights_only=True).to(device).to(weight_dtype).unsqueeze(0)
|
||||
negative_prompt_embeds = torch.zeros(256, 4096).to(device).to(weight_dtype).unsqueeze(0)
|
||||
negative_prompt_attention_mask = torch.zeros(256).bool().to(device).unsqueeze(0)
|
||||
generator = torch.Generator(device="cuda").manual_seed(12345)
|
||||
video = sample_validation_video(
|
||||
transformer,
|
||||
vae,
|
||||
scheduler,
|
||||
scheduler_type=scheduler_type,
|
||||
num_frames=args.num_frames,
|
||||
# Peiyuan TODO: remove hardcode
|
||||
height=480,
|
||||
width=848,
|
||||
num_inference_steps=validation_sampling_step,
|
||||
guidance_scale=validation_guidance_scale,
|
||||
generator=generator,
|
||||
prompt_embeds = prompt_embeds,
|
||||
prompt_attention_mask = prompt_attention_mask,
|
||||
negative_prompt_embeds = negative_prompt_embeds,
|
||||
negative_prompt_attention_mask = negative_prompt_attention_mask,
|
||||
)[0]
|
||||
if nccl_info.rank_within_group == 0:
|
||||
videos.append(video[0])
|
||||
# collect videos from all process to process zero
|
||||
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
# log if main process
|
||||
torch.distributed.barrier()
|
||||
all_videos = [None for i in range(int(os.getenv("WORLD_SIZE", '1')))] # remove padded videos
|
||||
torch.distributed.all_gather_object(all_videos, videos)
|
||||
if nccl_info.global_rank == 0:
|
||||
# remove padding
|
||||
videos = [video for videos in all_videos for video in videos]
|
||||
videos = videos[:num_embeds]
|
||||
# linearize all videos
|
||||
video_filenames = []
|
||||
for i, video in enumerate(videos):
|
||||
filename = os.path.join(args.output_dir, f"validation_step_{global_step}_sample_{validation_sampling_step}_guidance_{validation_guidance_scale}_video_{i}.mp4")
|
||||
export_to_video(video, filename, fps=30)
|
||||
video_filenames.append(filename)
|
||||
|
||||
logs = {
|
||||
f"{'ema_' if ema else ''}validation_sample_{validation_sampling_step}_guidance_{validation_guidance_scale}": [
|
||||
wandb.Video(filename)
|
||||
for i, filename in enumerate(video_filenames)
|
||||
]
|
||||
}
|
||||
wandb.log(logs, step=global_step)
|
||||
|
||||
@@ -1,31 +0,0 @@
|
||||
"""fastvideo2 — a post-training-to-serving substrate for video models, MVP.
|
||||
|
||||
Four surfaces (see README.md at the repo root):
|
||||
contracts fastvideo2.card / pipeline / loop — frozen data cards, enforced
|
||||
stage edges, the driven-loop protocol
|
||||
reference fastvideo2.<family>.reference — the standalone eager oracle
|
||||
verifier fastvideo2.verify — tiered gates + evidence ledger
|
||||
trace engine identity chain — request/stage/loop.step -> NVTX
|
||||
|
||||
``import fastvideo2`` is dependency-light: torch / diffusers / transformers
|
||||
load lazily, only when weights are actually touched.
|
||||
"""
|
||||
from fastvideo2.card import ModelCard, derive
|
||||
from fastvideo2.engine import Instance, Output, Request, run
|
||||
from fastvideo2.registry import resolve
|
||||
from fastvideo2.sdk import Model, Result, load
|
||||
|
||||
__version__ = "0.1.0"
|
||||
__all__ = ["ModelCard", "derive", "Model", "Result", "load", "resolve",
|
||||
"Instance", "Output", "Request", "run", "generate", "__version__"]
|
||||
|
||||
|
||||
def generate(model: str, prompt: str, *, root: str | None = None,
|
||||
device: str | None = None, **request_kwargs) -> Result:
|
||||
"""One-call convenience over the SDK: load then generate (loads per call —
|
||||
hold a :class:`Model` via :func:`load` to amortize residency).
|
||||
|
||||
>>> result = fastvideo2.generate("wan2.1-t2v-1.3b", "a cat surfing", seed=7)
|
||||
>>> result.video.shape # [T, H, W, C] uint8
|
||||
"""
|
||||
return load(model, root=root, device=device).generate(prompt, **request_kwargs)
|
||||
@@ -1,3 +0,0 @@
|
||||
from fastvideo2.cli import main
|
||||
|
||||
raise SystemExit(main())
|
||||
@@ -1,213 +0,0 @@
|
||||
"""Model cards — the contract surface.
|
||||
|
||||
A card is a frozen, pure-data description of one servable artifact: its
|
||||
components, the loops its weights assume, and the sampling defaults that are
|
||||
part of the trained artifact. Components and loops are declared as
|
||||
``"module:attr"`` reference strings and loop params must be plain JSON values
|
||||
— no callables — so a card is:
|
||||
|
||||
* **serializable** — ``to_json``/``from_dict`` are lossless (T0-gated), so a
|
||||
card ships as ``card.json`` beside a checkpoint or inside a deploy config;
|
||||
* **content-addressed** — ``digest()`` hashes the canonical JSON (think git
|
||||
object id) and is the card's identity in evidence records, T1 baselines,
|
||||
and environment manifests. It names the declaration only; weights, code,
|
||||
and environment drift are checked separately (T1, T2, env fingerprint).
|
||||
|
||||
Variants are expressed as diffs against a base card via :func:`derive` — never
|
||||
as builder functions with keyword arguments. Derivation is additive: a variant
|
||||
that needs to *remove* something picked the wrong base and should be declared
|
||||
fresh.
|
||||
|
||||
Import discipline: this module is stdlib-only. ``validate()`` imports the
|
||||
declared loop modules to check their ``semantics`` ids, so loop modules must be
|
||||
importable without torch.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import importlib
|
||||
import json
|
||||
from dataclasses import asdict, dataclass, field, fields, is_dataclass, replace
|
||||
from typing import Any
|
||||
|
||||
|
||||
class CardError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
def resolve_ref(ref: str) -> Any:
|
||||
"""Resolve a ``"module:attr"`` reference string to the live object."""
|
||||
mod, _, attr = ref.partition(":")
|
||||
if not mod or not attr:
|
||||
raise CardError(f"bad reference {ref!r} (expected 'module:attr')")
|
||||
return getattr(importlib.import_module(mod), attr)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ComponentSpec:
|
||||
"""One weight-bearing (or processing) component of the artifact."""
|
||||
component_id: str
|
||||
kind: str # dit | vae | text_encoder | tokenizer
|
||||
module: str # loader reference, e.g. "fastvideo2.wan21.model:WanModel"
|
||||
subfolder: str # subfolder in the checkpoint layout ("" = repo root)
|
||||
dtype: str = "bf16" # bf16 | fp32 | "" (dtype-less, e.g. tokenizer)
|
||||
source: str = "" # weights repo override; "" = the card-level `weights`
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class LoopSpec:
|
||||
"""One iterative computation the card can run.
|
||||
|
||||
``loop`` names the implementation class; ``params`` are plain JSON values
|
||||
passed to its constructor. The class carries a ``semantics`` id that
|
||||
provenance pins (see ``Provenance.assumes_loop``).
|
||||
"""
|
||||
loop_id: str
|
||||
loop: str # "fastvideo2.wan21.loop:WanDenoiseLoop"
|
||||
params: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SamplingDefaults:
|
||||
"""Per-model generation defaults that are part of the trained artifact."""
|
||||
num_steps: int
|
||||
guidance_scale: float
|
||||
height: int
|
||||
width: int
|
||||
num_frames: int
|
||||
fps: int
|
||||
shift: float
|
||||
negative_prompt: str = ""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Provenance:
|
||||
"""Where the weights came from and what they assume.
|
||||
|
||||
``assumes_loop`` is a *semantics id* (e.g. ``"wan.flow_euler.cfg/v1"``),
|
||||
not a loop_id: validation resolves every declared loop class and requires
|
||||
one whose ``semantics`` matches. A distilled student that requires a
|
||||
different sampler therefore cannot validate against a base card.
|
||||
``substitution`` classifies this artifact relative to ``parents``:
|
||||
``exact`` | ``bounded`` | ``quality-changing``.
|
||||
"""
|
||||
method: str = "base"
|
||||
parents: tuple[str, ...] = ()
|
||||
assumes_loop: str = ""
|
||||
precision: str = "bf16"
|
||||
substitution: str = "exact"
|
||||
tolerances: dict[str, float] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelCard:
|
||||
model_id: str
|
||||
family: str
|
||||
weights: str # canonical source (HF repo id)
|
||||
components: dict[str, ComponentSpec]
|
||||
loops: dict[str, LoopSpec]
|
||||
capabilities: tuple[str, ...]
|
||||
provenance: Provenance
|
||||
sampling_defaults: SamplingDefaults
|
||||
determinism: str = "tolerance" # bitwise | tolerance
|
||||
|
||||
# --- identity ---------------------------------------------------------- #
|
||||
def to_dict(self) -> dict:
|
||||
return asdict(self)
|
||||
|
||||
def to_json(self) -> str:
|
||||
return json.dumps(self.to_dict(), sort_keys=True, indent=2)
|
||||
|
||||
def digest(self) -> str:
|
||||
"""Content digest over the canonical JSON — the card's identity in the
|
||||
evidence ledger and every environment manifest."""
|
||||
canon = json.dumps(self.to_dict(), sort_keys=True, separators=(",", ":"))
|
||||
return hashlib.sha256(canon.encode()).hexdigest()[:16]
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, d: dict) -> "ModelCard":
|
||||
return cls(
|
||||
model_id=d["model_id"],
|
||||
family=d["family"],
|
||||
weights=d["weights"],
|
||||
components={k: ComponentSpec(**v) for k, v in d["components"].items()},
|
||||
loops={k: LoopSpec(**v) for k, v in d["loops"].items()},
|
||||
capabilities=tuple(d["capabilities"]),
|
||||
provenance=Provenance(**{**d["provenance"], "parents": tuple(d["provenance"]["parents"])}),
|
||||
sampling_defaults=SamplingDefaults(**d["sampling_defaults"]),
|
||||
determinism=d.get("determinism", "tolerance"),
|
||||
)
|
||||
|
||||
# --- validation -------------------------------------------------------- #
|
||||
def validate(self) -> "ModelCard":
|
||||
errs: list[str] = []
|
||||
if not self.components:
|
||||
errs.append("card declares no components")
|
||||
if not self.loops:
|
||||
errs.append("card declares no loops")
|
||||
for cid, spec in self.components.items():
|
||||
if cid != spec.component_id:
|
||||
errs.append(f"component key {cid!r} != component_id {spec.component_id!r}")
|
||||
semantics_seen: list[str] = []
|
||||
for lid, spec in self.loops.items():
|
||||
if lid != spec.loop_id:
|
||||
errs.append(f"loop key {lid!r} != loop_id {spec.loop_id!r}")
|
||||
try:
|
||||
cls = resolve_ref(spec.loop)
|
||||
except Exception as e: # unresolvable ref is a contract violation
|
||||
errs.append(f"loop {lid!r}: cannot resolve {spec.loop!r} ({e})")
|
||||
continue
|
||||
sem = getattr(cls, "semantics", None)
|
||||
if not sem:
|
||||
errs.append(f"loop {lid!r}: class {spec.loop!r} declares no `semantics` id")
|
||||
else:
|
||||
semantics_seen.append(sem)
|
||||
try:
|
||||
json.dumps(spec.params)
|
||||
except TypeError:
|
||||
errs.append(f"loop {lid!r}: params are not plain JSON values")
|
||||
# the teeth: weights may only be served under a loop whose semantics
|
||||
# they were trained for.
|
||||
if self.provenance.assumes_loop and self.provenance.assumes_loop not in semantics_seen:
|
||||
errs.append(
|
||||
f"provenance.assumes_loop={self.provenance.assumes_loop!r} matches no declared "
|
||||
f"loop semantics (have {semantics_seen}) — these weights cannot run on this card")
|
||||
if self.determinism not in ("bitwise", "tolerance"):
|
||||
errs.append(f"unknown determinism class {self.determinism!r}")
|
||||
if errs:
|
||||
raise CardError(f"ModelCard {self.model_id!r} failed validation:\n - " + "\n - ".join(errs))
|
||||
return self
|
||||
|
||||
|
||||
def _merge_field(old: Any, patch: Any) -> Any:
|
||||
"""One-level structural merge used by :func:`derive`.
|
||||
|
||||
dict field + dict patch -> merge by key (spec values replace; dict
|
||||
values patch the existing spec/dict)
|
||||
dataclass field + dict patch -> replace() with recursively merged fields
|
||||
anything else -> the patch value wins
|
||||
"""
|
||||
if is_dataclass(old) and isinstance(patch, dict):
|
||||
merged = {k: _merge_field(getattr(old, k), v) for k, v in patch.items()}
|
||||
return replace(old, **merged)
|
||||
if isinstance(old, dict) and isinstance(patch, dict):
|
||||
out = dict(old)
|
||||
for k, v in patch.items():
|
||||
out[k] = _merge_field(old[k], v) if k in old else v
|
||||
return out
|
||||
return patch
|
||||
|
||||
|
||||
def derive(base: ModelCard, **delta: Any) -> ModelCard:
|
||||
"""A variant as an explicit diff against a base card.
|
||||
|
||||
Additive only: keys merge, nothing is deleted. A variant that must remove a
|
||||
component or loop is a different architecture — declare it fresh. The
|
||||
derived card re-validates, so an invalid diff fails at declaration.
|
||||
"""
|
||||
valid = {f.name for f in fields(ModelCard)}
|
||||
unknown = set(delta) - valid
|
||||
if unknown:
|
||||
raise CardError(f"derive: unknown card fields {sorted(unknown)}")
|
||||
merged = {k: _merge_field(getattr(base, k), v) for k, v in delta.items()}
|
||||
return replace(base, **merged).validate()
|
||||
@@ -1,92 +0,0 @@
|
||||
"""CLI: describe / generate / verify — the three agent-facing verbs.
|
||||
|
||||
python -m fastvideo2 describe wan2.1-t2v-1.3b
|
||||
python -m fastvideo2 generate wan2.1-t2v-1.3b --prompt "a cat surfing" --out cat.mp4
|
||||
python -m fastvideo2 verify wan2.1-t2v-1.3b --tier 2 [--bless]
|
||||
|
||||
``describe`` prints the card as JSON plus its digest — machine-readable
|
||||
capability discovery. ``verify`` appends typed results to the evidence ledger
|
||||
and exits non-zero on any failed gate.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import sys
|
||||
|
||||
|
||||
def _describe(args) -> int:
|
||||
from fastvideo2.registry import resolve
|
||||
card, _ = resolve(args.model)
|
||||
print(card.to_json())
|
||||
print(f'// digest: {card.digest()}', file=sys.stderr)
|
||||
return 0
|
||||
|
||||
|
||||
def _generate(args) -> int:
|
||||
import fastvideo2
|
||||
kwargs = {k: getattr(args, k) for k in
|
||||
("seed", "num_steps", "guidance_scale", "height", "width", "num_frames", "shift")
|
||||
if getattr(args, k) is not None}
|
||||
model = fastvideo2.load(args.model, root=args.root, device=args.device)
|
||||
result = model.generate(args.prompt, **kwargs)
|
||||
result.save(args.out, fps=args.fps)
|
||||
steps = [t for t in result.trace if "/denoise." in t["label"]]
|
||||
print(f"video {result.video.shape} -> {args.out}")
|
||||
print(f"total {result.seconds:.1f}s; denoise steps {len(steps)}, "
|
||||
f"mean {sum(t['seconds'] for t in steps) / max(len(steps), 1):.2f}s/step")
|
||||
return 0
|
||||
|
||||
|
||||
def _verify(args) -> int:
|
||||
from fastvideo2.verify import LEDGER, verify
|
||||
results = verify(args.model, tier=args.tier, root=args.root, device=args.device,
|
||||
bless=args.bless, anchor=args.anchor)
|
||||
for r in results:
|
||||
mark = {"pass": "PASS ", "blessed": "BLESS", "fail": "FAIL "}[r.status]
|
||||
print(f" {mark} {r.gate:14s} {r.detail or json.dumps(r.metrics)[:120]}")
|
||||
print(f"ledger: {LEDGER}")
|
||||
return 0 if all(r.ok for r in results) else 1
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
p = argparse.ArgumentParser(prog="fastvideo2", description=__doc__)
|
||||
sub = p.add_subparsers(dest="cmd", required=True)
|
||||
|
||||
d = sub.add_parser("describe", help="print a card as JSON + digest")
|
||||
d.add_argument("model")
|
||||
d.set_defaults(fn=_describe)
|
||||
|
||||
g = sub.add_parser("generate", help="run one request, save an mp4")
|
||||
g.add_argument("model")
|
||||
g.add_argument("--prompt", required=True)
|
||||
g.add_argument("--out", default="out.mp4")
|
||||
g.add_argument("--root", default=None, help="local checkpoint dir (else HF cache)")
|
||||
g.add_argument("--device", default=None)
|
||||
g.add_argument("--seed", type=int, default=0)
|
||||
g.add_argument("--num-steps", dest="num_steps", type=int, default=None)
|
||||
g.add_argument("--guidance-scale", dest="guidance_scale", type=float, default=None)
|
||||
g.add_argument("--height", type=int, default=None)
|
||||
g.add_argument("--width", type=int, default=None)
|
||||
g.add_argument("--num-frames", dest="num_frames", type=int, default=None)
|
||||
g.add_argument("--shift", type=float, default=None)
|
||||
g.add_argument("--fps", type=int, default=16)
|
||||
g.set_defaults(fn=_generate)
|
||||
|
||||
v = sub.add_parser("verify", help="run tiered gates; append to the evidence ledger")
|
||||
v.add_argument("model")
|
||||
v.add_argument("--tier", type=int, default=3, choices=(0, 1, 2, 3))
|
||||
v.add_argument("--root", default=None)
|
||||
v.add_argument("--device", default=None)
|
||||
v.add_argument("--bless", action="store_true",
|
||||
help="write the T1 fingerprint baseline for this environment")
|
||||
v.add_argument("--anchor", action="store_true",
|
||||
help="also certify components against the official Wan2.1 goldens")
|
||||
v.set_defaults(fn=_verify)
|
||||
|
||||
args = p.parse_args(argv)
|
||||
return args.fn(args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -1,4 +0,0 @@
|
||||
"""Dreamverse runtime on fastvideo2 — session WS + fMP4 segments. See server.py."""
|
||||
from fastvideo2.dreamverse.server import build_app, main
|
||||
|
||||
__all__ = ["build_app", "main"]
|
||||
@@ -1,3 +0,0 @@
|
||||
from fastvideo2.dreamverse.server import main
|
||||
|
||||
main()
|
||||
@@ -1,100 +0,0 @@
|
||||
"""Dreamverse runtime anchor: boot the server with DUMMY prompt keys, drive
|
||||
one full session over the protocol, and assert the streaming contract:
|
||||
segment_start -> live step_complete x3 -> media_init -> binary fMP4 chunks
|
||||
(first chunk carries an ISO-BMFF `ftyp` box) -> media_segment_complete ->
|
||||
segment_complete with a latents sha.
|
||||
|
||||
Usage (cluster): python -m fastvideo2.dreamverse.gates.dreamverse_anchor
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
import urllib.request
|
||||
|
||||
MODEL = "fastwan-qad-fp8-1.3b"
|
||||
PORT = 8019
|
||||
PROMPT = "A raccoon in a field of sunflowers, warm light, mid-shot."
|
||||
|
||||
|
||||
def main() -> int:
|
||||
env = dict(os.environ, CEREBRAS_API_KEY="dummy", GROQ_API_KEY="dummy")
|
||||
server = subprocess.Popen(
|
||||
[sys.executable, "-m", "fastvideo2.dreamverse", "--model", MODEL,
|
||||
"--port", str(PORT)], env=env,
|
||||
stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
|
||||
try:
|
||||
for _ in range(360):
|
||||
try:
|
||||
with urllib.request.urlopen(
|
||||
f"http://127.0.0.1:{PORT}/health", timeout=5) as r:
|
||||
if json.loads(r.read())["model"] == MODEL:
|
||||
break
|
||||
except Exception:
|
||||
time.sleep(2)
|
||||
else:
|
||||
raise RuntimeError("server never became healthy")
|
||||
|
||||
import websockets
|
||||
|
||||
async def session() -> dict:
|
||||
counts = {"steps": 0, "chunks": 0, "ftyp": False}
|
||||
async with websockets.connect(
|
||||
f"ws://127.0.0.1:{PORT}/ws", max_size=None) as ws:
|
||||
await ws.send(json.dumps({"type": "session_init_v2",
|
||||
"enhancement": True}))
|
||||
for expected in ("queue_status", "gpu_assigned", "stream_start"):
|
||||
got = json.loads(await ws.recv())["type"]
|
||||
assert got == expected, (got, expected)
|
||||
await ws.send(json.dumps({"type": "segment_prompt_source",
|
||||
"prompt": PROMPT, "seed": 7}))
|
||||
while True:
|
||||
raw = await asyncio.wait_for(ws.recv(), timeout=900)
|
||||
if isinstance(raw, bytes):
|
||||
if counts["chunks"] == 0:
|
||||
counts["ftyp"] = b"ftyp" in raw[:64]
|
||||
counts["chunks"] += 1
|
||||
continue
|
||||
msg = json.loads(raw)
|
||||
t = msg["type"]
|
||||
if t == "step_complete":
|
||||
counts["steps"] += 1
|
||||
elif t == "segment_complete":
|
||||
counts["latents_sha"] = msg["latents_sha"]
|
||||
counts["frames"] = msg["frames"]
|
||||
break
|
||||
elif t == "error":
|
||||
raise RuntimeError(msg)
|
||||
await ws.send(json.dumps({"type": "leave"}))
|
||||
assert json.loads(await ws.recv())["type"] == "stream_complete"
|
||||
return counts
|
||||
|
||||
c = asyncio.run(session())
|
||||
finally:
|
||||
server.terminate()
|
||||
server.wait(timeout=30)
|
||||
|
||||
ok = (c["steps"] == 3 and c["chunks"] >= 1 and c["ftyp"]
|
||||
and c.get("frames", 0) == 81)
|
||||
print(f"steps={c['steps']} chunks={c['chunks']} ftyp={c['ftyp']} "
|
||||
f"frames={c.get('frames')} sha={c.get('latents_sha')} "
|
||||
f"{'OK' if ok else 'FAIL'}")
|
||||
|
||||
from fastvideo2.verify import GateResult, append_ledger, env_fingerprint
|
||||
append_ledger([GateResult(gate="anchor.dreamverse-runtime",
|
||||
status="pass" if ok else "fail", model_id=MODEL,
|
||||
card_digest="-",
|
||||
metrics={"steps": float(c["steps"]),
|
||||
"chunks": float(c["chunks"]),
|
||||
"ftyp": 1.0 if c["ftyp"] else 0.0},
|
||||
tolerances={}, env=env_fingerprint(),
|
||||
detail=f"latents {c.get('latents_sha')}")])
|
||||
return 0 if ok else 1
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -1,269 +0,0 @@
|
||||
"""Dreamverse runtime on fastvideo2 — the realtime video session server.
|
||||
|
||||
Ported from ``apps/dreamverse`` (fastvideo-main): the WebSocket session
|
||||
protocol (``session_init_v2`` → per-segment prompts → fMP4 fragments over
|
||||
the socket), the ffmpeg fragmented-MP4 encoder (verbatim flags from
|
||||
``entrypoints/streaming/stream.py``: libx264, zerolatency,
|
||||
``empty_moov+default_base_moof+frag_keyframe+faststart``), and an optional
|
||||
Cerebras/Groq prompt enhancer (boots with dummy keys — enhancement simply
|
||||
stays off, the same bring-up shortcut the GB200 deploys used).
|
||||
|
||||
Deliberately re-based for v2.1 (the original is LTX2-specific — audio,
|
||||
refine stage, continuation-state, LoRA stack): segments generate through the
|
||||
fastvideo2 SDK on the FastWan 3-step DMD student (seconds per segment on
|
||||
GB200), and per-step progress is LIVE via the engine's ``on_step`` hook
|
||||
(``step_complete`` per denoise step — the original emits one terminal event
|
||||
per segment). Message names follow the upstream protocol so their web client
|
||||
schema maps directly; LTX2-only fields are ignored.
|
||||
|
||||
Run: python -m fastvideo2.dreamverse --model fastwan-qad-fp8-1.3b --port 8009
|
||||
ffmpeg: FASTVIDEO_FFMPEG_BIN, or PATH, or the imageio-ffmpeg bundled binary.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
import threading
|
||||
import urllib.request
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
|
||||
def find_ffmpeg() -> str:
|
||||
p = os.environ.get("FASTVIDEO_FFMPEG_BIN") or shutil.which("ffmpeg")
|
||||
if p:
|
||||
return p
|
||||
try: # the GB200 bring-up shortcut: pip-installed bundled binary
|
||||
import imageio_ffmpeg
|
||||
return imageio_ffmpeg.get_ffmpeg_exe()
|
||||
except ImportError as e:
|
||||
raise RuntimeError("no ffmpeg (set FASTVIDEO_FFMPEG_BIN, install "
|
||||
"ffmpeg, or `pip install imageio-ffmpeg`)") from e
|
||||
|
||||
|
||||
def fmp4_encode(frames: Any, *, fps: int, ffmpeg: str) -> list[bytes]:
|
||||
"""One segment -> fragmented-MP4 byte chunks (upstream's exact flags)."""
|
||||
t, h, w, _ = frames.shape
|
||||
args = [ffmpeg, "-hide_banner", "-loglevel", "error",
|
||||
"-f", "rawvideo", "-pix_fmt", "rgb24", "-s", f"{w}x{h}",
|
||||
"-r", str(fps), "-i", "-",
|
||||
"-c:v", "libx264", "-preset", "ultrafast", "-tune", "zerolatency",
|
||||
"-pix_fmt", "yuv420p",
|
||||
"-movflags", "empty_moov+default_base_moof+frag_keyframe+faststart",
|
||||
"-f", "mp4", "-"]
|
||||
proc = subprocess.Popen(args, stdin=subprocess.PIPE, stdout=subprocess.PIPE,
|
||||
stderr=subprocess.DEVNULL, bufsize=0)
|
||||
out: list[bytes] = []
|
||||
|
||||
def _read() -> None:
|
||||
while True:
|
||||
chunk = proc.stdout.read(65536)
|
||||
if not chunk:
|
||||
break
|
||||
out.append(chunk)
|
||||
|
||||
reader = threading.Thread(target=_read, daemon=True)
|
||||
reader.start()
|
||||
for i in range(t):
|
||||
proc.stdin.write(frames[i].tobytes())
|
||||
proc.stdin.close()
|
||||
proc.wait(timeout=120)
|
||||
reader.join(timeout=30)
|
||||
return out
|
||||
|
||||
|
||||
class PromptEnhancer:
|
||||
"""Cerebras-or-Groq chat call (upstream's provider pair, gpt-oss-120b).
|
||||
Dummy/missing keys or any failure -> pass the prompt through unchanged."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.cerebras = os.environ.get("CEREBRAS_API_KEY", "")
|
||||
self.groq = os.environ.get("GROQ_API_KEY", "")
|
||||
self.enabled = any(k and k != "dummy" for k in (self.cerebras, self.groq))
|
||||
|
||||
def enhance(self, prompt: str, history: list[str]) -> str:
|
||||
if not self.enabled:
|
||||
return prompt
|
||||
targets = []
|
||||
if self.cerebras and self.cerebras != "dummy":
|
||||
targets.append(("https://api.cerebras.ai/v1/chat/completions",
|
||||
self.cerebras, "gpt-oss-120b"))
|
||||
if self.groq and self.groq != "dummy":
|
||||
targets.append(("https://api.groq.com/openai/v1/chat/completions",
|
||||
self.groq, "openai/gpt-oss-120b"))
|
||||
system = ("Rewrite the user's next-video-segment prompt into one vivid, "
|
||||
"concrete shot description. Prior segments: "
|
||||
+ " | ".join(history[-3:]))
|
||||
for url, key, model_name in targets:
|
||||
try:
|
||||
req = urllib.request.Request(
|
||||
url, method="POST",
|
||||
headers={"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json"},
|
||||
data=json.dumps({"model": model_name, "temperature": 1.0,
|
||||
"messages": [{"role": "system", "content": system},
|
||||
{"role": "user", "content": prompt}]
|
||||
}).encode())
|
||||
with urllib.request.urlopen(req, timeout=20) as r:
|
||||
return json.loads(r.read())["choices"][0]["message"]["content"].strip()
|
||||
except Exception:
|
||||
continue
|
||||
return prompt
|
||||
|
||||
|
||||
def build_app(model: Any) -> Any:
|
||||
from fastapi import FastAPI
|
||||
from starlette.routing import WebSocketRoute
|
||||
from starlette.websockets import WebSocketDisconnect
|
||||
|
||||
from fastvideo2.engine import Request
|
||||
from fastvideo2.engine import run as engine_run
|
||||
from fastvideo2.sdk import Result
|
||||
|
||||
app = FastAPI(title="dreamverse-fv2", version="0.1")
|
||||
ffmpeg = find_ffmpeg()
|
||||
enhancer = PromptEnhancer()
|
||||
gen_lock = threading.Lock()
|
||||
|
||||
@app.get("/health")
|
||||
def health() -> dict:
|
||||
return {"status": "ok", "model": model.model_id,
|
||||
"enhancer": enhancer.enabled, "ffmpeg": ffmpeg}
|
||||
|
||||
@app.get("/readyz")
|
||||
def readyz() -> dict:
|
||||
return {"ready": True}
|
||||
|
||||
async def ws_session(ws) -> None:
|
||||
await ws.accept()
|
||||
try:
|
||||
init = await ws.receive_json()
|
||||
except WebSocketDisconnect:
|
||||
return
|
||||
if init.get("type") != "session_init_v2":
|
||||
await ws.send_json({"type": "error", "code": "bad_init",
|
||||
"error": "expected session_init_v2"})
|
||||
await ws.close()
|
||||
return
|
||||
session_id = uuid.uuid4().hex[:12]
|
||||
enhancement_on = bool(init.get("enhancement", False)) and enhancer.enabled
|
||||
history: list[str] = []
|
||||
segment_idx = 0
|
||||
await ws.send_json({"type": "queue_status", "position": 0})
|
||||
await ws.send_json({"type": "gpu_assigned", "session_id": session_id})
|
||||
await ws.send_json({"type": "stream_start", "session_id": session_id,
|
||||
"model": model.model_id})
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
while True:
|
||||
try:
|
||||
msg = await ws.receive_json()
|
||||
except WebSocketDisconnect:
|
||||
return
|
||||
mtype = msg.get("type")
|
||||
if mtype == "leave":
|
||||
await ws.send_json({"type": "stream_complete",
|
||||
"segments": segment_idx})
|
||||
await ws.close()
|
||||
return
|
||||
if mtype == "enhancement_updated":
|
||||
enhancement_on = bool(msg.get("enabled")) and enhancer.enabled
|
||||
continue
|
||||
if mtype != "segment_prompt_source":
|
||||
await ws.send_json({"type": "error", "code": "bad_message",
|
||||
"error": f"unsupported type {mtype!r}"})
|
||||
continue
|
||||
|
||||
prompt = str(msg.get("prompt", ""))
|
||||
if not prompt:
|
||||
await ws.send_json({"type": "error", "code": "bad_prompt",
|
||||
"error": "prompt required"})
|
||||
continue
|
||||
if enhancement_on:
|
||||
prompt = await loop.run_in_executor(
|
||||
None, enhancer.enhance, prompt, history)
|
||||
history.append(prompt)
|
||||
await ws.send_json({"type": "segment_start", "segment": segment_idx,
|
||||
"prompt": prompt})
|
||||
|
||||
q: asyncio.Queue = asyncio.Queue()
|
||||
|
||||
def on_step(label: str, seconds: float, meta: dict) -> None:
|
||||
loop.call_soon_threadsafe(
|
||||
q.put_nowait, {"type": "step_complete", "label": label,
|
||||
"seconds": round(seconds, 4)})
|
||||
|
||||
def generate(p: str = prompt, seed: Any = msg.get("seed", 0),
|
||||
steps: Any = msg.get("num_steps")) -> None:
|
||||
try:
|
||||
req = Request(prompt=p, request_id=f"{session_id}-{segment_idx}",
|
||||
seed=int(seed), num_steps=steps)
|
||||
with gen_lock:
|
||||
out = engine_run(model.instance, model.pipeline, req,
|
||||
on_step=on_step)
|
||||
result = Result(outputs=out.outputs, trace=out.trace,
|
||||
request=req.resolve(model.card),
|
||||
model_id=model.model_id,
|
||||
card_digest=model.card.digest(),
|
||||
fps=model.card.sampling_defaults.fps)
|
||||
chunks = fmp4_encode(result.video, fps=result.fps,
|
||||
ffmpeg=ffmpeg)
|
||||
import torch
|
||||
sha = hashlib.sha256(result.latents.detach().to(
|
||||
torch.float32).cpu().numpy().tobytes()).hexdigest()[:16]
|
||||
loop.call_soon_threadsafe(
|
||||
q.put_nowait, {"__chunks": chunks, "latents_sha": sha,
|
||||
"frames": int(result.video.shape[0])})
|
||||
except Exception as e:
|
||||
loop.call_soon_threadsafe(
|
||||
q.put_nowait, {"type": "error", "code": "generation",
|
||||
"error": f"{type(e).__name__}: {e}"})
|
||||
|
||||
threading.Thread(target=generate, daemon=True).start()
|
||||
while True:
|
||||
ev = await q.get()
|
||||
if "__chunks" in ev:
|
||||
await ws.send_json({"type": "media_init",
|
||||
"segment": segment_idx,
|
||||
"mime": 'video/mp4; codecs="avc1"'})
|
||||
for chunk in ev["__chunks"]:
|
||||
await ws.send_bytes(chunk)
|
||||
await ws.send_json({"type": "media_segment_complete",
|
||||
"segment": segment_idx})
|
||||
await ws.send_json({"type": "segment_complete",
|
||||
"segment": segment_idx,
|
||||
"frames": ev["frames"],
|
||||
"latents_sha": ev["latents_sha"]})
|
||||
segment_idx += 1
|
||||
break
|
||||
await ws.send_json(ev)
|
||||
if ev.get("type") == "error":
|
||||
break
|
||||
|
||||
app.router.routes.append(WebSocketRoute("/ws", ws_session))
|
||||
return app
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> None:
|
||||
import argparse
|
||||
|
||||
import uvicorn
|
||||
|
||||
import fastvideo2 as fv2
|
||||
|
||||
p = argparse.ArgumentParser("fastvideo2.dreamverse")
|
||||
p.add_argument("--model", default="fastwan-qad-fp8-1.3b")
|
||||
p.add_argument("--host", default="127.0.0.1")
|
||||
p.add_argument("--port", type=int, default=8009)
|
||||
p.add_argument("--device", default=None)
|
||||
args = p.parse_args(argv)
|
||||
model = fv2.load(args.model, device=args.device)
|
||||
uvicorn.run(build_app(model), host=args.host, port=args.port)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,178 +0,0 @@
|
||||
"""The engine — drives a pipeline over one resident instance, one request at a
|
||||
time, with the identity chain attached.
|
||||
|
||||
Identity chain: every unit of work is named ``request/stage`` for one-shot
|
||||
stages and ``request/stage/loop.step`` for loop steps. The same name goes to
|
||||
(a) the returned trace (typed timings, machine-readable) and (b) NVTX ranges
|
||||
when CUDA is present — so Nsight correlates kernels to model-level identity
|
||||
with no extra instrumentation.
|
||||
|
||||
Deliberately absent (this is the one-shot MVP): queueing, admission, batching,
|
||||
sessions, cancellation. Sessions with forkable state are the next consumer of
|
||||
the loop contract, not a reason to grow this file now.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
from dataclasses import dataclass, field, replace
|
||||
from typing import Any
|
||||
|
||||
from fastvideo2.card import ModelCard
|
||||
from fastvideo2.loading import load_component, resolve_weights
|
||||
from fastvideo2.loop import LoopRunner, build_loop
|
||||
from fastvideo2.pipeline import ComponentStage, LoopStage, Pipeline, run_component_stage
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Request:
|
||||
"""One generation request. ``None`` fields resolve from the card's
|
||||
sampling defaults via :meth:`resolve`."""
|
||||
prompt: str
|
||||
request_id: str = "req0"
|
||||
negative_prompt: str | None = None
|
||||
seed: int = 0
|
||||
num_steps: int | None = None
|
||||
guidance_scale: float | None = None
|
||||
height: int | None = None
|
||||
width: int | None = None
|
||||
num_frames: int | None = None
|
||||
shift: float | None = None
|
||||
capture_trajectory: bool = False
|
||||
|
||||
def resolve(self, card: ModelCard) -> "Request":
|
||||
d = card.sampling_defaults
|
||||
fill = {
|
||||
"negative_prompt": d.negative_prompt,
|
||||
"num_steps": d.num_steps,
|
||||
"guidance_scale": d.guidance_scale,
|
||||
"height": d.height,
|
||||
"width": d.width,
|
||||
"num_frames": d.num_frames,
|
||||
"shift": d.shift,
|
||||
}
|
||||
patch = {k: v for k, v in fill.items() if getattr(self, k) is None}
|
||||
return replace(self, **patch)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Output:
|
||||
request_id: str
|
||||
outputs: dict[str, Any]
|
||||
trace: list[dict] = field(default_factory=list) # [{label, seconds, ...meta}]
|
||||
|
||||
@property
|
||||
def seconds(self) -> float:
|
||||
return sum(t["seconds"] for t in self.trace)
|
||||
|
||||
|
||||
class Instance:
|
||||
"""A resident, loaded card: components materialize lazily and are shared
|
||||
by reference; loops are built from the card's declared specs."""
|
||||
|
||||
def __init__(self, card: ModelCard, root: str | None = None, device: str = "cpu"):
|
||||
self.card = card
|
||||
self.device = device
|
||||
self.root = resolve_weights(card, root)
|
||||
self._components: dict[str, Any] = {}
|
||||
self._source_roots: dict[str, str] = {}
|
||||
self._loops: dict[str, Any] = {}
|
||||
|
||||
def component(self, component_id: str) -> Any:
|
||||
if component_id not in self._components:
|
||||
spec = self.card.components.get(component_id)
|
||||
if spec is None:
|
||||
raise KeyError(f"component {component_id!r} not declared on card {self.card.model_id!r}")
|
||||
root = self._source_root(spec.source) if spec.source else self.root
|
||||
self._components[component_id] = load_component(spec, root, self.device)
|
||||
return self._components[component_id]
|
||||
|
||||
def _source_root(self, source: str) -> str:
|
||||
"""Resolve a per-component weights source (e.g. the official-layout
|
||||
transformer repo) through the same snapshot cache as card weights."""
|
||||
if source not in self._source_roots:
|
||||
from huggingface_hub import snapshot_download
|
||||
self._source_roots[source] = snapshot_download(source)
|
||||
return self._source_roots[source]
|
||||
|
||||
def loop(self, loop_id: str) -> Any:
|
||||
if loop_id not in self._loops:
|
||||
spec = self.card.loops.get(loop_id)
|
||||
if spec is None:
|
||||
raise KeyError(f"loop {loop_id!r} not declared on card {self.card.model_id!r}")
|
||||
self._loops[loop_id] = build_loop(spec)
|
||||
return self._loops[loop_id]
|
||||
|
||||
|
||||
def load(card: ModelCard, root: str | None = None, device: str | None = None) -> Instance:
|
||||
"""The public entrypoint: card + weights root -> resident instance."""
|
||||
card.validate()
|
||||
if device is None:
|
||||
device = _detect_device()
|
||||
return Instance(card, root=root, device=device)
|
||||
|
||||
|
||||
def _detect_device() -> str:
|
||||
try:
|
||||
import torch
|
||||
if torch.cuda.is_available():
|
||||
return "cuda"
|
||||
except ImportError:
|
||||
pass
|
||||
return "cpu"
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _nvtx(name: str):
|
||||
"""NVTX range when CUDA is live; free otherwise."""
|
||||
pushed = False
|
||||
try:
|
||||
import torch
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.nvtx.range_push(name)
|
||||
pushed = True
|
||||
except ImportError:
|
||||
pass
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
if pushed:
|
||||
import torch
|
||||
torch.cuda.nvtx.range_pop()
|
||||
|
||||
|
||||
def run(instance: Instance, pipeline: Pipeline, request: Request,
|
||||
on_step: Any = None) -> Output:
|
||||
"""Run one request through a pipeline to completion. ``on_step(label,
|
||||
seconds, meta)``, when given, fires live after every loop step (serving
|
||||
streams progress through it)."""
|
||||
import time
|
||||
req = request.resolve(instance.card)
|
||||
trace: list[dict] = []
|
||||
slots: dict[str, Any] = {}
|
||||
for name in pipeline.inputs: # request-provided slots, by attribute name
|
||||
slots[name] = getattr(req, name)
|
||||
|
||||
for stage in pipeline.stages:
|
||||
chain = f"{req.request_id}/{stage.stage_id}"
|
||||
if isinstance(stage, ComponentStage):
|
||||
with _nvtx(chain):
|
||||
t0 = time.perf_counter()
|
||||
run_component_stage(stage, instance, slots, req)
|
||||
trace.append({"label": chain, "seconds": time.perf_counter() - t0})
|
||||
elif isinstance(stage, LoopStage):
|
||||
loop = instance.loop(stage.loop_id)
|
||||
|
||||
def observe(label: str, seconds: float, meta: dict, _chain: str = chain) -> None:
|
||||
trace.append({"label": f"{_chain}/{label}", "seconds": seconds, **meta})
|
||||
if on_step is not None:
|
||||
on_step(f"{_chain}/{label}", seconds, meta)
|
||||
|
||||
inputs = {k: slots[k] for k in stage.reads}
|
||||
with _nvtx(chain):
|
||||
runner = LoopRunner(loop, req, instance, inputs, observe=observe)
|
||||
slots[stage.writes[0]] = runner.run()
|
||||
else:
|
||||
raise TypeError(f"unknown stage kind {type(stage).__name__}")
|
||||
|
||||
outputs = {name: slots[slot] for name, slot in pipeline.outputs.items()}
|
||||
return Output(request_id=req.request_id, outputs=outputs, trace=trace)
|
||||
@@ -1,18 +0,0 @@
|
||||
# Evidence ledger
|
||||
|
||||
Typed verification records, committed with the code they vouch for.
|
||||
|
||||
- `ledger.jsonl` — append-only `GateResult` records: gate, status, card digest,
|
||||
metrics, tolerances, environment fingerprint, timestamp. Written only by
|
||||
`python -m fastvideo2 verify`; never edited by hand.
|
||||
- `<model_id>.fingerprints.json` — the blessed T1 component baseline for one
|
||||
card digest in one environment.
|
||||
- `sample_*.mp4` — eyeballable artifacts from full-scale runs (e.g.
|
||||
`sample_wan21_seed7.mp4`, 50 steps / 81 frames on GB200, byte-identical
|
||||
between the production pipeline and `reference.py` at the same seed). Re-bless deliberately (`verify --bless`)
|
||||
when the card or environment legitimately changes; a digest mismatch is a
|
||||
failure, not a skip.
|
||||
|
||||
Ownership rule: baselines and gate tolerances are human-owned. Agents run the
|
||||
gates and append evidence; they do not re-bless baselines to make a failure
|
||||
disappear.
|
||||
@@ -1,82 +0,0 @@
|
||||
# FastWan variants — bitwise alignment vs fastvideo-main
|
||||
|
||||
Authority for FastWan artifacts is **fastvideo main** (they were distilled in
|
||||
that stack); the alignment target was bit-exactness against main's own serving
|
||||
path, pinned to main's exposed knobs so the goldens measure the artifact, not
|
||||
the accelerator stack: `FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN` (QAD) /
|
||||
`VIDEO_SPARSE_ATTN` (VSA), FP8 per-tensor dynamic quant (QAD), no torch.compile,
|
||||
no FSDP, single GB200.
|
||||
|
||||
Goldens captured at main commit `c459a1897899ffcec3be7765534d81000b9bb9c1`;
|
||||
the vendored forward (`wan21/model_fv.py`) was read at `e3f47dc2de2d…` — all
|
||||
10 numerics-relevant source files verified byte-identical between the two.
|
||||
|
||||
## Final anchor results (all rows target 0.0 — bitwise)
|
||||
|
||||
| row | fastwan-qad-fp8-1.3b | fastwan-t2v-1.3b (VSA) |
|
||||
|---|---|---|
|
||||
| dit bf16 probes t∈{1000,757,522} | 0.0 / 0.0 / 0.0 | — |
|
||||
| dit fp8 probes t∈{1000,757,522} | 0.0 / 0.0 / 0.0 | — |
|
||||
| dit vsa probes t∈{1000,757,522} | — | 0.0 / 0.0 / 0.0 |
|
||||
| text_encoder (e2e + probe prompts) | 0.0 / 0.0 | 0.0 / 0.0 |
|
||||
| e2e step-1 / step-2 latent chain | 0.0 / 0.0 | 0.0 / 0.0 |
|
||||
| e2e final latents (81f, 480×832, 3 steps) | 0.0 | 0.0 |
|
||||
|
||||
Ledger: `anchor.fastwan-qad-main` and `anchor.fastwan-vsa-main`, both `pass`
|
||||
(card digests `1c8e6f7d1380552d`, `0bd5c7771e10ce44`).
|
||||
|
||||
## Root causes found by the gates (in discovery order)
|
||||
|
||||
1. **fp8 quantization is device-sensitive.** main converts weights to fp8 on
|
||||
the GPU (post-materialization); quantizing the *identical* bf16 weights on
|
||||
CPU produces different fp8 codes often enough to move a full forward by
|
||||
~4e-2 rel (fp8's coarse grid amplifies conversion-tie differences).
|
||||
Fix: `FP8Linear` defers quantization to first forward on the serving
|
||||
device (`layers/fp8.py`).
|
||||
2. **0-dim sigma tensors demote the renoise mixing to bf16.** In torch type
|
||||
promotion 0-dim tensors act as scalars, so `(1-σ)*x0 + σ*ε` with a 0-dim
|
||||
fp32 σ ran in bf16 (two roundings); main's `[B,1,1,1]` fp32 σ promotes the
|
||||
arithmetic to fp32 with one final bf16 cast. 3.2e-3 per step, compounding
|
||||
to 8.2e-2 over 3 steps. Fix: non-0-dim σ in `WanDMDLoop`.
|
||||
3. **main's DMD sigma table is NOT the one the code appears to prepare.**
|
||||
`DmdDenoisingStage.__init__` hardcodes a fresh internal
|
||||
`FlowMatchEulerDiscreteScheduler(shift=8.0)`; the pipeline scheduler that
|
||||
`TimestepPreparationStage.set_timesteps(n)` configured is never consulted,
|
||||
and the config `flow_shift` is ignored. Lookups run against the 1000-entry
|
||||
warped **init** table: σ(1000)=1.0, σ(757)=0.7567567, σ(522)=0.5217391
|
||||
(confirmed in the capture manifests). `dmd_inference_table` reproduces
|
||||
this exactly; a canary T0 test guards it.
|
||||
|
||||
Also confirmed en route: the CPU-generator RNG stream (initial fp32 draw +
|
||||
bf16 renoise draws), the fp64 x0 math, and the flash/dense forward at full
|
||||
81-frame geometry are each independently bitwise (triage decomposition in
|
||||
session evidence).
|
||||
|
||||
## Caveats
|
||||
|
||||
- Text parity holds for ASCII prompts; main's ftfy cleaning diverges from
|
||||
official's on CJK width-folding (see wan21 report — main measured 4.19e-1
|
||||
vs official on the Chinese negative prompt). FastWan cards reuse the wan21
|
||||
text stage; DMD uses no negative prompt. Add main's clean fn + a CJK golden
|
||||
before serving non-ASCII prompts against these cards.
|
||||
- VAE decode is not bitwise-gated (shared component; wan21 anchors cover it);
|
||||
golden videos are in the goldens dirs for SSIM-level comparison.
|
||||
- Committed goldens are trimmed (per-step model outputs and the reproducible
|
||||
step-0 input dropped); the full set regenerates via
|
||||
`gates/capture_fastvideo_main.py {qad,vsa}` — one command, pinned config.
|
||||
|
||||
## SFWan (self-forcing causal) — added 2026-07-23
|
||||
|
||||
`sfwan-t2v-1.3b` (wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers) anchored bitwise on
|
||||
the FIRST complete run: all **35/35 chunk-rollout forwards** (7 blocks x
|
||||
4 warped DMD steps + context pass) hash-match main's CausalDMDDenosingStage
|
||||
exactly, text 0.0, e2e final latents 0.0 (`anchor.sfwan-main` pass).
|
||||
|
||||
Causal-specific semantics vendored (each different from BOTH other Wan
|
||||
forwards): per-frame temb `[B, T_temb, 6, dim]`; ALL-bf16 modulation (no
|
||||
fp32 promotion anywhere); plain bf16 LayerNorms; fp64 RoPE multipliers at
|
||||
absolute positions (start_frame offsets); block-causal KV cache
|
||||
(21-frame global window, `.detach()` on writes — training rollout reuses
|
||||
this same module); cached text cross-attention; warp table =
|
||||
SelfForcingFlowMatchScheduler(shift 5, extra_one_step) rows
|
||||
`[1000, 937.5, 833.33, 625]` self-indexing their own sigmas.
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -1,51 +0,0 @@
|
||||
{
|
||||
"repo": "FastVideo/FastWan-QAD-FP8-1.3B",
|
||||
"snapshot": "3de0eec0e2562923d38a87344127a86a35a3c11d",
|
||||
"fastvideo_commit": "c459a1897899ffcec3be7765534d81000b9bb9c1",
|
||||
"fastvideo_src": "/mnt/FastVideo",
|
||||
"torch": "2.12.0+cu130",
|
||||
"flash_attn": "2.8.3",
|
||||
"python": "3.12.13",
|
||||
"gpu": "NVIDIA GB200",
|
||||
"attention_backend": "FLASH_ATTN",
|
||||
"quant": "FP8 per-tensor (dynamic act, post-load weight quant from bf16)",
|
||||
"vsa_sparsity": null,
|
||||
"seed": 1234,
|
||||
"e2e_prompt": "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest. The playful yet serene atmosphere is complemented by soft natural light filtering through the petals. Mid-shot, warm and cheerful tones.",
|
||||
"probe_prompt": "A cat and a dog baking a cake together in a kitchen.",
|
||||
"probe_timesteps": [
|
||||
1000,
|
||||
757,
|
||||
522
|
||||
],
|
||||
"probe_latent_bcfhw": [
|
||||
1,
|
||||
16,
|
||||
5,
|
||||
60,
|
||||
104
|
||||
],
|
||||
"e2e_latent_btchw": [
|
||||
1,
|
||||
21,
|
||||
16,
|
||||
60,
|
||||
104
|
||||
],
|
||||
"dmd_denoising_steps": [
|
||||
1000,
|
||||
757,
|
||||
522
|
||||
],
|
||||
"scheduler": {
|
||||
"class": "FlowMatchEulerDiscreteScheduler",
|
||||
"shift": 8.0,
|
||||
"table_len": 1000,
|
||||
"sigma_lookup": {
|
||||
"1000": 1.0,
|
||||
"757": 0.7567567229270935,
|
||||
"522": 0.52173912525177
|
||||
}
|
||||
},
|
||||
"notes": "no compile, no fsdp, single GPU; DmdDenoisingStage's INTERNAL scheduler (hardcoded shift 8.0) is the sigma authority; committed goldens trimmed: e2e_step0 dropped (input is the seeded draw, reproducible) and per-step outputs dropped (triage-only) \u2014 full set regenerable via capture_fastvideo_main.py"
|
||||
}
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -1,51 +0,0 @@
|
||||
{
|
||||
"repo": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
"snapshot": "25e7ed7f41fd8ce2fdd108688c65e8caf0ce3aef",
|
||||
"fastvideo_commit": "c459a1897899ffcec3be7765534d81000b9bb9c1",
|
||||
"fastvideo_src": "/mnt/FastVideo",
|
||||
"torch": "2.12.0+cu130",
|
||||
"flash_attn": "2.8.3",
|
||||
"python": "3.12.13",
|
||||
"gpu": "NVIDIA GB200",
|
||||
"attention_backend": "VIDEO_SPARSE_ATTN",
|
||||
"quant": "none (bf16, VSA sparsity 0.80)",
|
||||
"vsa_sparsity": 0.8,
|
||||
"seed": 1234,
|
||||
"e2e_prompt": "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest. The playful yet serene atmosphere is complemented by soft natural light filtering through the petals. Mid-shot, warm and cheerful tones.",
|
||||
"probe_prompt": "A cat and a dog baking a cake together in a kitchen.",
|
||||
"probe_timesteps": [
|
||||
1000,
|
||||
757,
|
||||
522
|
||||
],
|
||||
"probe_latent_bcfhw": [
|
||||
1,
|
||||
16,
|
||||
5,
|
||||
60,
|
||||
104
|
||||
],
|
||||
"e2e_latent_btchw": [
|
||||
1,
|
||||
21,
|
||||
16,
|
||||
60,
|
||||
104
|
||||
],
|
||||
"dmd_denoising_steps": [
|
||||
1000,
|
||||
757,
|
||||
522
|
||||
],
|
||||
"scheduler": {
|
||||
"class": "FlowMatchEulerDiscreteScheduler",
|
||||
"shift": 8.0,
|
||||
"table_len": 1000,
|
||||
"sigma_lookup": {
|
||||
"1000": 1.0,
|
||||
"757": 0.7567567229270935,
|
||||
"522": 0.52173912525177
|
||||
}
|
||||
},
|
||||
"notes": "no compile, no fsdp, single GPU; DmdDenoisingStage's INTERNAL scheduler (hardcoded shift 8.0) is the sigma authority; committed goldens trimmed: e2e_step0 dropped (input is the seeded draw, reproducible) and per-step outputs dropped (triage-only) \u2014 full set regenerable via capture_fastvideo_main.py"
|
||||
}
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -1,282 +0,0 @@
|
||||
[
|
||||
{
|
||||
"x_hash": "03ceb634cfbdc8a5",
|
||||
"out_hash": "658efa11e9a88e87",
|
||||
"t": [
|
||||
1000.0
|
||||
],
|
||||
"start": 0
|
||||
},
|
||||
{
|
||||
"x_hash": "b986baa8885d6cb9",
|
||||
"out_hash": "dd84656c50a4b92b",
|
||||
"t": [
|
||||
937.5
|
||||
],
|
||||
"start": 0
|
||||
},
|
||||
{
|
||||
"x_hash": "c08abf4de0fc5fc2",
|
||||
"out_hash": "4001c9596842592d",
|
||||
"t": [
|
||||
833.3333129882812
|
||||
],
|
||||
"start": 0
|
||||
},
|
||||
{
|
||||
"x_hash": "a27dd2144f6a3c91",
|
||||
"out_hash": "2a0f48499caa3ca5",
|
||||
"t": [
|
||||
625.0
|
||||
],
|
||||
"start": 0
|
||||
},
|
||||
{
|
||||
"x_hash": "29769211170d71a2",
|
||||
"out_hash": "a4fe211f3fda4514",
|
||||
"t": [
|
||||
0.0
|
||||
],
|
||||
"start": 0
|
||||
},
|
||||
{
|
||||
"x_hash": "d96a639192b5aa24",
|
||||
"out_hash": "4b1f9dd86d718170",
|
||||
"t": [
|
||||
1000.0
|
||||
],
|
||||
"start": 4680
|
||||
},
|
||||
{
|
||||
"x_hash": "13c644ac8937ac8e",
|
||||
"out_hash": "d804417cc5fcc794",
|
||||
"t": [
|
||||
937.5
|
||||
],
|
||||
"start": 4680
|
||||
},
|
||||
{
|
||||
"x_hash": "2a9fd4b8c4362a57",
|
||||
"out_hash": "3f996d20f583eb87",
|
||||
"t": [
|
||||
833.3333129882812
|
||||
],
|
||||
"start": 4680
|
||||
},
|
||||
{
|
||||
"x_hash": "c0bc578106311f69",
|
||||
"out_hash": "1a2bf4336c45ebc8",
|
||||
"t": [
|
||||
625.0
|
||||
],
|
||||
"start": 4680
|
||||
},
|
||||
{
|
||||
"x_hash": "cd0c54679935f14d",
|
||||
"out_hash": "3fcda958d74bf79a",
|
||||
"t": [
|
||||
0.0
|
||||
],
|
||||
"start": 4680
|
||||
},
|
||||
{
|
||||
"x_hash": "cd20f035b46125c5",
|
||||
"out_hash": "9a976bbb313fc9b9",
|
||||
"t": [
|
||||
1000.0
|
||||
],
|
||||
"start": 9360
|
||||
},
|
||||
{
|
||||
"x_hash": "3332d96e01ae4004",
|
||||
"out_hash": "ddb5cde97146b967",
|
||||
"t": [
|
||||
937.5
|
||||
],
|
||||
"start": 9360
|
||||
},
|
||||
{
|
||||
"x_hash": "1f02a448f8c27f80",
|
||||
"out_hash": "86a4ac2b3e6624f7",
|
||||
"t": [
|
||||
833.3333129882812
|
||||
],
|
||||
"start": 9360
|
||||
},
|
||||
{
|
||||
"x_hash": "c34398a75af94118",
|
||||
"out_hash": "e19727744ee978ee",
|
||||
"t": [
|
||||
625.0
|
||||
],
|
||||
"start": 9360
|
||||
},
|
||||
{
|
||||
"x_hash": "b8229bafed68f655",
|
||||
"out_hash": "8a07b4778c138a91",
|
||||
"t": [
|
||||
0.0
|
||||
],
|
||||
"start": 9360
|
||||
},
|
||||
{
|
||||
"x_hash": "570ff151d9d1444e",
|
||||
"out_hash": "1a66ace7a1d60b51",
|
||||
"t": [
|
||||
1000.0
|
||||
],
|
||||
"start": 14040
|
||||
},
|
||||
{
|
||||
"x_hash": "61bd37b3e6382f8b",
|
||||
"out_hash": "6705d1c5591a147b",
|
||||
"t": [
|
||||
937.5
|
||||
],
|
||||
"start": 14040
|
||||
},
|
||||
{
|
||||
"x_hash": "5aaccd0e0c06d59c",
|
||||
"out_hash": "e43adfb21fa1bc9b",
|
||||
"t": [
|
||||
833.3333129882812
|
||||
],
|
||||
"start": 14040
|
||||
},
|
||||
{
|
||||
"x_hash": "65de10d92f0d12bd",
|
||||
"out_hash": "800293d62b41f4ae",
|
||||
"t": [
|
||||
625.0
|
||||
],
|
||||
"start": 14040
|
||||
},
|
||||
{
|
||||
"x_hash": "637bc60ba3c0b7a1",
|
||||
"out_hash": "63fb4744b7d48b10",
|
||||
"t": [
|
||||
0.0
|
||||
],
|
||||
"start": 14040
|
||||
},
|
||||
{
|
||||
"x_hash": "ca0cccb30274219c",
|
||||
"out_hash": "2fe6ebe902961d57",
|
||||
"t": [
|
||||
1000.0
|
||||
],
|
||||
"start": 18720
|
||||
},
|
||||
{
|
||||
"x_hash": "d4e78ae4fa1bbc1d",
|
||||
"out_hash": "24694df9ba77685c",
|
||||
"t": [
|
||||
937.5
|
||||
],
|
||||
"start": 18720
|
||||
},
|
||||
{
|
||||
"x_hash": "2b8ddf173f274fc8",
|
||||
"out_hash": "1156a77610e6dc1b",
|
||||
"t": [
|
||||
833.3333129882812
|
||||
],
|
||||
"start": 18720
|
||||
},
|
||||
{
|
||||
"x_hash": "6bbd2511192162cc",
|
||||
"out_hash": "dcb87decbf4d126f",
|
||||
"t": [
|
||||
625.0
|
||||
],
|
||||
"start": 18720
|
||||
},
|
||||
{
|
||||
"x_hash": "b144c64592c830fd",
|
||||
"out_hash": "09b00b37c56dc714",
|
||||
"t": [
|
||||
0.0
|
||||
],
|
||||
"start": 18720
|
||||
},
|
||||
{
|
||||
"x_hash": "3f11cee5b1830f29",
|
||||
"out_hash": "ac56b4506da4e278",
|
||||
"t": [
|
||||
1000.0
|
||||
],
|
||||
"start": 23400
|
||||
},
|
||||
{
|
||||
"x_hash": "a7bc9435cd702102",
|
||||
"out_hash": "c238bb0f2fa1b957",
|
||||
"t": [
|
||||
937.5
|
||||
],
|
||||
"start": 23400
|
||||
},
|
||||
{
|
||||
"x_hash": "6d9a5dcc9af79188",
|
||||
"out_hash": "1b7759348b2e0a06",
|
||||
"t": [
|
||||
833.3333129882812
|
||||
],
|
||||
"start": 23400
|
||||
},
|
||||
{
|
||||
"x_hash": "42fd32612cfc36b8",
|
||||
"out_hash": "d12e11bb03d918bb",
|
||||
"t": [
|
||||
625.0
|
||||
],
|
||||
"start": 23400
|
||||
},
|
||||
{
|
||||
"x_hash": "a7583660de097f48",
|
||||
"out_hash": "90473cb353fbb62d",
|
||||
"t": [
|
||||
0.0
|
||||
],
|
||||
"start": 23400
|
||||
},
|
||||
{
|
||||
"x_hash": "e5c85230c382e2f4",
|
||||
"out_hash": "6ed5d13fa54d6f85",
|
||||
"t": [
|
||||
1000.0
|
||||
],
|
||||
"start": 28080
|
||||
},
|
||||
{
|
||||
"x_hash": "f60119f243f5bbb2",
|
||||
"out_hash": "52741bbb32a4fa39",
|
||||
"t": [
|
||||
937.5
|
||||
],
|
||||
"start": 28080
|
||||
},
|
||||
{
|
||||
"x_hash": "c12f401fafe1e333",
|
||||
"out_hash": "a574f69e6525862d",
|
||||
"t": [
|
||||
833.3333129882812
|
||||
],
|
||||
"start": 28080
|
||||
},
|
||||
{
|
||||
"x_hash": "f68f10bfa7c33eae",
|
||||
"out_hash": "1e2aa2baa76fdfbb",
|
||||
"t": [
|
||||
625.0
|
||||
],
|
||||
"start": 28080
|
||||
},
|
||||
{
|
||||
"x_hash": "503b0c92faec0367",
|
||||
"out_hash": "21ea40e0d7d642e3",
|
||||
"t": [
|
||||
0.0
|
||||
],
|
||||
"start": 28080
|
||||
}
|
||||
]
|
||||
Binary file not shown.
Binary file not shown.
@@ -1,23 +0,0 @@
|
||||
{
|
||||
"repo": "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
|
||||
"snapshot": "4b44356635ae5e927ca552a220f768022be76004",
|
||||
"fastvideo_commit": "c459a1897899ffcec3be7765534d81000b9bb9c1",
|
||||
"torch": "2.12.0+cu130",
|
||||
"python": "3.12.13",
|
||||
"gpu": "NVIDIA GB200",
|
||||
"attention_backend": "FLASH_ATTN",
|
||||
"seed": 1234,
|
||||
"e2e_prompt": "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest. The playful yet serene atmosphere is complemented by soft natural light filtering through the petals. Mid-shot, warm and cheerful tones.",
|
||||
"probe_prompt": "A cat and a dog baking a cake together in a kitchen.",
|
||||
"dmd_denoising_steps": [
|
||||
1000,
|
||||
750,
|
||||
500,
|
||||
250
|
||||
],
|
||||
"warp_denoising_step": true,
|
||||
"scheduler": "SelfForcingFlowMatchScheduler(shift=5, extra_one_step, sigma_min=0)",
|
||||
"num_frames_per_block": 3,
|
||||
"context_noise": 0,
|
||||
"notes": "causal chunk rollout via main's CausalDMDDenosingStage; per-forward hashes for all 35 forwards, full tensors for the first two chunks; no compile/fsdp; FLASH_ATTN"
|
||||
}
|
||||
Binary file not shown.
@@ -1,71 +0,0 @@
|
||||
{
|
||||
"fastvideo_commit": "c459a1897899ffcec3be7765534d81000b9bb9c1",
|
||||
"mode": "dmd2",
|
||||
"seed": 42,
|
||||
"gen_losses": [
|
||||
0.0,
|
||||
0.33539342880249023,
|
||||
0.0,
|
||||
0.31753748655319214,
|
||||
0.0
|
||||
],
|
||||
"fake_losses": [
|
||||
0.003410761244595051,
|
||||
0.00596819119527936,
|
||||
0.013151273131370544,
|
||||
0.022480690851807594,
|
||||
0.0037461533211171627
|
||||
],
|
||||
"torch": "2.12.0+cu130",
|
||||
"gpu": "NVIDIA GB200",
|
||||
"config": "dmd2 legacy: interval2 gw3.5 lr2e-6 shift8 steps[1000,757,522] simulate nlt4 1gpu",
|
||||
"self_noise_runs": {
|
||||
"gen": [
|
||||
[
|
||||
0.0,
|
||||
0.3357574939727783,
|
||||
0.0,
|
||||
0.3176053762435913,
|
||||
0.0
|
||||
],
|
||||
[
|
||||
0.0,
|
||||
0.33570683002471924,
|
||||
0.0,
|
||||
0.3168887794017792,
|
||||
0.0
|
||||
],
|
||||
[
|
||||
0.0,
|
||||
0.33539342880249023,
|
||||
0.0,
|
||||
0.31753748655319214,
|
||||
0.0
|
||||
]
|
||||
],
|
||||
"fake": [
|
||||
[
|
||||
0.003410761244595051,
|
||||
0.0059582265093922615,
|
||||
0.013146793469786644,
|
||||
0.022459683939814568,
|
||||
0.0037503200583159924
|
||||
],
|
||||
[
|
||||
0.003410761244595051,
|
||||
0.005963137373328209,
|
||||
0.013137918896973133,
|
||||
0.022351054474711418,
|
||||
0.0037392042577266693
|
||||
],
|
||||
[
|
||||
0.003410761244595051,
|
||||
0.00596819119527936,
|
||||
0.013151273131370544,
|
||||
0.022480690851807594,
|
||||
0.0037461533211171627
|
||||
]
|
||||
]
|
||||
},
|
||||
"self_noise_max": 0.0007165968418121338
|
||||
}
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -1,54 +0,0 @@
|
||||
[
|
||||
{
|
||||
"targets": [
|
||||
2
|
||||
],
|
||||
"dmd_t": null,
|
||||
"critic_t": 737.588623046875,
|
||||
"gen_loss": 0.0,
|
||||
"fake_loss": 0.003410761244595051,
|
||||
"x0_student_hash": "96f01be58584097b"
|
||||
},
|
||||
{
|
||||
"targets": [
|
||||
1,
|
||||
0
|
||||
],
|
||||
"dmd_t": 386.4990234375,
|
||||
"critic_t": 716.4179077148438,
|
||||
"gen_loss": 0.33539342880249023,
|
||||
"fake_loss": 0.00596819119527936,
|
||||
"x0_student_hash": "34f51407f1b630e5"
|
||||
},
|
||||
{
|
||||
"targets": [
|
||||
2
|
||||
],
|
||||
"dmd_t": null,
|
||||
"critic_t": 967.0870971679688,
|
||||
"gen_loss": 0.0,
|
||||
"fake_loss": 0.013151273131370544,
|
||||
"x0_student_hash": "2d65be80d333fd5d"
|
||||
},
|
||||
{
|
||||
"targets": [
|
||||
0,
|
||||
1
|
||||
],
|
||||
"dmd_t": 662.4631958007812,
|
||||
"critic_t": 980.0,
|
||||
"gen_loss": 0.31753748655319214,
|
||||
"fake_loss": 0.022480690851807594,
|
||||
"x0_student_hash": "822e22773ec52d1e"
|
||||
},
|
||||
{
|
||||
"targets": [
|
||||
2
|
||||
],
|
||||
"dmd_t": null,
|
||||
"critic_t": 456.45648193359375,
|
||||
"gen_loss": 0.0,
|
||||
"fake_loss": 0.0037461533211171627,
|
||||
"x0_student_hash": "472a8020d513b4ea"
|
||||
}
|
||||
]
|
||||
@@ -1,41 +0,0 @@
|
||||
{
|
||||
"fastvideo_commit": "c459a1897899ffcec3be7765534d81000b9bb9c1",
|
||||
"dataset": "wlsaidhi/crush-smol_processed_t2v",
|
||||
"seed": 42,
|
||||
"steps": 5,
|
||||
"num_latent_t": 8,
|
||||
"config": "1gpu bs1 accum1 lr5e-5 wd1e-4 betas(0.9,0.999) clip1.0 uniform-t cfg_rate0 dit_fp32 mixed_bf16 flow-match target=noise-latents flash_attn",
|
||||
"torch": "2.12.0+cu130",
|
||||
"gpu": "NVIDIA GB200",
|
||||
"losses": [
|
||||
0.1940414160490036,
|
||||
0.94484943151474,
|
||||
0.1001732274889946,
|
||||
0.9069435000419617,
|
||||
0.10525074601173401
|
||||
],
|
||||
"self_noise_runs": [
|
||||
[
|
||||
0.1940414160490036,
|
||||
0.9445295333862305,
|
||||
0.10048552602529526,
|
||||
0.9095064997673035,
|
||||
0.10758557915687561
|
||||
],
|
||||
[
|
||||
0.1940414160490036,
|
||||
0.9446913599967957,
|
||||
0.10070198774337769,
|
||||
0.9106038808822632,
|
||||
0.10818696022033691
|
||||
],
|
||||
[
|
||||
0.1940414160490036,
|
||||
0.94484943151474,
|
||||
0.1001732274889946,
|
||||
0.9069435000419617,
|
||||
0.10525074601173401
|
||||
]
|
||||
],
|
||||
"self_noise_max": 0.0036603808403015137
|
||||
}
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -1,82 +0,0 @@
|
||||
[
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"latents_hash": "93ae01bc47c0817b",
|
||||
"embeds_hash": "53f384dfed477e28",
|
||||
"noise_hash": "e1f135db4ff97ff8",
|
||||
"noisy_hash": "6a0cabf3c2e98794",
|
||||
"timesteps": [
|
||||
118.0
|
||||
],
|
||||
"sigmas": [
|
||||
0.1181640625
|
||||
],
|
||||
"pred_hash": "2903c217cc9279df",
|
||||
"loss": 0.1940414160490036,
|
||||
"grad_norm": 0.23514027893543243
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a colorful sponge being flattened as if it were under a hydraulic press, with the sponge being compressed and eventually flattened into a thin layer.",
|
||||
"latents_hash": "b9d5abc838e815da",
|
||||
"embeds_hash": "82c3f45fec89f678",
|
||||
"noise_hash": "ea2e0427b4584180",
|
||||
"noisy_hash": "5ff0bac18dd8e80e",
|
||||
"timesteps": [
|
||||
85.0
|
||||
],
|
||||
"sigmas": [
|
||||
0.0849609375
|
||||
],
|
||||
"pred_hash": "688def082dcbfa60",
|
||||
"loss": 0.94484943151474,
|
||||
"grad_norm": 10.1902437210083
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is composed of a large, cylindrical metal cylinder with yellow and black stripes, and a metal base. The objects being flattened are two cylindrical blocks of cotton candy, one pink and one blue. The press is positioned on a metal table, and the background features a green wall with a yellow and red sign.",
|
||||
"latents_hash": "840ddc5d78da4b7b",
|
||||
"embeds_hash": "ffce0c1bf664fb3d",
|
||||
"noise_hash": "2f512661cebc8a51",
|
||||
"noisy_hash": "05b1746f31f8d83d",
|
||||
"timesteps": [
|
||||
618.0
|
||||
],
|
||||
"sigmas": [
|
||||
0.6171875
|
||||
],
|
||||
"pred_hash": "963b9731aafbe2ea",
|
||||
"loss": 0.1001732274889946,
|
||||
"grad_norm": 0.523212730884552
|
||||
},
|
||||
{
|
||||
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
|
||||
"latents_hash": "c5d8005785df08ab",
|
||||
"embeds_hash": "88eea058d58c26f7",
|
||||
"noise_hash": "1022b2e68b615e5f",
|
||||
"noisy_hash": "ee1750fb9fdb7094",
|
||||
"timesteps": [
|
||||
41.0
|
||||
],
|
||||
"sigmas": [
|
||||
0.041015625
|
||||
],
|
||||
"pred_hash": "5fd1a77f3716b63b",
|
||||
"loss": 0.9069435000419617,
|
||||
"grad_norm": 3.6921098232269287
|
||||
},
|
||||
{
|
||||
"caption": "The video shows a large, industrial press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing a pile of pink objects into a pile of crumbs. The press is large and metallic, with a yellow and black striped pattern on its side. The background is a green wall with a yellow warning sign.",
|
||||
"latents_hash": "57045eb31ae1e81a",
|
||||
"embeds_hash": "71cee36d5ca3083d",
|
||||
"noise_hash": "cdc1cbc4a1fed9ac",
|
||||
"noisy_hash": "53ada1e537af47ac",
|
||||
"timesteps": [
|
||||
610.0
|
||||
],
|
||||
"sigmas": [
|
||||
0.609375
|
||||
],
|
||||
"pred_hash": "315edc5acc47767a",
|
||||
"loss": 0.10525074601173401,
|
||||
"grad_norm": 0.9913093447685242
|
||||
}
|
||||
]
|
||||
@@ -1,22 +0,0 @@
|
||||
{
|
||||
"fastvideo_commit": "c459a1897899ffcec3be7765534d81000b9bb9c1",
|
||||
"mode": "qad",
|
||||
"seed": 42,
|
||||
"gen_losses": [
|
||||
0.0,
|
||||
0.1784808486700058,
|
||||
0.0,
|
||||
0.17068849503993988,
|
||||
0.0
|
||||
],
|
||||
"fake_losses": [
|
||||
0.0031898675952106714,
|
||||
0.005877670831978321,
|
||||
0.015015869401395321,
|
||||
0.023572081699967384,
|
||||
0.0025676116347312927
|
||||
],
|
||||
"torch": "2.12.0+cu130",
|
||||
"gpu": "NVIDIA GB200",
|
||||
"config": "dmd2 legacy: interval2 gw3.5 lr2e-6 shift8 steps[1000,757,522] simulate nlt4 1gpu"
|
||||
}
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user