Compare commits

..
664 Commits
Author SHA1 Message Date
rlsu9 34a64597f4 update 2024-12-04 01:42:35 +04:00
rlsu9 4998121d6d update 2024-12-04 01:32:29 +04:00
rlsu9 1c25e3ed55 update 2024-12-04 01:24:30 +04:00
rlsu9 e98886f32f update 2024-12-03 10:03:26 +04:00
rlsu9 8619d2942d add l2 2024-12-03 09:53:44 +04:00
rlsu9 1876244531 update 2024-12-03 09:22:01 +04:00
rlsu9 34f709725e update 2024-12-03 01:44:07 +04:00
rlsu9 c5d6341f98 update 2024-12-02 22:45:29 +04:00
rlsu9 5325aa2cd2 xMerge branch 'main' of https://github.com/jzhang38/FastVideo-OSP into peiyuan 2024-12-02 22:31:48 +04:00
rlsu9 74d6fc92a9 update 2024-12-02 22:30:34 +04:00
haoanyscale daaefc5f63 fix finetune code bug 2024-12-02 11:45:04 +00:00
rlsu9 3eee42e917 update 2024-12-02 11:06:16 +04:00
rlsu9 28c0a5dce9 update 2024-12-02 11:03:44 +04:00
rlsu9 275dce4700 update 2024-12-02 10:53:14 +04:00
rlsu9 83fc3449e8 update 2024-12-02 10:44:00 +04:00
rlsu9 e96f359236 update 2024-12-02 05:12:47 +04:00
rlsu9 e89273a0da update 2024-12-02 04:18:38 +04:00
rlsu9 407e918676 runlong 2024-12-02 02:55:26 +04:00
rlsu9 94d86c168b wandb offline and dir 2024-12-02 01:31:24 +04:00
rlsu9 5a4c227572 update 2024-12-02 00:50:03 +04:00
rlsu9 554f934d37 update 2024-12-02 00:46:34 +04:00
rlsu9 51b72e7434 update 2024-12-01 23:05:29 +04:00
rlsu9 ebfbf6046a update 2024-12-01 21:56:36 +04:00
rlsu9 892324ae59 typo 2024-12-01 21:44:29 +04:00
rlsu9 b66e1a0427 new script 2024-12-01 21:40:25 +04:00
Hao Zhang 02ff44fcd1 update experiment 10 2024-12-01 05:43:24 +00:00
Hao Zhang eb551efe7f update scripts 2024-12-01 05:38:39 +00:00
rlsu9 4af3a99f9e typo 2024-12-01 09:36:42 +04:00
rlsu9 b3d16e0284 typo 2024-12-01 07:58:24 +04:00
rlsu9 c466fe9d73 Merge branch 'main' of https://github.com/jzhang38/FastVideo-OSP into peiyuan 2024-12-01 07:38:25 +04:00
rlsu9 494129344f debug gradient accumulation loss 2024-12-01 07:38:13 +04:00
Hao Zhang 8b9ecfdbcb update gitignore 2024-12-01 02:04:18 +00:00
rlsu9 58440f43de Merge branch 'main' of https://github.com/jzhang38/FastVideo-OSP into peiyuan 2024-12-01 05:26:04 +04:00
rlsu9 c8238e1f7b typo 2024-12-01 05:25:23 +04:00
Zhang Peiyuan 5ff08541f6 [Debug] Typo (#65) 2024-11-30 17:14:14 -08:00
rlsu9 f771899001 update 2024-12-01 05:13:51 +04:00
rlsu9 c609c28e63 update 2024-12-01 05:04:33 +04:00
rlsu9 c19f72f722 typo 2024-12-01 04:44:42 +04:00
rlsu9 43b4778474 Merge branch 'main' of https://github.com/jzhang38/FastVideo-OSP into peiyuan 2024-12-01 04:44:29 +04:00
Zhang Peiyuan 08f5489406 [Feat & Debug] fix uncond; multi guidance validaiton; multiphase schedule; linear range (#64) 2024-11-30 16:43:24 -08:00
rlsu9 7e6e52e02a update 2024-12-01 04:39:44 +04:00
rlsu9 e4114d8fd4 merge main 2024-12-01 04:35:01 +04:00
rlsu9 a19ed188c8 add script 2024-12-01 04:32:22 +04:00
rlsu9 6d3a71b681 linear range; fix uncond; multi guidance validation 2024-12-01 04:31:53 +04:00
Zhang Peiyuan 81103cd48d [Feat] EMA Distill; Distributed validation (#63) 2024-11-29 21:46:40 -08:00
rlsu9 a3d4a1cef4 add gupload 2024-11-30 09:41:37 +04:00
rlsu9 954dc8e35c distributed validation 2024-11-30 09:39:00 +04:00
rlsu9 278e8bb8f3 Merge branch 'main' into peiyuan 2024-11-30 05:28:37 +04:00
rlsu9 b378896b82 readme 2024-11-30 05:28:02 +04:00
rlsu9 1d2a7f233c distributed validation 2024-11-30 05:24:05 +04:00
rlsu9 b4aed1721f add aws efo env 2024-11-30 03:44:43 +04:00
rlsu9 3447ded21b revert to sp=4, sp bs=2, full shard 2024-11-30 03:43:39 +04:00
rlsu9 0c01223d3d update script; no sp 2024-11-30 02:51:22 +04:00
rlsu9 f8089afbf1 ema 2024-11-30 02:44:36 +04:00
rlsu9 2b0c1e66cb update env 2024-11-29 22:52:19 +04:00
rlsu9 79e7dac0ac ok 2024-11-29 21:50:49 +04:00
rlsu9 2b3f3cddd4 remove hardcode 2024-11-29 19:35:30 +04:00
rlsu9 064195363c remove harcode 2024-11-29 19:34:02 +04:00
rlsu9 ed1e17f37f add ema transformer 2024-11-29 11:42:17 +04:00
rlsu9 e32d137aec add ema_transform 2024-11-29 11:11:46 +04:00
rlsu9 62a5530c21 add 2024-11-29 11:08:32 +04:00
rlsu9 b3be22e1e2 add upload command 2024-11-29 09:44:13 +04:00
rlsu9 7156b8d7d5 typo 2024-11-29 09:12:22 +04:00
Zhang Peiyuan 6dd7980ec2 [Feat] Refactor GAN; State saving & Resume; Experiments script (#59) 2024-11-28 21:09:31 -08:00
Zhang Peiyuan b3c54b9e5b [Feat] Training precision (#57) 2024-11-27 16:25:07 -08:00
Zhang Peiyuan 333488ac09 [Feat][Debug] linear quadratic distill; HF precision bug (#56) 2024-11-27 15:13:20 -08:00
Zhang Peiyuan c7526707cc [Fix] Squeeze bug (#55) 2024-11-26 13:29:03 -08:00
Zhang Peiyuan a0ffbe2927 [Feat] PCM Distill; Refactor FM logit to be compatible with all SD3/Flux scheduler. (#54) 2024-11-26 12:26:24 -08:00
rlsu9andRunlong faef037f2e [feat]: Add batchy data preprocess (#53)
Co-authored-by:Runlong <rlsu9@ucsd.edu>
2024-11-25 21:52:36 -08:00
rlsu9andrunlong 9c45a15f53 [feat]: Add Image-Video Mixture training to main repo (#50)
Co-authored-by: runlong <r3su@ucsd.edu>
2024-11-24 14:36:45 -08:00
Zhang Peiyuan e6ff7ae2a3 Resolve config bug and seed (#51)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2024-11-15 20:29:48 -08:00
Zhang Peiyuan 897afb8b94 Add LADD (#45)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2024-11-15 19:28:41 -08:00
f8178a5a14 add lr scheduler; precision bug fix; add naive dataloader resume (#49)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
Co-authored-by: Yongqi Chen <yongqich@gl1712.arc-ts.umich.edu>
Co-authored-by: Yongqi Chen <yongqich@gl-login3.arc-ts.umich.edu>
Co-authored-by: Yongqi Chen <yongqich@gl-login1.arc-ts.umich.edu>
2024-11-14 19:57:38 -08:00
4c76b3cc7a [Feat] Lora resume (#48)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
Co-authored-by: Yongqi Chen <yongqich@gl1712.arc-ts.umich.edu>
Co-authored-by: Yongqi Chen <yongqich@gl-login3.arc-ts.umich.edu>
Co-authored-by: Yongqi Chen <yongqich@gl-login1.arc-ts.umich.edu>
2024-11-12 22:30:58 -08:00
3e0e8cf534 Add lora (#47)
Co-authored-by: Yongqi Chen <144848849+BrianChen1129@users.noreply.github.com>
Co-authored-by: Yongqi Chen <yongqich@gl1712.arc-ts.umich.edu>
Co-authored-by: Yongqi Chen <yongqich@gl-login3.arc-ts.umich.edu>
Co-authored-by: Yongqi Chen <yongqich@gl-login1.arc-ts.umich.edu>
2024-11-11 13:12:14 -08:00
Zhang Peiyuan 635fb2350d [Refactor] Switch to FSDP (#42) 2024-11-09 15:46:25 -08:00
Zhang Peiyuan a52f29374c No checkout mochi 2024-11-07 20:49:14 -08:00
rlsu9andhaoanyscale 31827c9e0a [feat]: Add adaptive fps dataloader and remove redundant code (#41)
Co-authored-by: haoanyscale <c-hao.zhang@anyscale.com>
2024-11-06 19:52:48 -08:00
rlsu9andhaoanyscale 74688e56b5 [feat]: Add vae encoder embedded generator to main (#30)
Co-authored-by: haoanyscale <c-hao.zhang@anyscale.com>
2024-11-06 11:58:15 -08:00
Zhang Peiyuan e56b2ae0c4 Add validation logging with SP (#36) 2024-11-05 14:46:20 -08:00
Zhang Peiyuan 6c6d4b34e7 [Feat] Sequence Parallel (#31) 2024-11-05 08:13:01 -08:00
Zhang Peiyuan 93d6ee49ba Merge pull request #25 from jzhang38/peiyuan
SP inference and Overfit training done.
2024-11-01 17:27:30 -07:00
rlsu9 e678eeb2dc Merge branch 'main' of https://github.com/jzhang38/FastVideo-OSP into peiyuan 2024-11-02 04:26:41 +04:00
rlsu9 432403b12d SP inference done! 2024-11-02 04:21:45 +04:00
rlsu9 8e899953ff Merge branch 'hl/diffusers' of https://github.com/jzhang38/FastVideo-OSP into peiyuan 2024-11-01 20:41:49 +04:00
rlsu9 fd80dc653d switch to deepspeed dummyoptim 2024-11-01 20:41:05 +04:00
foreverpiano 9fb510b618 fix some bugs; still output green 2024-11-01 12:30:38 +00:00
rlsu9 97f48ef433 Merge branch 'hl/diffusers' of https://github.com/jzhang38/FastVideo-OSP into peiyuan 2024-11-01 04:44:23 +04:00
rlsu9 5975deed62 OK 2024-11-01 04:26:21 +04:00
rlsu9 60c420731d Add zero3 2024-10-31 23:42:49 +04:00
rlsu9 a9a66a22a0 Debug successful! 2024-10-31 23:20:46 +04:00
Zhang Peiyuan fd415488df Merge pull request #21 from jzhang38/data_preprocess
[feat]: update data processor with multi GPU data parallel
2024-10-30 17:38:37 -07:00
haoanyscale a496c1c656 update data preprocess 2024-10-31 00:13:05 +00:00
rlsu9 2788f9351a Merge pull request #20 from jzhang38/main
Merge pull request #18 from jzhang38/peiyuan
2024-10-30 17:10:17 -07:00
rlsu9 5209b6ebdd debugging .. 2024-10-30 21:25:48 +04:00
foreverpiano db2c68e879 update inference sp code / can run / still has bug / don't output normal mp4 2024-10-30 16:36:50 +00:00
rlsu9 682774d730 Merge pull request #18 from jzhang38/peiyuan
[feat]: Merge training code into main branch
2024-10-30 09:29:50 -07:00
foreverpiano 2f4d6caa90 rename 2024-10-30 12:43:53 +00:00
haoanyscale dcd8e95c7b fix typos in readme and latent dataset debug file 2024-10-30 05:47:51 +00:00
a1286225768@gmail.com 9e2c9eaa01 overfitting .... 2024-10-30 04:14:04 +00:00
Peiyuan Zhang 567f9f66f8 debug 2024-10-29 21:10:50 +00:00
foreverpiano 48a82aba37 sp enable & still has bug 2024-10-29 15:28:27 +00:00
foreverpiano 4143cd1f95 random seed args 2024-10-29 15:27:52 +00:00
foreverpiano 5b3bffcca9 small bug 2024-10-29 14:17:54 +00:00
foreverpiano 37b091f2ef load 2024-10-29 12:47:48 +00:00
foreverpiano 70f3407032 update optimizer 2024-10-29 12:42:04 +00:00
Peiyuan Zhang d22d40c74a 14 2024-10-29 03:47:08 +00:00
Peiyuan Zhang 6fe5521667 Update generate_synthetic.sh and deepspeed_zero2_config.yaml 2024-10-29 01:00:31 +00:00
Peiyuan Zhang d1422b646d ok 2024-10-29 00:39:26 +00:00
Peiyuan Zhang d4e071be20 training 2024-10-29 00:38:40 +00:00
Peiyuan Zhang 7455bc9e7f commit first 2024-10-28 23:39:15 +00:00
Peiyuan Zhang 45c89c46c5 commit first 2024-10-28 21:46:01 +00:00
Peiyuan Zhang 7ec6f8f720 Deleted unnecessary files 2024-10-28 19:21:00 +00:00
Peiyuan Zhang cd6f8f14f1 remove files 2024-10-28 17:56:34 +00:00
Peiyuan Zhang 8824f74609 clean up name changing 2024-10-28 17:54:38 +00:00
Peiyuan Zhang 1ff9d45044 Remove OSP modeling 2024-10-28 17:53:13 +00:00
Peiyuan Zhang ffe11218dd Change name to fast video 2024-10-28 17:52:49 +00:00
Peiyuan Zhang 6c723e75f3 amend log validation 2024-10-28 17:50:34 +00:00
Peiyuan Zhang 4c3a83a0b3 Delete Open Sora Plan Modeling 2024-10-28 17:20:12 +00:00
Peiyuan Zhang d8489bf368 Generate synthetic dataset 2024-10-28 00:44:02 +00:00
Peiyuan Zhang f6a6366c58 Original pipeline 2024-10-27 23:30:04 +00:00
Peiyuan Zhang 5f8b952ede typo 2024-10-27 22:57:54 +00:00
Peiyuan Zhang a07aee80c3 Mochi Inferenfce & Diffusers 2024-10-27 22:57:14 +00:00
Peiyuan Zhang 235b6498b0 Refactor OpenSora sample_t2v.py and update download_hf 2024-10-26 21:59:22 +00:00
Peiyuan Zhang a89be390e7 remove vae loss 2024-10-26 19:20:04 +00:00
Peiyuan Zhang 3b8c76ea64 Remove files in causalvae 2024-10-26 19:06:46 +00:00
Peiyuan Zhang 63950da0d4 normalize 255; vae reconstruct 2024-10-26 18:57:07 +00:00
Peiyuan Zhang f245e4b4f8 Add mochi download & Change output dir 2024-10-26 18:08:50 +00:00
Peiyuan Zhang 1a8f9196ea Delete merge_data.txt 2024-10-26 18:00:16 +00:00
Peiyuan Zhang ea372da285 Remove unused adaptor files 2024-10-26 17:59:04 +00:00
Peiyuan Zhang fac510d043 Remove unused arguments in train_t2v_diffusers.py 2024-10-26 17:54:43 +00:00
Peiyuan Zhang b6b1a2aa92 Update T5Base 2024-10-26 17:53:51 +00:00
Peiyuan Zhang 464c8a421f Update PyTorch installation command 2024-10-26 17:36:23 +00:00
Peiyuan Zhang 602bf19d93 Update model path and cache directory 2024-10-25 10:21:31 +00:00
Peiyuan Zhang 84cf3b6795 Update t2v_debug_multi.sh with video_length_tolerance_range and dataloader_num_workers 2024-10-25 09:58:41 +00:00
Peiyuan Zhang 83a0280a14 update version 2024-10-25 09:22:16 +00:00
Peiyuan Zhang 182f4082f5 Update max height and width for video processing 2024-10-25 04:28:31 +00:00
Peiyuan Zhang 8d3f651217 Add pretrained model for OpenSoraT2V-ROPE-L 2024-10-25 04:26:57 +00:00
Peiyuan Zhang 330717768d include pretrained open-sora 2024-10-25 04:17:36 +00:00
Peiyuan Zhang ef653171dd Merge branch 'peiyuan' of https://github.com/jzhang38/FastVideo-OSP into peiyuan 2024-10-25 03:57:05 +00:00
Peiyuan Zhang 15fbad6dd5 Remove UDIT and inpaint 2024-10-25 03:55:15 +00:00
Peiyuan Zhang c00cc4b531 Update PyTorch index URLs and video length tolerance range 2024-10-25 03:18:16 +00:00
Peiyuan Zhang 1f423fca06 Fix warning with dataset handling and model loading 2024-10-25 02:31:07 +00:00
rlsu9 2cc0e17ade fix typo for dataset download 2024-10-25 01:40:26 +00:00
Peiyuan Zhang b1564eb141 Remove compress kv 2024-10-25 00:34:07 +00:00
Peiyuan Zhang 97f2b7fdaa Delete npu related stuff and remove inpaint module 2024-10-24 22:57:31 +00:00
Peiyuan Zhang 83825e2060 Update EMA model and t2v_debug.sh script 2024-10-24 22:19:48 +00:00
Peiyuan Zhang da9c1ab7c1 Update torchvision imports 2024-10-24 22:11:25 +00:00
Peiyuan Zhang 7c5b0a6ccd Remove all npu code 2024-10-24 21:35:27 +00:00
Peiyuan Zhang 5f21e3a979 Delete unused files and code 2024-10-24 21:09:24 +00:00
Peiyuan Zhang b47585c68d Update dependencies; Setup code for debugging 2024-10-24 21:01:32 +00:00
Peiyuan Zhang eb1669dab7 Add environment setup and training instructions to README.md 2024-10-24 17:47:34 +00:00
jzhang38 c1eb81f0f9 Remove NPU related code and update training process 2024-10-24 02:36:18 +00:00
jzhang38 8fb5c38381 Remove unused scripts and update TODO list 2024-10-24 02:27:03 +00:00
jzhang38 6e313fad97 Delete unnecessary files 2024-10-24 02:23:33 +00:00
Guangyi Liuandguangyi 294993ca78 [fix] fix the path typo for google/mt5-xxl in gradio_web_server.py (#405)
* [fix] fix the path typo for google/mt5-xxl in gradio_web_server.py

* [fix] fix the issue: the cache of pretrained mt5-xxl weights is inconsistent.

---------

Co-authored-by: guangyi <guangyi.liu@mbz-h100-029.core42.ai>
2024-08-23 11:14:30 +08:00
lb203 de7fc3150a Update train_inpaint.sh
do not set seed
2024-08-20 15:10:56 +08:00
lb203 eb1311b9a2 Update README.md 2024-08-20 11:23:51 +08:00
lb203 d3ea50fc05 Update Report-v1.2.0.md
add 29x480p link
2024-08-20 11:22:06 +08:00
yunyang GeandLinB203 842435c016 release Open-Sora Plan v1.2.0 i2v (#389)
* gitignore

* inpaint

* Create condition_image_path.txt

* Rename condition_image_path.txt to condition_images_path.txt

* Update sample_inpaint.py

* Update sample_inpaint.sh

* inpaint

* fix bug

* Update README.md

* inpainting

* release v1.2.0 i2v

* release v1.2.0 i2v

* Update Report-v1.2.0.md

* Update pipeline_inpaint_sp.py

* Update sample_inpaint_ddp.py

* Update sample_inpaint_sp.py

* Update sample_inpaint.py

* Grammar Error Correction

* Grammar Error Correction

---------

Co-authored-by: LinB203 <2267330597@qq.com>
2024-08-14 12:39:51 +08:00
lb203 394c6444aa Update train_t2v_diffusers.py 2024-08-05 15:23:05 +08:00
lb203 7ef00d7749 Update Report-v1.2.0.md 2024-08-01 13:06:37 +08:00
lb203 16e3afed62 Update README.md 2024-08-01 13:05:36 +08:00
lb203 3c08cb2fef Delete opensora/train/train_t2v_diffusers_lora.py
lora bug
2024-07-27 23:20:25 +08:00
LinB203 74744be314 fix sample 2024-07-27 07:59:13 +00:00
lb203 cc4ba38e1a Update README.md 2024-07-26 21:42:19 +08:00
lb203 39e00fd1c4 Update Report-v1.2.0.md 2024-07-26 21:40:26 +08:00
lb203 adb2a20a3d Merge pull request #352 from cxh0519/main
Update Report-v1.2.0.md
2024-07-25 14:12:33 +08:00
lb203 a6eb95f471 Merge branch 'main' into main 2024-07-25 14:11:05 +08:00
lb203 f0667d8db2 Update Report-v1.2.0.md 2024-07-25 14:10:17 +08:00
lb203 20b395624b Update README.md 2024-07-25 14:09:12 +08:00
Xinhua Cheng 40e0423a1e Update Report-v1.2.0.md 2024-07-25 14:05:45 +08:00
lb203 7e42eef228 Update Report-v1.2.0.md 2024-07-25 13:26:35 +08:00
lb203 08ae0a8379 Update t2v_datasets.py 2024-07-24 21:38:47 +08:00
lb203 b50b51ad92 Update Report-v1.2.0.md 2024-07-24 19:48:17 +08:00
lb203 a4300446d7 Update Report-v1.2.0.md 2024-07-24 19:46:43 +08:00
lb203 ebdfdc49e7 Update Report-v1.2.0.md 2024-07-24 19:46:14 +08:00
lb203 fbf13b7688 Update pyproject.toml 2024-07-24 19:45:28 +08:00
lb203 3a8c34efd0 Update README.md 2024-07-24 19:14:11 +08:00
lb203 ae5679185c Update Report-v1.2.0.md 2024-07-24 18:32:15 +08:00
lb203 067870f28f Update README.md 2024-07-24 18:32:02 +08:00
LinB203 535cb6330f release v1.2.0 2024-07-24 10:31:04 +00:00
lb203 b08681f697 Update Report-v1.1.0.md 2024-06-11 16:40:10 +08:00
lb203 d3ea240481 Update README.md 2024-06-07 18:35:08 +08:00
lb203 40bb7a55d7 Rename sample_video_513.sh to sample_video_221.sh 2024-06-01 21:25:35 +08:00
lb203 27c7046472 Update sample_video_513.sh 2024-06-01 21:25:27 +08:00
lb203 8d62c7273e Update gradio_utils.py 2024-06-01 18:28:18 +08:00
lb203 9dd3c8c494 Update README.md 2024-06-01 18:27:03 +08:00
lb203 46566a713d Update gradio_web_server.py 2024-06-01 18:24:50 +08:00
lb203 fb21d9938a Merge pull request #290 from Linzy19/0528lzy
fix the bug
2024-05-28 18:08:32 +08:00
lb203 be284a0274 Update README.md 2024-05-28 18:07:33 +08:00
ZongyingLin c58d27384a fix the bug 2024-05-28 15:16:49 +08:00
lb203 c421a6ed11 Merge pull request #286 from qqingzheng/causalvideovae-docs
Update CausalVideoVAE docs
2024-05-28 00:40:19 +08:00
New User 336403dd98 fix vis 2024-05-28 00:34:03 +08:00
New User 23aeb03f91 fix train bug 2024-05-27 23:50:36 +08:00
lb203 1a0872ac8e Update LICENSE 2024-05-27 23:40:13 +08:00
Zongjian 0360c1b6be [docs] update CausalVideoVAE docs 2024-05-27 23:36:18 +08:00
New User 559fbf3c36 fix wrong code 2024-05-27 23:33:49 +08:00
lb203 183b05173a Update t2v_datasets.py 2024-05-27 23:01:17 +08:00
lb203 ebefa85284 Update train_t2v.py 2024-05-27 23:00:56 +08:00
lb203 b2976643f8 Update train_t2v.py 2024-05-27 22:59:19 +08:00
lb203 b5b2c0968e Update modeling_latte.py 2024-05-27 22:57:34 +08:00
lb203 dc12340706 Update dataset_utils.py 2024-05-27 22:54:38 +08:00
lb203 b2a50079ac Update README.md 2024-05-27 21:01:02 +08:00
lb203 c9a7881dc5 Update pyproject.toml 2024-05-27 19:44:45 +08:00
lb203 71597994d1 Merge pull request #285 from cxh0519/patch-1
Update Report-v1.1.0.md
2024-05-27 19:44:26 +08:00
Xinhua Cheng 5b78e8e00f Update Report-v1.1.0.md
solve typos
2024-05-27 19:37:15 +08:00
YuanLi 858696d591 Update LICENSE 2024-05-27 19:20:26 +08:00
lb203 2a8b2328a5 Update README.md 2024-05-27 17:42:43 +08:00
lb203 f920d640d0 Update README.md 2024-05-27 17:12:39 +08:00
lb203 1e16f6be29 Update README.md 2024-05-27 16:39:16 +08:00
lb203 194ee0307a Update README.md 2024-05-27 16:33:23 +08:00
lb203 06eedf5ea3 Update README.md 2024-05-27 16:28:09 +08:00
lb203 a19488e33a Update README.md 2024-05-27 16:27:58 +08:00
lb203 60773f4860 Update README.md 2024-05-27 16:11:40 +08:00
lb203 da48eca111 Update README.md 2024-05-27 16:08:09 +08:00
lb203 e3e81fd4f1 Update README.md 2024-05-27 15:57:42 +08:00
lb203 91be347ad2 Update README.md 2024-05-27 15:51:22 +08:00
lb203 b3b0d60536 Update README.md 2024-05-27 15:47:46 +08:00
lb203 6e5515df28 Create Report-v1.1.0.md 2024-05-27 15:44:58 +08:00
lb203 732c672cb4 Update README.md 2024-05-27 15:44:27 +08:00
lb203 e2ae46b718 Update README.md 2024-05-27 15:44:00 +08:00
lb203 9cd5a906fd Merge pull request #284 from PKU-YuanGroup/dev
released v1.1.0
2024-05-27 15:43:39 +08:00
LinB203 27a98335b9 update prompt 2024-05-27 07:41:41 +00:00
LinB203 65ec02d036 5.27 2024-05-27 03:47:55 +00:00
LinB203 c21ea81533 fix demo 2024-05-26 07:39:11 +00:00
LinB203 8aebbc64c0 prepare scripts 2024-05-26 04:17:03 +00:00
root de3f7d8a9f prepare release 2024-05-25 02:54:40 +00:00
root 9b2951969f 5.15 2024-05-15 12:29:03 +00:00
root cfd61fdbb5 fix mask 2024-05-04 12:48:09 +00:00
root 78272b0941 update 2024-05-04 01:49:56 +00:00
root aa6d088e91 train vis 2024-04-30 15:23:43 +00:00
node106 594d98265b fix dataset 2024-04-29 13:15:41 +00:00
LinB203 8dc49b8a6b multi-data 2024-04-29 19:01:10 +08:00
LinB203 5e3a3c6f78 mask loss 2024-04-27 14:53:23 +08:00
LinB203 e767ef3a2e update dataset 2024-04-22 11:10:56 +08:00
LinB203 e92a28ba49 fix rope with compress 2024-04-21 23:30:46 +08:00
root 8358b3014c compress kv and rope pi 2024-04-21 11:44:26 +00:00
LinB203 3287e525c5 abs and rope 2024-04-17 15:45:54 +00:00
LinB203 c59023066e vae temporal tiling 2024-04-16 16:49:43 +00:00
LinB203 350138480f support newvae 2024-04-15 15:04:35 +00:00
LinB203 f1086e192c img training 2024-04-15 12:27:54 +00:00
LinB203 1323daa7b3 read image from folder 2024-04-14 12:03:04 +00:00
LinB203 8a6db0b399 refactor dynamic training 2024-04-14 11:37:16 +00:00
LinB203 8fa3e614a1 add 2drope and dynamic training 2024-04-14 02:04:14 +00:00
lb203 bec0e85238 Update pyproject.toml 2024-04-13 19:58:11 +08:00
lb203 098ecbbd5d Update Report-v1.0.0.md 2024-04-12 14:21:12 +08:00
YuanLi 8d3cd692e0 Update LICENSE 2024-04-12 09:49:31 +08:00
lb203 0183e89cb7 Merge pull request #218 from qqingzheng/vae_rec
Updated content related to VAE reconstruction
2024-04-11 19:06:38 +08:00
lb203 58b64d2662 Update README.md 2024-04-11 16:10:40 +08:00
qqingzheng f58e57d8f5 [refactor] fix hardcode 2024-04-11 03:56:55 +00:00
qqingzheng cdc626b14e [docs] fix typo 2024-04-11 02:36:10 +00:00
qqingzheng 8349763576 [feat] add time chunk inference 2024-04-11 02:31:20 +00:00
qqingzheng 02ad56275d [docs] update inference example. 2024-04-11 02:30:59 +00:00
lb203 461a4b2c97 Merge pull request #214 from JJJYmmm/fix_name_error_emavq
[fix]: Fix variable naming errors
2024-04-10 23:20:28 +08:00
lb203 65a3f0fe8a Merge pull request #207 from AlonzoLeeeooo/main
Fix typos
2024-04-10 23:17:53 +08:00
lb203 0cf62a3e0c Merge pull request #206 from digger-yu/patch1
fix typo
2024-04-10 23:17:15 +08:00
JJJYmmm e2492578c7 [fix]: Fix variable naming errors 2024-04-10 19:40:55 +08:00
USTC-liuchang d581ba621a Fix typos
line 24: `github -> GitHub`
line 54: `a -> an`
line 85: `re-organizes -> re-organize, modulizes -> modulize`
line 87: `opened -> open` (modified according to the changelogs of other dates)
line 94: `a -> an`
2024-04-10 11:51:18 +08:00
digger yu d65994469a fix typo 2024-04-10 08:59:00 +08:00
lb203 193b0398a8 Update README.md 2024-04-10 01:50:40 +08:00
lb203 e08fe61f2b Delete assets/we_want_you.jpg 2024-04-10 01:45:49 +08:00
lb203 c3cd4da606 Merge pull request #201 from qqingzheng/make_vae_better
[docs] update docs and train.sh

Former-commit-id: af001aec7d9bbf38d1e7e1f6a8ace5a77670b85e [formerly e6c3c58c77c2d8fdab1d5b464cd861ff78cc2e43]
Former-commit-id: 6ad032635deb1854b76d5a63a67f10442db0e655
2024-04-09 21:04:29 +08:00
qqingzheng a4fe6a634d [docs] update train.sh
Former-commit-id: eeb00ac046692bc97a064ec1b47c994436798ab1 [formerly 6d0fd85e1b13a7824b8181ba9ccebec2a3c76d40]
Former-commit-id: 1e8e669ad9967aeeeb7622bc04a0759f348cb26f
2024-04-09 13:03:34 +00:00
qqingzheng 933141c5ea [docs] add evaluation detail
Former-commit-id: c15a717cc61159e5f6adafb718237ab578ac81e7 [formerly 9e92772776fa4dd18772c7fec939a42ad6dfbc0a]
Former-commit-id: 98c2a38af5064f27888d1b102c9fded87d2420d5
2024-04-09 12:58:13 +00:00
qqingzheng 3f92ec2c53 [fix] fix bug in #185
Former-commit-id: d74efbd0195380ca383598c06d3d8185dd9cc9bc [formerly 0f2687266bc2271e11b2f06edfe514f466e57b0f]
Former-commit-id: cb6bb234befc49b3870c5690653469a0714890c5
2024-04-09 12:48:32 +00:00
qqingzheng 71541c766a [docs] update docs and train.sh
Former-commit-id: 1d9358e84d12863681b396c7621ae7244caa4dc2 [formerly b801491889b3abad1211225d7f07af3ae0389895]
Former-commit-id: 9e6a59ae2f246b50ace0a033cbd6548660241204
2024-04-09 12:43:19 +00:00
lb203 0623ee9a36 Update train_t2v_feature.py
Former-commit-id: 4a4f7ea2366f1dff5669d6d4834238ebd790e6e2 [formerly ca7d3158fb01f8e7228f833cea8e6741be127c2a]
Former-commit-id: 389f0e7dc66826026046b8e12d7a9f8692d32e42
2024-04-09 19:56:14 +08:00
lb203 4e54046c5f fix using pretrained bug
Former-commit-id: 2dfbb0683b2f3bd1a1d619e5e4fabb5bd4c92d79 [formerly 27f352e822f732cb351d64ac5726c5df3b62c1bb]
Former-commit-id: 722d60ae5ef28499c47292e0a4807a87b9658b01
2024-04-09 19:55:49 +08:00
lb203 326b6cbbe6 Update train_videoae_65x512x512.sh
Former-commit-id: 8e40db106efa2fd879cca42bd9ce2d9b79d777f8 [formerly 3e11781b045a838a5aa98ea75e274198f2b26a44]
Former-commit-id: f3b0804a7cf063e7ff04ee57dccee0320c1f32a5
2024-04-09 18:33:53 +08:00
lb203 71eb96c306 Update train_videoae_65x256x256.sh
Former-commit-id: b2f5cfe72452dfb369aba166e4a101dfcf5758c9 [formerly 557b78ab38a61d7520e91044137f818e50ec24cb]
Former-commit-id: 0e46c9f261e7f5fd4d08e3d7735c17c7975b0e72
2024-04-09 18:33:17 +08:00
lb203 a20766c262 Update README.md
Former-commit-id: be00a17729ac889493f63fd1d649054d7980ca42 [formerly 9b4d8dcfc8e6e2abe4ca727e01561d27d1ba3968]
Former-commit-id: f1b6f5013b99d0d288af0f05c46cb8300d79a8aa
2024-04-09 17:13:42 +08:00
YuanLi 72aef089ff Update README.md
Former-commit-id: 68eaf2e6b00266cf0858b33f9b4d09681efbd2bc [formerly d8dcad250ed8357964ab8ee570b9973b45e29fcb]
Former-commit-id: 75f51b3d84203be3672f6b07bb3c70ba0ae5556c
2024-04-09 16:56:19 +08:00
YuanLi 9b011137b3 Update README.md
Former-commit-id: c48eba0bb15d594007809d4cc605dc46ba029cec [formerly 7deb8a0decf7a9757af4475e3029d415d78d185f]
Former-commit-id: 7f43bb2f1d9308bd08154171d78e9b36e00a837f
2024-04-09 16:50:28 +08:00
YuanLi 14dacb1ca1 Update README.md
Former-commit-id: 54bad83bf054a2ff574199931121c686b4e28008 [formerly 06bb190a39a9b98c8de5e385fc304a86cc8c9fb2]
Former-commit-id: 23e8a235d6828d9ab7c02ef59d40fe13a46fa4b1
2024-04-09 16:47:05 +08:00
lb203 fd3198d7db Update README.md
Former-commit-id: 6fd60c73b6ae69d6ad9dc14c7371bd9bff086b44 [formerly 2ed5d5b7bdcd88044da26ca94237c50af7ab47dd]
Former-commit-id: 443c333353963bc1ec64a4946afbe5e17f9b5a87
2024-04-09 15:26:19 +08:00
lb203 1556d341ed Add files via upload
Former-commit-id: 943dd9a063baaa05d7f0d0bd52cbde35cf2b3d98 [formerly 5542359738243a8e948a8f08ed4e8dedf6b4b818]
Former-commit-id: fbee07f6fef2ac86bc15ff66f2e2a1a68b64ed01
2024-04-09 15:25:46 +08:00
lb203 3794a468fa Update Report-v1.0.0-cn.md
Former-commit-id: 83124efeaa5a8cd528a15760bdf88acc403142bd
2024-04-09 13:36:35 +08:00
lb203 b21d386205 Update Report-v1.0.0-cn.md
Former-commit-id: 59edce70edbaed957ba40f1d95b4a5050ece12bf
2024-04-09 11:34:43 +08:00
lb203 9b72cb3e6f Update Report-v1.0.0-cn.md
Former-commit-id: 0edac24fcaa447a534755bf3aac1f2027a491c06
2024-04-09 11:34:17 +08:00
lb203 5b5159e8fc Update Report-v1.0.0.md
Former-commit-id: 3c8eb75dc2045154e8654f860b1fabd7e597094b
2024-04-09 11:33:34 +08:00
lb203 113bb8c5df Update Data.md
Former-commit-id: 15defc1a3846335fd68e18fa9b8ffbe38e3e7100
2024-04-09 11:32:24 +08:00
lb203 4ad069ffdc Update Report-v1.0.0-cn.md
Former-commit-id: d5cf6aed7c1673562a6d2022c3fd7f7b16006cfd
2024-04-09 11:09:02 +08:00
lb203 421aa9b66e Update README.md
Former-commit-id: 0cc27a5b065b61f37d44dc2c90ba2782075eefb2
2024-04-09 11:06:18 +08:00
lb203 8c40057697 Update README.md
Former-commit-id: c437a18dbec7ac6e02287565a84d75f96664ee99
2024-04-09 11:06:05 +08:00
lb203 791a48a088 Update README.md
Former-commit-id: e386ec112337271bb9f55ea6ec872eaba94aff50
2024-04-09 11:02:06 +08:00
lb203 b85795c263 Create Report-v1.0.0-cn.md
Former-commit-id: d484754e080bd65f6b35619b840040581fa6b69c
2024-04-09 11:00:43 +08:00
lb203 bd895c4d90 Update README.md
Former-commit-id: 88e7f2062c494c199caff468f5f14f3165f763de
2024-04-09 09:46:28 +08:00
lb203 4731a994b0 Merge pull request #187 from qqingzheng/add_causalvae_docs
add causalvae doc

Former-commit-id: d4166b1cf6356e4415bbdba315c34e23ffdd0231
2024-04-09 09:39:48 +08:00
Chestnut bf4364b50a 更新 README.md
Former-commit-id: 42607786f3f3d0eba9336f774b0a2ae076cce793
2024-04-09 00:25:46 +08:00
Chestnut 1817a15c17 Update README.md
Former-commit-id: 567ade18047a9cc81648997c5d9287fa8900614d
2024-04-08 23:43:00 +08:00
Chestnut 8bc9abb45b Update Train_And_Eval_CausalVideoVAE.md
Former-commit-id: 156739c35b5dda594138b2e3b1fd6ec99556af23
2024-04-08 23:41:49 +08:00
Chestnut 2dc20a2cd1 Update Train_And_Eval_CausalVideoVAE.md
Former-commit-id: 589657b089930379bb75f3d49063ebb72671efc6
2024-04-08 23:40:07 +08:00
qqingzheng 7d9f03db56 add causalvae doc
Former-commit-id: 4e6d9e038c8edb81f8402834bd01eeeb823924f9
2024-04-08 12:25:43 +00:00
lb203 5e5c6a64dc Update Data.md
Former-commit-id: d9a5c8fbd36a34b7baf720bc27d69604ae5c9d57
2024-04-08 19:28:52 +08:00
lb203 9906697152 Update pyproject.toml
Former-commit-id: bb037420b75303b58486ba79da9180b2bf63cf88
2024-04-08 19:00:26 +08:00
LinB203 1f7edca9fd clean
Former-commit-id: a7f2c1c8b2243587cedd6e152e64298842287095
2024-04-08 18:57:06 +08:00
lb203 a7c3aebf27 Update README.md
Former-commit-id: 828be4267f0795969293e4a321e4e0bfa8e1f1e6
2024-04-08 16:08:12 +08:00
lb203 029d7eba30 Update README.md
Former-commit-id: dc923fe2413b3be31d8d88e7dd83a931fcbf9458
2024-04-08 14:44:58 +08:00
lb203 22b86b0e32 Update Report-v1.0.0.md
Former-commit-id: 37f181f34f64ca1141a1345b7c1f2af283be1cad
2024-04-08 14:27:11 +08:00
lb203 0e8162191f Update Report-v1.0.0.md
Former-commit-id: 400f708a080f0799e02697bb8e66d8430ffb1f56
2024-04-08 10:54:11 +08:00
lb203 8139823ff5 Update README.md
Former-commit-id: 2873d0eaedfad392c627a714448118949d7b7e13
2024-04-08 10:40:23 +08:00
lb203 b37a6ff3a7 Update Report-v1.0.0.md
Former-commit-id: 6ebe2c62c0904142026d3796027e95ce43afe16a
2024-04-08 10:37:04 +08:00
lb203 172e1bcc58 Update Report-v1.0.0.md
Former-commit-id: 38b779022418cc41cdf9fc508e461f2f6299fd60
2024-04-07 23:48:08 +08:00
lb203 22f61c0272 Update README.md
Former-commit-id: 9c0f82ae9953b56f2540e9c7ada021e4e1bc6165
2024-04-07 23:47:36 +08:00
lb203 cd99a6759a Merge pull request #183 from chaojie/patch-2
Update README.md fix colab link

Former-commit-id: 257c9275a0bdd8ae931592b7e5cf302f0d97d900
2024-04-07 23:07:21 +08:00
YuanLi 8cd4d9671e Update README.md
Former-commit-id: 5731e810e2efb7ca7eda1aa8954eb7fd260f4250
2024-04-07 22:55:49 +08:00
chaojie 459369d3b6 Update README.md
Former-commit-id: 1b54444249fa91b29334b8fc364fb1978c25fb2d
2024-04-07 22:52:01 +08:00
YuanLi 15113e2e7a Update README.md
Former-commit-id: bbf3cef43a129a25a80ce134a8895796e24353c6
2024-04-07 22:50:49 +08:00
YuanLi 9016110a90 Update README.md
Former-commit-id: c4a976e9b93536299b2e350454c780fec54fbe95
2024-04-07 22:45:48 +08:00
lb203 154070c2d0 Update README.md
Former-commit-id: c1492211ced3cdbec93ef28678a1467245e122e9
2024-04-07 22:26:58 +08:00
lb203 ed570f7b18 Merge pull request #178 from qqingzheng/causalvae_release_pr
Fix some compatibility issues

Former-commit-id: 11af500907302d699b761799408462a0ebfb9b31
2024-04-07 19:45:14 +08:00
qqingzheng 492dbab0d6 fix bugs in inference
Former-commit-id: e17c7a7b7704920cb86894419bf573d0fb0c06fa
2024-04-07 11:36:57 +00:00
qqingzheng 0044ec6623 Fix bugs in inference
Former-commit-id: c251a86ac7c91d6f2bb12ae645628e9cf4b83d3d
2024-04-07 11:11:16 +00:00
lb203 e1fc7291d5 Update README.md
Former-commit-id: ce2bbf82045041495ed448b46d5f4bdba84aa740
2024-04-07 18:46:43 +08:00
lb203 d3bbaa7a55 Update Report-v1.0.0.md
Former-commit-id: 7cec40420601f7da57449586508e50119bd02a61
2024-04-07 18:11:34 +08:00
lb203 d58cc4504a Update README.md
Former-commit-id: 9f926abb08f156807091dd044c133788649a6bac
2024-04-07 16:57:38 +08:00
lb203 d30e3653f2 Update README.md
Former-commit-id: 9bae31c73ea902581e1ee94cfd9012efdf78d349
2024-04-07 16:51:14 +08:00
lb203 b08bf8d907 Add files via upload
Former-commit-id: 9a207435d9efac0b19771df5d2288160af1bdd6e
2024-04-07 16:49:45 +08:00
lb203 d86ff4f933 Update README.md
Former-commit-id: f709d75e7a7d8a08ec60132aa906f6b880af8e58
2024-04-07 16:38:47 +08:00
lb203 79bf57e77c Update Report-v1.0.0.md
Former-commit-id: aea1f4a82a525d3e6e1d1a1bb7631326b294cdcc
2024-04-07 16:33:28 +08:00
lb203 3afaf0c5fe Update Report-v1.0.0.md
Former-commit-id: bd9aa7991212d98a6224b855a3bfb2c12feb5950
2024-04-07 16:31:48 +08:00
lb203 fcd94894da Update README.md
Former-commit-id: f8cb7f38751d7a6ef586c293fb769f5d60ad5c59
2024-04-07 16:30:47 +08:00
lb203 a3cf68c883 Update README.md
Former-commit-id: 30f61958e221c3d41c4aa5917f4511e8c0482cf9
2024-04-07 16:28:34 +08:00
lb203 5df5af0629 Update README.md
Former-commit-id: d5ab411b2e2fd5d6abeda500895e3339167572ab
2024-04-07 16:27:37 +08:00
lb203 deaea23969 Add files via upload
Former-commit-id: a232ec2d98b245a9b502a2c49fc4f468a86fbd64
2024-04-07 16:22:43 +08:00
lb203 3e6abd244d Update README.md
Former-commit-id: 3210474a46845ae3a6c0bde7fb6c7f9845866628
2024-04-07 16:13:48 +08:00
lb203 760029dbfe Update README.md
Former-commit-id: 7d4118566b77e09de4a3615612a5b368c76c293e
2024-04-07 16:12:52 +08:00
lb203 f08fa0b9d5 Update README.md
Former-commit-id: 037d7d3e084d7ff426347db49c4b210a2ee3927d
2024-04-07 15:58:09 +08:00
YuanLi c335b0b70c Update Report-v1.0.0.md
Former-commit-id: 92d906b340f5970a77377586d7295f46b5de99ec
2024-04-07 15:57:37 +08:00
lb203 cdd41c2aa3 Add files via upload
Former-commit-id: 292cbcfa4e4f384fe6958badc27d3067173119f9
2024-04-07 15:57:35 +08:00
lb203 d02d16e6e2 Add files via upload
Former-commit-id: 1ae8512b7846e9faafdb38ca532898d7f1bdcce2
2024-04-07 15:50:18 +08:00
lb203 7434308e92 Update README.md
Former-commit-id: a2ec90a19eacb099cb9f36c1b5dd6e24be307fae
2024-04-07 15:08:09 +08:00
lb203 511ea03617 Update README.md
Former-commit-id: a61200bea2eea4fd22bd8095f192fef4c22c49b4
2024-04-07 15:04:51 +08:00
lb203 8d5f0281ae Update README.md
Former-commit-id: e6fb5d864ac8af8214cdc89c71b42f296dda14e7
2024-04-07 14:59:24 +08:00
lb203 68ec69bc6c Update and rename CausalVideoVAE.md to Report-v1.0.0.md
Former-commit-id: 3b848f8de33cf4f0ed802428e9ad31d2fa26de63
2024-04-07 14:57:39 +08:00
LinB203 ef0157a168 update gradio demo
Former-commit-id: 80c2cc0bcc2a7532bc2440df16c1f811b0b30516
2024-04-07 09:44:23 +08:00
LinB203 20c3149b62 clean
Former-commit-id: 56ed87088e13cdd93929e2ceedbfae693e2f013a
2024-04-06 23:41:01 +08:00
qqingzheng 35f5e4cff8 Fix some compatibility issues
Former-commit-id: 942d8598b28516369780c3519607929869b427b7
2024-04-06 14:08:02 +00:00
LinB203 282f973d9b v1.0.0
Former-commit-id: 93f44f3b425c090d9bc92687a9c50a26b4f0d7c8
2024-04-06 22:04:52 +08:00
lb203 4b8b80835b Merge pull request #177 from qqingzheng/causalvae_release_pr
CausalVideoVAE release version

Former-commit-id: 50dd03b07b4a57e708ca01f703d806e9fc568d2a
2024-04-06 19:48:26 +08:00
qqingzheng a554dcc696 release
Former-commit-id: 52fe2ac6610661d6a13a08390c84a8237f37986a
2024-04-06 11:40:05 +00:00
LinB203 a4a1b93bbc refactor vae to hf
Former-commit-id: a938d3f7f962709f7ecff5e9d49c9bf1ce40d7fa
2024-04-06 19:26:08 +08:00
LinB203 cea0c041df update train script
Former-commit-id: ac16f77d19488cd46130587369cd7a5ea378052b
2024-04-06 13:28:29 +08:00
LinB203 8137d40a92 add tile training
Former-commit-id: 216079dc925beefe1ec722a4e87e46a4e1ca56be
2024-04-04 22:00:38 +08:00
LinB203 e202525517 sample pipeline
Former-commit-id: 8964b16c0dd65445a20050c9a0e82d909b55eaff
2024-04-04 10:50:29 +08:00
LinB203 8084949b14 released v1.0.0
Former-commit-id: 8df5897a4ffc341ced45faf617687fb8ebf1286c
2024-04-04 10:46:29 +08:00
lb203 b3f6d06941 tile only2d
Former-commit-id: 7f31ad03fdbaac97aea51a98bd29dc7327b700c9
2024-04-01 12:27:40 +08:00
lb203 b7f0d770b8 Merge pull request #172 from SamitHuang/fix_attn3d
[Bug fix] Fix reshape bugs in AttnBlock3D in CausalVideoVAE

Former-commit-id: 8a6b0947c126c87d13809adbd0adfdf65716c27f
2024-03-31 15:07:13 +08:00
Samit ee0b749421 fix reshape bugs in AttnBlock3D
Former-commit-id: 299622e168a6cbee14ef2056c47cae722c20e005
2024-03-31 14:52:28 +08:00
LinB203 0c86a4ff11 tile conv
Former-commit-id: 5133615bcba221c4b86ae4e8b09b82ee5d34c66d
2024-03-30 20:42:20 +08:00
LinB203 2aaf448a24 update sample
Former-commit-id: 116726b9d565c689496123d4b6e62b98400536a6
2024-03-30 15:30:29 +08:00
lb203 a50abb7b49 Update __init__.py
Former-commit-id: 2bdc29d86ebff3d6f29c38ded57c7b0dd1e6bbdb
2024-03-28 22:04:27 +08:00
lb203 b88dfd6d1e Update README.md
Former-commit-id: 2d9ae56309717a91e0690bd5c354fe3cbe5d5ed6
2024-03-28 21:16:34 +08:00
lb203 5dc588f47c Update CausalVideoVAE.md
Former-commit-id: fcde4d0ff783448c2006d6f4751a0e4f9cc23d0f
2024-03-28 20:16:49 +08:00
LinB203 f48b3da80a fix videovae trainer
Former-commit-id: 79e1feb412b6eeead3306f7f0d55a01f9b2a543e
2024-03-28 20:04:15 +08:00
LinB203 9462239286 update vae
Former-commit-id: 24298df5d374d5aaa5696f6200d4e23ae887715f
2024-03-28 19:57:15 +08:00
LinB203 c70648b740 update model
Former-commit-id: e4e24650376d9dce974b290752efb416d5061983
2024-03-28 19:10:24 +08:00
LinB203 45187a0d22 fix training bug
Former-commit-id: 09a6283ddd9ab11479f30324de788bd7415702ae
2024-03-28 19:07:46 +08:00
lb203 b88ae0d76d Merge pull request #165 from qqingzheng/add_eval
[feat] add eval code

Former-commit-id: 2b555c515ed71b8f699627423b492a5c43218b84
2024-03-28 19:02:32 +08:00
qqingzheng c564dc569d [feat] eval
Former-commit-id: d2c3cf5ef8d18feb7b0544d68d9796656e90ba8d
2024-03-28 11:00:20 +00:00
lb203 55a595ece8 Update README.md
Former-commit-id: 02b0422fd745b4b7ec091cc6ee81d1d237245fe6
2024-03-27 23:13:29 +08:00
lb203 c7a9f094c8 Update README.md
Former-commit-id: f294a0a28f288ae592e7de004320d3eb5e48e945
2024-03-27 23:13:01 +08:00
lb203 2dbd3bd700 Update README.md
Former-commit-id: 90d6d298eaee6255098b342b31444e7a347fc50e
2024-03-27 23:10:14 +08:00
lb203 c19b099618 Update CausalVideoVAE.md
Former-commit-id: 1cd665f699632e9dc9357379225c085d9c1278d5
2024-03-27 22:59:36 +08:00
lb203 7db5b81522 Update README.md
Former-commit-id: a41674f5ef71ced59b374226dddfa09c15a99f25
2024-03-27 22:57:59 +08:00
lb203 be1eb20b95 Update README.md
Former-commit-id: 8d70a890fc4191bfc98c3aca6dd3c8a70534fb3d
2024-03-27 22:57:40 +08:00
lb203 ce836de93d Rename causalvideovae.md to CausalVideoVAE.md
Former-commit-id: 0a78e4193fcaec4662b7c6465a9d80643b354cd6
2024-03-27 22:55:50 +08:00
lb203 44eb12b157 Create causalvideovae.md
Former-commit-id: 98414c2ad7dd39f740440b2d9059140f93492e83
2024-03-27 22:55:32 +08:00
lb203 5f9ff24a3d Merge pull request #162 from qqingzheng/add_causal_vae
[feat] add causalvae ✨

Former-commit-id: 4c1b27c277f6dee020d7aa9d6f719415f0eed898
2024-03-27 21:43:31 +08:00
qqingzheng ec81e2427c add causalvae
Former-commit-id: ac936276fd3213ee1815dfe096cd64cecb601428
2024-03-27 12:49:30 +00:00
LinB203 f78a599fa3 scripts
Former-commit-id: 047046f684c2831e68e0cf256cb46fc3790cda33
2024-03-25 16:33:14 +08:00
LinB203 6c12429fd2 train t2v feature
Former-commit-id: 7f16b162cae0e3533c8ba9aa03186a846d9e02bf
2024-03-25 16:25:01 +08:00
lb203 c6d7ed5c3f Update Data.md
Former-commit-id: d1c72e3209769d0204efb8e7d705bdf4f22b82ea
2024-03-22 10:30:04 +08:00
lb203 b3fb8cf3c5 Update README.md
Former-commit-id: 859c0f60dac1562b8a6bbb72b793a4709dbc3597
2024-03-21 16:02:23 +08:00
lb203 ebaa135355 Merge pull request #152 from Ytimed2020/main
Add CLIP support and example

Former-commit-id: 3c3f80caa5d24fac1f6bf5c01b4cb8df86f2edfd
2024-03-21 15:59:24 +08:00
lb203 211c7409bc Merge branch 'main' into main
Former-commit-id: e4c99f4b5d084171e3097bc8939d4806f32485a8
2024-03-21 15:59:03 +08:00
lb203 dfef484db5 t2v attention_mode
Former-commit-id: 4fd8bd87e0d0a98d5d0a3201d5db9a5ba055b894
2024-03-20 23:04:02 +08:00
lb203 4d476bc759 Update README.md
Former-commit-id: 95b28706d257d7769049abe30a7121b60e39a53f
2024-03-20 22:35:11 +08:00
lb203 a4b273f255 Update README.md
Former-commit-id: 9065d12ee867c76452ce370fb71cc83bdc01470b
2024-03-20 22:33:39 +08:00
LinB203 6b51d0ccfd support attention_mode
Former-commit-id: c72f354fbfd948e8cd4b0cc63c157bc4688c7665
2024-03-20 22:23:24 +08:00
LinB203 ff36d291d8 train with image
Former-commit-id: 209b9a8ad2f19a50e46f9b7ec68dd781898e491a
2024-03-19 23:21:57 +08:00
Ytimed2020 73ad660277 Update clip.py
Former-commit-id: 1ee53d152630b6acc69369be33287abfb143224d
2024-03-18 23:03:33 +08:00
Ytimed2020 0817d81290 Create clip.py
Former-commit-id: e3397c567efcfeef6f58dfef2eb4251bedbbf33a
2024-03-18 23:01:34 +08:00
Ytimed2020 f2ac960431 Update __init__.py
Former-commit-id: 7c8af2909c4549e96bc0893d09fcc53d71fb6cc7
2024-03-18 23:01:04 +08:00
lb203 81cc4190fd Merge pull request #145 from qqingzheng/add_casual_vqvae
[feat] add casual vqvae ✨

Former-commit-id: f034a4cfe8bc84fee5a50a1da833a39c9499f213
2024-03-18 13:10:12 +08:00
lb203 1f34030723 Update pyproject.toml
Former-commit-id: 5637656a79faba1b397b56918c42022cc209e731
2024-03-17 20:22:03 +08:00
lb203 757268eccd Merge pull request #146 from glgh/doc-fix
[docs]: fix VQVAE training script path

Former-commit-id: 6ea832b958e6c8e5494a7b199874ffce061325aa
2024-03-17 09:41:01 +08:00
gl 83cf1dbd8e [docs]: fix VQVAE script path
Former-commit-id: d5f5c681b8c50622a84256881d193d5803ff3c28
2024-03-16 12:20:59 -07:00
qqingzheng b3c17829e9 [feat] add casual vqvae
Former-commit-id: 89087970b267908c8667984aba42b1ade2505467
2024-03-17 02:46:33 +08:00
lb203 862f627607 Update README.md
Former-commit-id: 6729039f059a23bf71faa944e2498657ef2c7442
2024-03-17 00:00:45 +08:00
lb203 73161c3248 Update train_t2v.py
Former-commit-id: d25b39ffac4474abd40b551fe0e51c421cb7bf5e
2024-03-16 22:23:36 +08:00
lb203 1d1a6a5b63 Update train.py
Former-commit-id: 783faf5cd86c148a713ea986ceed8527b108f2b8
2024-03-16 22:23:22 +08:00
LinB203 085062a53f refactor and fix resume bug
Former-commit-id: 1b1136671b9a7e0b15b70168e2c44dbff04a22e2
2024-03-16 22:21:15 +08:00
LinB203 516082e06c resume and compress kv
Former-commit-id: 0af7d58f88c1b65b4089da6f3e5c70456614da7c
2024-03-15 22:48:09 +08:00
lb203 054bf4bfb0 Update README.md
Former-commit-id: d97edd0e714a237a54f191dcebb6fd39f73185d6
2024-03-15 22:22:42 +08:00
lb203 510434d9a1 Update README.md
Former-commit-id: 34e529fd02712d4466ab2bc1a63f9bd69976b6c3
2024-03-15 20:14:47 +08:00
lb203 befdbd5943 Update pyproject.toml
Former-commit-id: 7fb921440b608e47781d2f541b2a8d893a48193e
2024-03-15 20:13:18 +08:00
LinB203 daee16c393 update scripts
Former-commit-id: 8b2a794095f2deb9b1e8090486cb8242bf81031d
2024-03-15 19:09:18 +08:00
lb203 04034c2e50 Merge pull request #143 from anapple-hub/my-script-branch
[fix]: fix sample.sh

Former-commit-id: 7ceacc5f9544d0e27d8cea4981ed12ee784efd6f
2024-03-15 19:00:35 +08:00
LinB203 d2a8e64a83 del t2v
Former-commit-id: 5daa2952997a29e74e2b5d235001d6d9dabf9ff6
2024-03-15 18:43:44 +08:00
anapple-hub 99bbe2ecbe [fix]: fix sample.sh
Former-commit-id: 09916dfcfb0b80a4217016919e0394a2cb665cfd
2024-03-15 18:30:49 +08:00
lb203 210690e431 Update README.md
Former-commit-id: 7b96fd89eaf3743a4a9fbc995e6a2a791c734f2e
2024-03-15 13:47:29 +08:00
lb203 17c430042e Update sample.sh
Former-commit-id: b149a300f1aff53d9ab6e66a91cc2bcd37900347
2024-03-14 13:21:23 +08:00
lb203 4b8c154784 Update sample.py
Former-commit-id: eb906691abec030d7aa2078b63fc9c7d678662ed
2024-03-14 13:21:04 +08:00
LinB203 99a789bd5d del t2v
Former-commit-id: 910398ca297be8a188d59c59c121857666774af1
2024-03-14 08:51:10 +08:00
lb203 d82c3b9e60 Merge pull request #139 from qqingzheng/fix_vqvae
[bug] add quantization loss

Former-commit-id: 58bb160e834dcd74650741e2003720d6ae5cad5f
2024-03-13 18:47:15 +08:00
qqingzheng 08d8ef6e72 [bug] fix quantization loss
Former-commit-id: 111fcf55f969ffe4c9be284a317ae5712570a527
2024-03-13 18:16:28 +08:00
lb203 51581271b4 Update README.md
Former-commit-id: 2d5b6815f95e94a98847b8a0376473f01bb15d14
2024-03-13 14:02:33 +08:00
lb203 014e5b4933 Merge pull request #130 from sennnnn/dit_deepspeed
⭐ [Feature] Support deepspeed training for DiT

Former-commit-id: 9b25f038f32a03dd56552a132118757d0835c428
2024-03-13 13:40:03 +08:00
lb203 18a652afff Update README.md
Former-commit-id: cf1c56c0058288427c1dfd651c8c0e9a981c5653
2024-03-13 13:33:48 +08:00
lb203 8aa4e60cf4 Update README.md
Former-commit-id: 7e74356a3279b02d062bab9c9f361f61eb756179
2024-03-13 13:22:16 +08:00
sennnnn 9750ead38b Fix typos.
Former-commit-id: 6add5ee7ef39a973e91b7a7c1b3a4b0980130b88
2024-03-12 21:23:44 +08:00
sennnnn 66f085252e merge main branch.
Former-commit-id: 5f78da5eae2b18305bb4a0a6dc7ab771e9b96f60
2024-03-12 21:11:56 +08:00
lb203 f8b0701a4b Merge pull request #123 from yunyangge/main
[docs]: frame interpolation update

Former-commit-id: ced46982cd41903a212a841500618a614e7c9b8a
2024-03-12 14:44:21 +08:00
yunyang Ge 1c72bea528 Merge branch 'PKU-YuanGroup:main' into main
Former-commit-id: 314be102744ec83be3bc37750ea4d6674e383de4
2024-03-12 14:21:45 +08:00
lb203 c8811aa2d5 Update train_256.sh
Former-commit-id: a2dfc9a8c70845446e731c88acd222e9b2db9122
2024-03-12 14:21:14 +08:00
yunyang Ge 07d818282d Update readme.md
Former-commit-id: e287f94838007f2eda434a46bcd250099b627097
2024-03-12 14:19:52 +08:00
yunyang Ge e4ea15e66a Update and rename Frame Interpolation.md to readme.md
Former-commit-id: 87bdf1ebecb002c086497b3c47acb3f3cb2aaf34
2024-03-12 14:19:16 +08:00
yunyang Ge 59e2050bc3 Add files via upload
Former-commit-id: b707b1968fa59d85891040604e7d54bef61b4549
2024-03-12 14:18:00 +08:00
lb203 a0f4a001a3 Update README.md
Former-commit-id: 8ec7524f4f171e650d7a25cc88ed29f7b5017454
2024-03-12 14:16:18 +08:00
yunyang Ge 78b9cb8775 Update interpolation.py
Former-commit-id: e10e3416433e5f396aa13e7d8beaf451ff47a383
2024-03-12 14:13:54 +08:00
lb203 ff617764b7 Merge pull request #104 from sennnnn/videogpt_deepspeed
⭐ [Feature] Support deepspeed for videogpt training.

Former-commit-id: dc639030e3544ca54e9c579624375bdea486a0ef
2024-03-12 14:12:51 +08:00
lb203 f8ca07a1d8 Update sample.py
Former-commit-id: 6e09fc266a626f465109abc199d9fb452657929e
2024-03-12 14:10:10 +08:00
lb203 61306c7510 Update README.md
Former-commit-id: 0d7dd6233811ec3c8741affb436d8004d8c676b8
2024-03-12 14:06:10 +08:00
lb203 429391840f Merge pull request #121 from sysuyy/add_multi_node_script
[feat]: use accelerate on multi-node

Former-commit-id: 67b01dc1d934d3d86d7ecca00c5dd1a11fd4d22d
2024-03-12 14:05:08 +08:00
lb203 aa3ff1e9c3 Update README.md
Former-commit-id: c5eea8d3492872255a6a34d422230509d1c3531c
2024-03-12 13:35:03 +08:00
lb203 e0dac4fb40 Create train_256.sh
Former-commit-id: bf62151617feb49663f8dfe0c78aa01f41e62092
2024-03-12 12:55:41 +08:00
lb203 486218b252 Update README.md
Former-commit-id: 5717aab628fceba311e284d1bdddeacc4c43c7c5
2024-03-12 11:52:28 +08:00
lb203 21383acca7 Update README.md
Former-commit-id: dd00e04220ff5232c5c4503a095e409471476a51
2024-03-12 11:49:20 +08:00
lb203 e4949048e0 Update README.md
Former-commit-id: 266f9f97ea8f327fdefb2ca959c23d08b7084137
2024-03-12 11:47:12 +08:00
lb203 c39a13d6b0 Merge pull request #120 from touale/main
[fix]: correct attention_mode to attention-mode

Former-commit-id: 6deb9393f771edbf8db8a95e7b9f2107b2682c76
2024-03-12 11:00:50 +08:00
sennnnn ddf7e82b61 Update training arguments of videogpt training.
Former-commit-id: c9660af891c0dae1f51a1075ef5b2df012f65e06
2024-03-12 10:30:41 +08:00
sennnnn 064a503c31 use imageio for write video which has better compatibility.
Former-commit-id: 005119dc9f78ea2a611d22a53772fd59d97d7139
2024-03-12 10:30:01 +08:00
sennnnn 91409c89cb Can't fix zero loss bug.
Former-commit-id: 1f32f9d889743b7961b34b823b276bfd03974191
2024-03-12 10:17:24 +08:00
sysuyy 330e46e9e8 [feat]: use accelerate on multi-node
Former-commit-id: 0b12fc31f1d2a63f9da38821174ab63d9b3a2d72
2024-03-11 14:45:23 +00:00
touale 76e172e08c [fix]: correct attention_mode to attention-mode
Former-commit-id: 9957e2eebb2166b02600c6052fb54908f5d12bad
2024-03-11 22:14:34 +08:00
sennnnn 896d5e3657 Fix deepspeed zero loss bug.
Former-commit-id: e24af415bc5d6fd38d72a24d6312f4e65d3627dd
2024-03-11 21:37:23 +08:00
sennnnn a360da1ec7 update videogpt.
Former-commit-id: 2e211563c684a8b6ef62dfc0c4f1897224dfb58d
2024-03-11 20:35:01 +08:00
sennnnn 46126baadf use fp16 temporarily.
Former-commit-id: 93381968e10aae7bbec23ec1ef40e7cb311fddf0
2024-03-11 17:03:23 +08:00
lb203 2eba66af88 Update README.md
Former-commit-id: c2c64665d4f281e264d87c20cc1cfd38e7e5f168
2024-03-11 16:58:01 +08:00
sennnnn a5d891a8bc Add deepspeed script for videogpt training.
Former-commit-id: 3bda2a05cb789fe785b2355a613385b3eb7abbbc
2024-03-11 16:29:11 +08:00
lb203 d1c7015eaa Update README.md
Former-commit-id: abb8ed57a29b1f66adec66d854429f36f9199540
2024-03-11 16:10:31 +08:00
sennnnn 56ed84acbd Add hidden_size argument for videogpt config.
Former-commit-id: 1be9f22767fe9c80ad90cb17184751b97aff9012
2024-03-11 15:46:54 +08:00
sennnnn 9fbd96d66c Merge branch 'main' into videogpt_deepspeed
Former-commit-id: c815b5eae7d60f89df0366f762a9a4afc8c7655e
2024-03-11 15:45:15 +08:00
LinB203 0e2a86b2ac clean dit
Former-commit-id: 117ec4196a6f2cd83da223b993b14e28f7a693f8
2024-03-11 15:44:53 +08:00
sennnnn 59030ac60c Merge branch 'resolution_typo' into videogpt_deepspeed
Former-commit-id: acaac6af0e35fcc611a9ecd15a39d6bb41aed4aa
2024-03-11 14:36:07 +08:00
sennnnn 2b17f4fac2 Merge main branch.
Former-commit-id: 800ca74fad0d9021f1e96ddf32cd93cdb37cbbe0
2024-03-11 14:35:17 +08:00
lb203 523670a478 Merge pull request #117 from sennnnn/resolution_typo
Fix resolution typos

Former-commit-id: b9acfa97f5514513f6c983f3e815300b435d3076
2024-03-11 14:32:59 +08:00
sennnnn 1365640065 Fix resolution typos.
Former-commit-id: 3ad55d28eaf13f9c8fd23912d2691965737527e9
2024-03-11 14:26:05 +08:00
lb203 59e2407366 Merge pull request #97 from HowardLi1984/refiner
[feat]: Caption Refiner

Former-commit-id: 6af3b3689679a99acb741218416dc3a9cde29f45
2024-03-11 14:05:35 +08:00
sennnnn 1321e1f977 refactor dit for supporting deepspeed.
Former-commit-id: e1909745cb05ea2f28f3a391c194f7656593a236
2024-03-11 11:51:46 +08:00
lb203 14e217a698 Update README.md
Former-commit-id: 1f94f2721f35670aeae8dbfb7c0a5abd06cd2aee
2024-03-10 21:37:01 +08:00
lb203 d4aa360c50 Update README.md
Former-commit-id: 2808edf11005e04aa730971a786d1a4e3bca88b5
2024-03-10 21:16:15 +08:00
lb203 93b624e62d Update README.md
Former-commit-id: 64d147863368aeef479ad25f2e5463fdcb76de46
2024-03-10 20:27:46 +08:00
lb203 f423fc7538 Update README.md
Former-commit-id: 6d41cfe906f97dbed1cea8c9ed06a1f6dc84308e
2024-03-10 20:25:48 +08:00
LinB203 a7165b506e clean script
Former-commit-id: 02c61f63a8539a31516b92f5bb7c2d2e428ca7d7
2024-03-10 18:34:37 +08:00
lb203 740363460f Update README.md
Former-commit-id: 7a62dec060c91b02ca556451266e6f3ef28e03e7
2024-03-10 18:27:50 +08:00
lb203 aeb3081341 Update pyproject.toml
Former-commit-id: 5c379b6e4515ad75b89cdebbc5cd490c9a821d78
2024-03-10 18:21:50 +08:00
lb203 fb62b8ff7e Update README.md
Former-commit-id: 059d9e863958cfd679dbb8d365de6cd933e284fc
2024-03-10 18:20:28 +08:00
lb203 74494966bd Update README.md
Former-commit-id: 87636c1dc2275c4da77c6368a0f4425bea15dc12
2024-03-10 18:15:21 +08:00
LinB203 b57709108a train 1080p video
Former-commit-id: 74625142ef892f9e7fa459b2f98f2f7ae5b36d6a
2024-03-10 18:08:23 +08:00
lb203 1f4d4cbdfd Merge pull request #112 from sennnnn/diffusion_deepspeed_fix
🐛 [BUG] Add base config for Latte

Former-commit-id: 825aa78ded6a5436d53d488b2c2250ac41b2bff9
2024-03-10 15:29:50 +08:00
sennnnn 995454b518 Add base config.
Former-commit-id: d37c551a0421d6b73dee26060b517c46d8217e4f
2024-03-10 15:24:40 +08:00
lb203 3d3ea0c818 Update README.md
Former-commit-id: 582070d4f97c4a4398dcb9a6699f4ab6e980c812
2024-03-10 15:17:18 +08:00
LinB203 6dae0b1c0c deepspeed
Former-commit-id: 6179a4b0e7d06caed4201a27679abc64615febe6
2024-03-10 15:08:12 +08:00
lb203 557b062028 Merge pull request #110 from sennnnn/diffusion_deepspeed
⭐[Feature] Support deepspeed for latte training.

Former-commit-id: 59e731cb98915a2635547cad904ced7d9fcf73f4
2024-03-10 14:48:03 +08:00
sennnnn 28b7fc4bae Fix typos.
Former-commit-id: c10330d373466f3c1ab455bfb391569b9175bc3a
2024-03-10 13:17:44 +08:00
sennnnn 61f9c5bced Add some comments.
Former-commit-id: d5b5bd7743edde6a63dc74457bded0f8c56be1c2
2024-03-10 12:23:42 +08:00
sennnnn f3274235d8 Support deepspeed zero2 and zero2_offload for latte training.
Former-commit-id: c87e95dd357e6fa6d6fb2dcfee781bc49c039e13
2024-03-10 11:27:33 +08:00
sennnnn 81c455e878 Refactor latte.
Former-commit-id: 7408fd0ee16247afc3d4ba43b8c8e7dfb1091a97
2024-03-10 11:25:56 +08:00
sennnnn 443bab1364 Refactor latte.
Former-commit-id: 304534269f105578097bc16ab497636ad59fd9b2
2024-03-10 11:25:50 +08:00
sennnnn bfe4f15694 Fix mixed_precision bugs of latte.
Former-commit-id: d70a551972e04d95444d46cdbe2bfb16ad099fcb
2024-03-10 11:25:07 +08:00
sennnnn ee3810a585 Decouple mixed_precision and gradient_accumulation to command args.
Former-commit-id: 903d473d1432eb492a87efe184d03439f643b92d
2024-03-10 11:24:39 +08:00
lb203 d6d54da131 typo
Former-commit-id: 9e87366e906ed19726c8905afb27b5df4d4e1512
2024-03-10 11:15:14 +08:00
LinB203 f034826c6d fixed attn_mask in flash-attn
Former-commit-id: 5e2504f429a342b75bc7a9004cee7b3f75582ee3
2024-03-10 11:12:12 +08:00
lb203 74ee558f49 Merge pull request #107 from jpthu17/fix_2d_RoPE_init_bug
[fix] 2d RoPE init

Former-commit-id: e873ffa66f01268cce600f8da9ad297dc3e8aaaf
2024-03-10 10:21:43 +08:00
jpthu17 75938cfaeb [fix] 2d RoPE init
Former-commit-id: ff5931bae6efa94675b9ef6bc25f1b76217d128e
2024-03-09 22:04:15 +08:00
LinB203 9a0d0bc843 safe dit
Former-commit-id: f407cab7dc0e53257972d430e873f4fd77bb039d
2024-03-09 21:53:43 +08:00
LinB203 5ebb382d85 xformers and flashattn
Former-commit-id: 002cad95253f7da36344237ce8c8410ffb517711
2024-03-09 21:52:03 +08:00
lb203 c08f985c8c Update README.md
Former-commit-id: 6630c22baa2340a8b55badd61472daab7177bc63
2024-03-09 21:19:47 +08:00
lb203 c891f85022 Merge pull request #106 from jpthu17/add_2d_RoPE
[feat] Add 2D RoPE

Former-commit-id: 85c46fa0e728c9cfc5af4cf6523a7e6a64f96a93
2024-03-09 21:18:17 +08:00
lb203 ad54bd8f94 Merge pull request #105 from sennnnn/einops_bug
Fix einops bug: module 'keras.backend' has no attribute 'is_tensor'

Former-commit-id: c4f3a3a1172b68a4a34ddfac8ea69be34d11038b
2024-03-09 21:13:29 +08:00
jpthu17 aef95b043c [feat] Add 2D RoPE
Former-commit-id: ccae80d339e34fb23706827565b5095ae2c2f320
2024-03-09 19:53:48 +08:00
sennnnn 5834bc0a8b Merge branch 'main' into videogpt_deepspeed
Former-commit-id: 04dbc4c5f04d1e82815e02c6215d8ff6e098049f
2024-03-09 19:24:12 +08:00
sennnnn 29e167fecf Fix einops bug: module 'keras.backend' has no attribute 'is_tensor'
Former-commit-id: 849f8ed815295bac78c3c10d319e6731aa26ca02
2024-03-09 18:18:26 +08:00
sennnnn efad4348d5 Fix config typos.
Former-commit-id: 455973b57bda19379d41e59aa418421fd269281e
2024-03-09 17:39:03 +08:00
sennnnn 3b4e881b11 Support deepspeed for videogpt training.
Former-commit-id: 1f3dfe038845d91b81b0ecc9b6f3cabb06eed5d5
2024-03-09 17:33:30 +08:00
lb203 e415e25897 Update README.md
Former-commit-id: 513883df1ad853c9a5e7457c364227261df17a83
2024-03-09 17:13:18 +08:00
lb203 e2b563ab2d Merge pull request #91 from RuslanPeresy/typecheck
[refactor] add static type checking

Former-commit-id: 969dfa6dedaeb64cf776d977b466b14568cc4b1c
2024-03-09 16:54:19 +08:00
LinB203 6df13b07bb flash-attn train and sample
Former-commit-id: e336565bd5ff2aaf4d0b4ad5f42732e89cd216b3
2024-03-09 16:30:09 +08:00
lb203 39e398e815 Merge pull request #103 from rain305f/new_eval
[docs]: update the eval code

Former-commit-id: 5ee073e7493a608c27c257b2a341585192c5b68f
2024-03-09 15:47:54 +08:00
rain305f e780368bb6 [docs]:update the eval code
Former-commit-id: a8c9e6cb9e39a9e929dab34077280a8d80d85573
2024-03-09 06:22:04 +00:00
rain305f 8a987de2a8 [docs]:update the eval code
Former-commit-id: 937f2fc87d34ab9d4724cb95a6457b54a6bf4ebc
2024-03-09 05:21:43 +00:00
lb203 a60217616e Update README.md
Former-commit-id: 38e0e290e7a5909706c2cce2e4dd047647134fb8
2024-03-09 11:49:17 +08:00
lb203 fa401d1d3f Update README.md
Former-commit-id: 88e9f21cfc8ec775dbc9a65f740043e91e26bd91
2024-03-09 11:44:12 +08:00
lb203 81f5e207d0 Update README.md
Former-commit-id: 826cc5e9162888bed82413088c6be662c278176b
2024-03-09 11:41:37 +08:00
lb203 7774b0d1eb Update pyproject.toml
Former-commit-id: 36da5d0e950625e629aebe954f89f92827397d0d
2024-03-09 11:37:47 +08:00
lb203 82cc72a739 Merge pull request #96 from Jason-fan20/my-docs-branch
Update README.md

Former-commit-id: 426cbdfa78282e098739dc1740d802d8e9c7d6de
2024-03-09 10:58:50 +08:00
lb203 39848c3dff Merge pull request #102 from sennnnn/vqvae_data
🐛 [BUG] Fix the memory error bug when loading metadata pickle file of vqvae dataset

Former-commit-id: 3b458d91fdc646e63f004031ebab9984a34934d2
2024-03-09 10:58:18 +08:00
Li Hao c8ba21979e [BUG]: Shorten the refined caption
Shorten the refined caption using GPT-3.5 summary


Former-commit-id: bb159f0d5fa928258475b2058983397ceb91af9d
2024-03-09 10:22:09 +08:00
lb203 38c2c2ccaf Merge pull request #99 from jialin-zhao/jialin-zhao-xformers-inputs-revised
Update latte.py about the xformers input dims

Former-commit-id: 421b49c351124170ab476e17cba6e45bd7cf706e
2024-03-09 09:58:51 +08:00
lb203 c42a86310d Merge pull request #94 from rain305f/eval_code
[feat]:update the eval_code for calculating the FVD, clip_score etc. metric

Former-commit-id: 363b237f71387e51164d02ed622c5b997fd5dde3
2024-03-09 09:43:25 +08:00
LinB203 b0a9e72008 extract script
Former-commit-id: 4139f2f68ae09695302ce49afdb7318df83cabfd
2024-03-09 01:03:05 +08:00
lb203 7573495d70 Update README.md
Former-commit-id: 99b0066895af072a8887737f2b603ec7fc893ccd
2024-03-09 00:33:09 +08:00
rain305f 130fcd397a [docs]: update the evaluation
Former-commit-id: 02e3789b74daa036fbbff6d13dfc7e98372260ed
2024-03-08 15:38:54 +00:00
sennnnn 41b14c6a13 Fix the memory error bug when loading metadata pickle file of vqvae dataset.
Former-commit-id: 97ab243bd12915cea7aa2deecfac4b99417f78a3
2024-03-08 22:54:17 +08:00
lb203 f751508193 Merge pull request #55 from SimonLeeGit/main
[refactor]: add docker development support

Former-commit-id: 159477551119f68a07adeab349d37083b390e533
2024-03-08 20:51:55 +08:00
lb203 8eef9df1b4 Merge pull request #61 from Mon-ius/main
[feat] - add CI for docker autobuild

Former-commit-id: f9b004a99e67a32be46c1bb3267ad20378b1e89b
2024-03-08 20:50:48 +08:00
lb203 61643ba9e8 Merge pull request #100 from Linzy19/230308
[Updata] updata the video super resolution

Former-commit-id: ebfe0c526f0ee977a91ef0214658783972afc0ac
2024-03-08 20:49:49 +08:00
ZongyingLin 80d857ac86 [Updata] updata the video super resolution
updata the run.py and add README.md


Former-commit-id: 1e6d6ace7a0963149e228813801210b022aceacf
2024-03-08 20:45:47 +08:00
Li Hao 99fe427ac8 [feat]: Caption Refiner
Caption Refiner for Video Caption


Former-commit-id: 97c16bccb9979acf19fb91ae8d7b84ac4a3b4d97
2024-03-08 20:04:41 +08:00
chialin cd2ffd7229 Update latte.py
revise the xformers inputs

Former-commit-id: 60db66c912192c3fbafe55ebe16157514b884e5b
2024-03-08 20:02:31 +08:00
Jasonfan 825fe9eada Update README.md
Former-commit-id: 8f7b41118a3f491d1d5b5789c53790c46b023191
2024-03-08 19:59:51 +08:00
lb203 e3f002e280 Update README.md
Former-commit-id: 70f6ebf421f0830a51dcc2838c2c3783cfb67376
2024-03-08 19:24:44 +08:00
LinB203 a23d24fd11 fix extract feature
Former-commit-id: d4dce5ec1aaa058e452a41f15cd35fe1713f0b78
2024-03-08 19:22:36 +08:00
LinB203 2fd41ff003 update sample.py
Former-commit-id: 10d8621051d1510270c6a55abe686eac86297170
2024-03-08 17:28:16 +08:00
LinB203 bc4e4f6fa6 del use_fp16
Former-commit-id: 5424c0112e04b56d1df30ac76cbc104e848406a3
2024-03-08 17:27:22 +08:00
LinB203 eaee962fa9 update train.py
Former-commit-id: 844173711342f2dcc021294e87971ca39854bfe4
2024-03-08 17:25:48 +08:00
LinB203 b07538bf8a feature dataset
Former-commit-id: 9683caae8648735552270b871de22fc540ba2b55
2024-03-08 17:24:23 +08:00
LinB203 21fad0d34a update script
Former-commit-id: 43f539b68af150c45040d79918a84de1dce47c33
2024-03-08 17:21:06 +08:00
rain305f 77a79e67b6 [docs]:update the eval_code for calculate the FVD, clip_score, ssim, lpips, psnr
Former-commit-id: b504623d7ec2ac0b508530d21ffc88978f8a4f71
2024-03-08 09:19:42 +00:00
rain305f cf2e4f48d0 [docs]:update the eval_code for calculate the FVD, clip_score, ssim, lpips, psnr
Former-commit-id: 1ee226c338126fbdc5178020b12271b980a7ce0c
2024-03-08 09:11:05 +00:00
SimonLee 9c406ca0d7 Merge branch 'main' of https://github.com/PKU-YuanGroup/Open-Sora-Plan
Former-commit-id: aa4a88029885ff0b6e4e96174c07927116150434
2024-03-08 15:44:23 +08:00
RuslanPeresy 7e7dcbe406 [refactor] add static type checking
Former-commit-id: 12c4240dc28888c61f1f79243627e7cdb445c840
2024-03-08 15:37:47 +08:00
SimonLee 9711db6896 update docker scripts and configs
Former-commit-id: 4567e6368f1bf413d8f6319a97a852b8289ca770
2024-03-08 15:32:17 +08:00
lb203 e37dae81ac Delete scripts/train_vqvae.sh
Former-commit-id: 952a4ce3a3b0da159d599ff3a0a052265f48c2dd
2024-03-08 11:21:00 +08:00
lb203 20ca4c089c Update README.md
Former-commit-id: c6b1ea8f1ef01096aa8c3e5272bc9fb255b5532b
2024-03-08 11:18:43 +08:00
SimonLee 518ce56f09 [refactor]: add docker development support
Former-commit-id: fded3baff7915c0cfc9cec3765c271259359831c
2024-03-08 10:58:28 +08:00
lb203 997738438f Update README.md
Former-commit-id: 97fe914aaa2b9bd7c2ff5d4239381631c47e8356
2024-03-08 10:01:16 +08:00
lb203 b76c71e634 Update __init__.py
Former-commit-id: 03a45aaf60d294b7a68d8fc44132d3d671e4472b
2024-03-08 09:46:17 +08:00
Liu Hanchen 203105e504 Update README.md
Former-commit-id: b37da2202e730a4930a194a80a703b591a535adf
2024-03-08 01:27:35 +08:00
Liu Hanchen 76ce11567a Update README.md
Former-commit-id: bca7070fac79d5ae5c37830cd55ca1e81a289197
2024-03-08 01:24:29 +08:00
Liu Hanchen d6980acdd1 Update README.md
Former-commit-id: c4d7aafa3ceb6e4d48bf85b402ae2b2ce9394e53
2024-03-08 01:16:28 +08:00
Liu Hanchen 2448aa7c8c commit
Former-commit-id: 81359cf67248b6c594ff79f121a8a38558c1cc70
2024-03-08 01:04:16 +08:00
lb203 bb167c3a70 Update Data.md
Former-commit-id: 86be6b2bf5a9be1030bb3af661a2ec85e710a5c3
2024-03-07 23:00:54 +08:00
lb203 85232cc966 Update Data.md
Former-commit-id: 9531e33666f842125b331cb266eb26a3a03a4009
2024-03-07 22:35:54 +08:00
lb203 ad5632cc50 Delete opensora/models/super_resolution/placeholder
Former-commit-id: 8dd5a6a3c85dc28989ff95e6ad41be7fdb0f0811
2024-03-07 22:19:00 +08:00
LinB203 9c8b50a29a train script
Former-commit-id: 3fd1ccf099c2deed20af50bf9556e2fc3b7a4382
2024-03-07 22:17:39 +08:00
lb203 68cf1dd673 Merge pull request #88 from Linzy19/2307
[updata] updata video super resolution

Former-commit-id: 7f512b4f186a795358b7e2c922068e2f881aabae
2024-03-07 21:53:22 +08:00
ZongyingLin ae766cb6dd [updata] updata video super resolution
updata video super resolution


Former-commit-id: 82397545859791d112d6f0569e33fe23437563df
2024-03-07 21:27:28 +08:00
lb203 2e84d17e50 Update README.md
Former-commit-id: 7cd398721e556db6199f1d06c7a0d1de2a9720dd
2024-03-07 21:22:24 +08:00
lb203 9f4a42baf9 Update README.md
Former-commit-id: 076e6e70e717888f24265d624c6a7a8ffa2f5d1b
2024-03-07 20:48:32 +08:00
lb203 7a03fe4ccc Update README.md
Former-commit-id: deddba1f03bf697fb1252a0ce4fdb5df5bcae8de
2024-03-07 20:44:55 +08:00
lb203 af8ade6655 Update README.md
Former-commit-id: a1d1f07c83eb0c03d93b023b5c9162a6a3c26cd9
2024-03-07 20:44:12 +08:00
LinB203 af49b2a9e9 init pos embed
Former-commit-id: 28f8076136ec214c66039ead7547c9b25cfe7bfc
2024-03-07 20:01:10 +08:00
lb203 0ca3104739 Update README.md
Former-commit-id: e458366e47d4b13737a71bbf866e20ec10961ade
2024-03-07 19:52:52 +08:00
lb203 781c8ca81d Update README.md
Former-commit-id: 3b4eb14dcd8862b8251b72f30cba523a5dc1def8
2024-03-07 19:52:06 +08:00
lb203 2c7fa77589 Merge pull request #84 from Nyx-177/main
[docs] Modify README.md to fix spelling and grammar issues

Former-commit-id: 096a602e57b5dc3a99d98fc55d608f7569b17af6
2024-03-07 19:42:57 +08:00
Nyx 4cd6d293ca [docs] fix merge conflict
Former-commit-id: d8b1afe957765f3844bd9c1d88dd925efa57330e
2024-03-07 21:30:20 +10:00
Nyx 5149e05147 [docs] Modify README.md to fix other grammatical errors
Former-commit-id: f0de27fcc38bb4dab0eeafac225a2924080a3218
2024-03-07 21:15:56 +10:00
lb203 fbd9b89c01 Update README.md
Former-commit-id: 8a191abc1903ad5998fe83149c9c9ef5c8226b9e
2024-03-07 19:14:49 +08:00
lb203 a0a8e4fb76 Update README.md
Former-commit-id: a881ff5198e1125032418f4f02fd95d2179ffb3a
2024-03-07 19:14:12 +08:00
Nyx d57553d520 [docs] Modify README.md to fix translation error in name "CloseGPT"
Former-commit-id: dd36043df4b0408e592a82f9046f6d654a5375c3
2024-03-07 20:55:02 +10:00
lb203 8e0e8ba91e Update README.md
Former-commit-id: 2a68a1aa399a53f58246df4f9f64dfdcdd130bcf
2024-03-07 18:46:14 +08:00
lb203 28ea2cc1b2 Update README.md
Former-commit-id: 7bc544af55c0a987d0df426848cb06cb330fd52b
2024-03-07 18:16:23 +08:00
lb203 4a07be0b0e Merge pull request #82 from sennnnn/star_his
Rapidly increasing stars are the best honor of a great open-source project.

Former-commit-id: 6aeda408cc64765f00092e61b6ad9a3f0b48fe4d
2024-03-07 17:51:10 +08:00
sennnnn ff69d06080 Fix conflict.
Former-commit-id: 6d4a3950cb135718b3f122657733cb5f890db971
2024-03-07 17:46:31 +08:00
sennnnn 390d7f1896 complement emoji.
Former-commit-id: 50679cc482d8ecf4e3083c0f5a791f195a58a855
2024-03-07 17:39:51 +08:00
lb203 ca9bbce122 Update README.md
Former-commit-id: 7af6b82b29c265e53e9c48477083b4cff0e0a3d3
2024-03-07 17:37:24 +08:00
sennnnn 0fde7b4edd refine readme.
Former-commit-id: 391d9511b72aa9c09c7ad8992ae5df8c89c88c63
2024-03-07 17:37:22 +08:00
lb203 74115d65c4 Update README.md
Former-commit-id: f4c4b3b4b1ce3c1d0b56ba5a109cd6cc57da900d
2024-03-07 17:27:18 +08:00
lb203 a261d31f3a Update README.md
Former-commit-id: 2a286e5aab48aa68e7d3f04df23ec1cdfe71c6de
2024-03-07 17:25:52 +08:00
lb203 da2c84cf50 Update README.md
Former-commit-id: 4a04e71636dabccbef446623674f3a514395b7f9
2024-03-07 17:21:53 +08:00
lb203 3dc3c31fda Update README.md
Former-commit-id: da5507e570a27becab754ba24a0671d669b7e08e
2024-03-07 17:03:17 +08:00
lb203 d14563a9d4 Update README.md
Former-commit-id: 4d83f1feab7d6872aca1c6b41be95950ec3cf465
2024-03-07 17:00:33 +08:00
LinB203 784c246e9c videogpt inference
Former-commit-id: c5d617450a8271ad3b8a5e8f76ee2d77a822e31b
2024-03-07 16:44:19 +08:00
LinB203 11996dc9b4 fixed dit forward
Former-commit-id: 4cd537fab12231f7471107762b02b36217ca42d8
2024-03-07 16:41:31 +08:00
lb203 a03025a8f3 Merge pull request #80 from qqingzheng/videogpt
[docs] Modify README.md and add VQVAE documentation.

Former-commit-id: a92ddf62b7bbd15820401f52edf853dbbb513301
2024-03-07 16:39:28 +08:00
lb203 035214520e Update README.md
Former-commit-id: 89f9fc8eed3ca8cc2434ec227acbb845bfbcff82
2024-03-07 16:37:47 +08:00
lb203 93e98f1ef5 Update README.md
Former-commit-id: 4a893ff6177d59219b6755c4a89764ef54fc7f62
2024-03-07 16:36:16 +08:00
lb203 8b4ad5bf4f Update README.md
Former-commit-id: 0898e4b917f5e245c4394118511bcb7639973817
2024-03-07 16:33:17 +08:00
lb203 afdd297bd5 Update README.md
Former-commit-id: 8b553a4477c64b5452f8424ca6987c0fb8c111ce
2024-03-07 16:31:19 +08:00
lb203 8aa1181ab2 Update README.md
Former-commit-id: 8e9658ed993380abc43b518fcbac21a02d8ba9ba
2024-03-07 16:26:12 +08:00
Chestnut b64293c723 Merge branch 'PKU-YuanGroup:main' into videogpt
Former-commit-id: 339857e1b672a1a688f86c1b8e1d8dae29f10301
2024-03-07 15:03:43 +08:00
qqingzheng d681598e6e [docs] Modify README.md and add VQVAE documentation.
Former-commit-id: a5bfe70f8a8998db969d9c4938fda123aa153bc0
2024-03-07 06:59:40 +00:00
qqingzheng cb86e18190 [refactor] rename train_videogpt.sh
Former-commit-id: e2ea2caacaf5be046f48cf50a53ccafa7a3cc351
2024-03-07 06:15:14 +00:00
lb203 9d9bad199c Merge pull request #76 from luo3300612/dev-rec-video
[fix]: disable gradient computation to save GPU memory when reconstructing a video

Former-commit-id: b8cda6636444080d0cf00452be75a70c1e79f32a
2024-03-07 13:49:39 +08:00
lb203 b6bc9b1455 Update README.md
Former-commit-id: 36e69275598b0cff38fdc5006926e422cda069b3
2024-03-07 13:47:45 +08:00
lb203 5caa97e9a9 Update README.md
Former-commit-id: 4179ed075e68d48094655c38dbde784836fd6e0f
2024-03-07 13:46:58 +08:00
lb203 f507bb4947 Update README.md
Former-commit-id: a1cc0b26a3b156d7708e8df0cfe42aad33b1c946
2024-03-07 13:45:01 +08:00
luo3300612 0379dbb349 [fix]: disable gradient computation to save GPU memory when reconstructing a video
Former-commit-id: 06da17790b53cf593a0c264dda17c76f68d2a541
2024-03-07 13:39:56 +08:00
LinB203 c521999f80 attn_mask with bf16
Former-commit-id: d9667a8c15a5344e3c777197fd6badec15dcddbd
2024-03-07 13:35:01 +08:00
lb203 fcc8009570 Merge pull request #64 from khan-yin/main
[feat]: add sit model

Former-commit-id: acacfd670282b709e6c769717b5295e230861724
2024-03-07 13:29:24 +08:00
lb203 5a7e9c8822 Update README.md
Former-commit-id: 29f24a819fe2d4593c48afe6ca990659784bd776
2024-03-07 12:58:47 +08:00
khan-yin c716b3281e feat: update sample
Former-commit-id: e7dc492b9b4764db28aa31cc1eede914052b4ca6
2024-03-07 04:13:57 +00:00
khan-yin 94b71a5ea0 Merge remote-tracking branch 'upstream/main' into main
Former-commit-id: 7ef9b5e76a91b444f87955c3e0e7ea3d72315396
2024-03-07 04:05:06 +00:00
khan-yin e90373b161 feat: Incorporating SiT sample
Former-commit-id: 91e2931e412f8c96ef748cb2322ee956309fb6aa
2024-03-07 03:40:24 +00:00
lb203 7cf1bcc93d Update README.md
Former-commit-id: b3f6bb7f252505efb0ffd8195003c79206fac313
2024-03-07 11:38:23 +08:00
LinB203 cca84f1c6c denorm fun
Former-commit-id: 05707a07b14295f5e1932d0c10bb84a20a0cd328
2024-03-07 10:40:06 +08:00
lb203 1ebf531e1e Update README.md
Former-commit-id: 1a7178a61547c33a744d032302fb76ee0a3c6671
2024-03-07 10:33:20 +08:00
lb203 c9adc3e9a5 Update README.md
Former-commit-id: d097e40998a3396f74c2bfbcf26d6ace8080d0b6
2024-03-07 10:20:51 +08:00
Kehan Yin 3f19169f55 Merge branch 'PKU-YuanGroup:main' into main
Former-commit-id: a2dea6886ce90634fb7d1c2e0fbfce95a500eefc
2024-03-07 10:04:08 +08:00
lb203 006ad5353a Merge pull request #70 from qqingzheng/videogpt
[refactor] Performed code abstraction on the VideoGPT code and enabled support for accelerate training.

Former-commit-id: 822c1a173f5bc792868fae9e6e4cd7defe2ea183
2024-03-07 09:23:28 +08:00
Chestnut 589da07fec Merge branch 'main' into videogpt
Former-commit-id: 714cc434cea5823fc09ef6fe502b1fe83170cae2
2024-03-07 01:25:58 +08:00
qqingzheng b409ef399a [refactor] added more methods in abstract classes
Former-commit-id: a535e3a06eb766bebe7baa3d9df64b04d04d40e6
2024-03-06 17:23:40 +00:00
lb203 8944956249 Update README.md
Former-commit-id: 604a29d1de43e0fcb90ada4dd97f5f080e37caea
2024-03-07 01:11:28 +08:00
lb203 921d927a27 Update train.sh
Former-commit-id: 66aa3709f80d89a30d0a57c5e033ce1ae46f998f
2024-03-07 01:06:20 +08:00
lb203 06423d785d Update train.sh
Former-commit-id: 6c4957d715eaa489691be21dd0abdb9309660e0d
2024-03-07 00:57:14 +08:00
LinB203 38bf4afdfa support accelerate training
Former-commit-id: b32671c9991b9a9db66cb148822ccb3e8be7dc11
2024-03-07 00:55:35 +08:00
qqingzheng 7a78ef2a97 [refactor] configuration default values
Former-commit-id: d7ca7685d6a8c6d015082dda241bc98ab156770d
2024-03-06 15:59:13 +00:00
qqingzheng 4549eeeab3 Merge branch 'videogpt' of github.com:qqingzheng/Open-Sora-Plan into videogpt
Former-commit-id: ca5f92952c7a0daae6deaca40f6dab1d00a0c0b2
2024-03-06 15:48:26 +00:00
qqingzheng da08e767db [refactor] rename
Former-commit-id: ec317605d24223527b8c2d4b4d2ee876f86abbaf
2024-03-06 15:48:19 +00:00
lb203 0118dd8f1d Merge pull request #69 from luo3300612/dev-rec-video
[fix]: disable gradient computation to save GPU memory when reconstructing a video

Former-commit-id: f9afcda4130e1709d6f201ed310adcc8335480d9
2024-03-06 23:46:10 +08:00
lb203 4efdbc2aca Update README.md
Former-commit-id: 59a595a1e31b156500a944f244afec40988e872d
2024-03-06 23:39:17 +08:00
Chestnut 140eb267d9 Merge branch 'main' into videogpt
Former-commit-id: 0a45d76fdcc3faf167239dfe5dc055b30f3f4e05
2024-03-06 23:36:14 +08:00
LinB203 1b99cf2789 add sample script
Former-commit-id: 9ce847651bb6dba28c51ce5807e222ae2902a78a
2024-03-06 23:23:17 +08:00
lb203 ae2646e213 Merge pull request #74 from Linzy19/2307
[update]: update super_resolution

Former-commit-id: a2586e672bd5d60f1a5e937a796d5f6b6ac1da4d
2024-03-06 23:19:44 +08:00
ZongyingLin 36e0f8813a update
add RGT


Former-commit-id: a491111e48cec588095381eb70ffca94600ca4f7
2024-03-06 23:10:26 +08:00
qqingzheng 0622ea11dc [fix] fix a bug when using ucf101_stride4x4x4
Former-commit-id: f541e723d65ea44ae4560c68908123c7e6781bbc
2024-03-06 15:08:36 +00:00
qqingzheng df625f66cc [refactor] remove test.ipynb
Former-commit-id: e73c3a9e0d41e51a0fd755e57739424213a76312
2024-03-06 15:07:11 +00:00
qqingzheng 107f140b56 [fix] fix a bug when using ucf101_stride4x4x4
Former-commit-id: 9911403c688321db731450d9383bc02923a4be2d
2024-03-06 14:51:35 +00:00
Chestnut 2fd71909a4 Merge branch 'main' into videogpt
Former-commit-id: f8de10e627ce27f1511ddbbc5d6b29f3daf094ab
2024-03-06 22:13:26 +08:00
qqingzheng 2b272863b0 [refactor] adapt to old training code
Former-commit-id: 358d362ecbc100b21ea02895cb0426cf5b386a72
2024-03-06 13:58:13 +00:00
qqingzheng 1bed11aee3 [refactor] reformat videogpt, support training videogpt on accelerate
Former-commit-id: 75385ed47fe60008b61b1bcf89f61ea696d095bc
2024-03-06 13:45:16 +00:00
lb203 5dc0b3cb41 Update README.md
Former-commit-id: 673659495b7503ba402fbcb6c46204499761b4cf
2024-03-06 21:24:58 +08:00
root 80d394d64a [fix]: disable gradient computation to save GPU memory when reconstructing a video
Former-commit-id: b7bf763e7f8800959ed82bb89f41684db216dd8e
2024-03-06 21:04:29 +08:00
lb203 f58d14ba24 Update README.md
Former-commit-id: 398cfcd7b19cf1ba827eb49fc11223f72fa93a85
2024-03-06 20:46:10 +08:00
lb203 d3acf63c4e Update README.md
Former-commit-id: 9fd88eb8ac5122b0999ccb0071f195aaeb0979a0
2024-03-06 20:30:57 +08:00
LinB203 193d7ef7a0 support latte training
Former-commit-id: d471a9414842f422d1200c2fd3f3ddfcf9e97fd6
2024-03-06 20:29:31 +08:00
khan-yin e47f894f69 fix: rename sit model filename
Former-commit-id: f3b2cd7dcf644f680616eee0571db90478fede23
2024-03-06 12:08:04 +00:00
khan-yin 7cee5443e4 fix: remove replicated diffusion module in sit
Former-commit-id: 997cb049185421d303b14d59ab6c0603f1d1a66f
2024-03-06 12:06:55 +00:00
Kehan Yin edfe14e9fa Merge branch 'PKU-YuanGroup:main' into main
Former-commit-id: c7d4f9c2cfaa61e3babc711f13a095ed237978b1
2024-03-06 19:46:26 +08:00
lb203 31f5ce25d4 Update README.md
Former-commit-id: 898d35f15ce0fd124cd48078f0a8af2764bf55d1
2024-03-06 19:33:06 +08:00
LinB203 5a80da7ec1 add vae and fix bug in dit
Former-commit-id: 2dd4254c22ce671d5cd5e7d22a2386b7660fbdc4
2024-03-06 19:32:20 +08:00
lb203 ec0dd141e1 Merge pull request #63 from kabachuha/alternative-attentions
Support for less VRAM heavy alternative attentions (ReBased and Ring attention)

Former-commit-id: c90f9d3c780843edccc72d29f68f5ed7c655a721
2024-03-06 19:28:59 +08:00
khan-yin 54d3069702 [feat]: add sit model
Former-commit-id: 2eda854d6c25254b71d0105a3ed092953a9d6f68
2024-03-06 09:52:27 +00:00
kabachuha abffb358f9 support for ring attention
Former-commit-id: 0533c73b99217e88da121bd0f68e0949bc88016e
2024-03-06 12:24:10 +03:00
kabachuha e1459ff3d0 option to use rebased linear attention
Former-commit-id: f1b39b34df5fbdef14ef3f6a0e6d0972a8570445
2024-03-06 12:18:03 +03:00
Monius f2a826a2f8 add CI for docker autobuild
Former-commit-id: c4e2d5ff1dce8662a42f1ed6a7786855d40173f7
2024-03-06 16:58:03 +08:00
lb203 29409555f3 Update README.md
Former-commit-id: 10ee64c122f28efea39fc1a4542d278683a58ce1
2024-03-06 16:27:06 +08:00
lb203 d878709456 Update README.md
Former-commit-id: 0f2bdd2bcd819d528f2a41a96ba93f5c3c04dc8c
2024-03-06 16:23:53 +08:00
qqingzheng 3810227fd6 first commit
Former-commit-id: ec3481bc35a8120d519082dc76585428f0896470
2024-03-06 08:16:58 +00:00
lb203 8b3920706a Update README.md
Former-commit-id: f1aaeca19bc9e4304961c8c903f9bf03cb42fd9f
2024-03-06 16:13:26 +08:00
lb203 932c7fd80b Create LICENSE
Former-commit-id: 369fe085e91e55b9be6165da4a1eb7f9c0cf5a7f
2024-03-06 16:12:41 +08:00
lb203 cfb34d011b Delete LICENSE.txt
Former-commit-id: b10d8b2c6620e3a8412527bd03ea0adc26c981a6
2024-03-06 16:11:05 +08:00
lb203 0147265637 Update README.md
Former-commit-id: 437bc5c68b5c7a85865ad97d544c4818be45cb52
2024-03-06 15:38:04 +08:00
LinB203 cf40ce1e7a add latte
Former-commit-id: 01ba10e5d65c9d2d957472fbe7958e1b07a971e9
2024-03-06 15:28:20 +08:00
yanyang1024 4f608035e2 Update README.md
Former-commit-id: c5a79c7e993d9e6a77c6c70efa1c0ac2bb3a063b
2024-03-06 12:12:58 +08:00
yanyang1024 70ee26f0ec fix_readme_requirements
Former-commit-id: eb5a5c69d12cc412c12a4fea7f00baafd02701ca
2024-03-06 12:05:21 +08:00
lb203 e14c09b3da Update README.md
Former-commit-id: 1436d8feea7f82d3d747c632a68c47b3ea74258f
2024-03-05 22:11:38 +08:00
lb203 3ecbc2eab6 Merge pull request #50 from yunyangge/interpolation
[feat]: frame_interpolation

Former-commit-id: 59ac3fd0b48bcd89f2b1f816b74a4e17c83328b2
2024-03-05 22:10:18 +08:00
yunyangge 21b257c25e [feat]: frame_interpolation
Former-commit-id: 711fc4bba80d5f051f712e3460e2c0685117cf68
2024-03-05 21:53:22 +08:00
lb203 458d1d7770 Update README.md
Former-commit-id: 3e3d674134498512b1ef1b0dd72d6e19afd6ef65
2024-03-05 15:52:05 +08:00
lb203 34c1d8d2d0 Update README.md
Former-commit-id: 164e76cf496b1157eebea299ed58ca6c1799d70d
2024-03-05 15:50:45 +08:00
lb203 192abaeb9e Update README.md
Former-commit-id: be7ecc3d56f7289052fc66d535038e7111f7ac84
2024-03-05 15:46:32 +08:00
lb203 8b2f93325f Update Contribution_Guidelines.md
Former-commit-id: dc87e459ca5a8a090800331111118acb65214f65
2024-03-05 15:01:08 +08:00
lb203 2ca658809a Update README.md
Former-commit-id: 6ecaba8e03d7c7f99c068a57c923f11c5456a8c9
2024-03-05 14:54:38 +08:00
lb203 86332d9c0a Update README.md
Former-commit-id: 1dd9c945be2a170d87d568bb00e1ec30e7be592b
2024-03-05 13:46:09 +08:00
lb203 f26658b7de Merge pull request #42 from mio2333/patch-2
Update requirements.txt

Former-commit-id: 0c76835e4a4681ee317de45a18549dc0d9b9fef9
2024-03-05 13:39:23 +08:00
Junwu Zhang 5a56c37298 reformat code
Former-commit-id: ce23b280fbb419c08347ebcbf1ed4075bdf9af60
2024-03-05 05:38:23 +00:00
Junwu Zhang e01856047d Update README.md
Former-commit-id: 351be764d8648ac60cf170c3fe4e6ed9c544c90e
2024-03-05 13:34:35 +08:00
Junwu Zhang c14ef52685 Rename data.md to Data.md
Former-commit-id: 27484f7a2d8634bd1cd9781a2f10afa7f304f5cb
2024-03-05 13:31:50 +08:00
Junwu Zhang c62e40c802 Update README.md
Former-commit-id: c652acd1fc5a003b4d6eefcfb83dfb310f5b5e81
2024-03-05 13:31:23 +08:00
Junwu Zhang 370be675e8 Create Contribution_Guidelines.md
Former-commit-id: 897dbd63dc6f53095e1dda3f10a86dd40e4bd263
2024-03-05 13:28:32 +08:00
mio e72ef94a7f Update requirements.txt
As in previous step,the three packages have been installed, these packages should be deleted here because it can not be installed by " pip install -r requirements.txt" command and will abort the command

Former-commit-id: 0655fa28cf8bb761f1a3d2def62f44197b257b2c
2024-03-05 13:22:28 +08:00
Junwu Zhang e777c2e617 Update README.md
Former-commit-id: f295f4623448dd4ffa9be747a5cc6733cf48c663
2024-03-05 12:51:07 +08:00
Zhenyu Tang e4eee7bc20 Update README.md
Former-commit-id: 0ff9314bddef1dfbbf429726ed4240b823b221b3
2024-03-04 22:45:06 +08:00
Zhenyu Tang b2dbc2099e Update README.md
Former-commit-id: e27cc5d49a19e9580ba1a29d948051b918aede74
2024-03-04 22:41:37 +08:00
YuanLi 24006388be Update README.md
Former-commit-id: 1cd7d8b5df34ab9f61d9346b2647b643f78a8f5d
2024-03-04 22:41:04 +08:00
Zhenyu Tang cf5bafaee6 Update README.md
Former-commit-id: 5fd1e6afe1277ccffdfeaf507477d944a6a35f03
2024-03-04 22:28:08 +08:00
Zhenyu Tang 028c52fd9f Update README.md
Former-commit-id: fb46c80c1a2663c09849472e93b2da6cd068a98c
2024-03-04 22:24:21 +08:00
Zhenyu Tang b7dfe57c17 Update README.md
Former-commit-id: aee63a3a8422e029f1550e47825ca4beb916874b
2024-03-04 22:23:44 +08:00
Junwu Zhang a02832797a reformat code
Former-commit-id: 7a3dccf2c738eaf535e62996211c8bc5f3de5683
2024-03-04 14:06:39 +00:00
291 changed files with 11747 additions and 15983 deletions
+29 -112
View File
@@ -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
-10
View File
@@ -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
-89
View File
@@ -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.
+17 -183
View File
@@ -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.
+101 -65
View File
@@ -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.
-10
View File
@@ -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.
-47
View File
@@ -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()
+85
View File
@@ -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')
+83
View File
@@ -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()
+312
View File
@@ -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
+591
View File
@@ -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))
+719
View File
@@ -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)
+105
View File
@@ -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
+308
View File
@@ -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
+687
View File
@@ -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)
+148
View File
@@ -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
+123
View File
@@ -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
+39
View File
@@ -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
+654
View File
@@ -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)
+124
View File
@@ -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
+756
View File
@@ -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)
+155
View File
@@ -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)
+107
View File
@@ -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)
+207
View File
@@ -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)
+482
View File
@@ -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)
+244
View File
@@ -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
+298
View File
@@ -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)
+293
View File
@@ -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)
+328
View File
@@ -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")
+23
View File
@@ -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
+52
View File
@@ -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()
+471
View File
@@ -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)
# &quot;
caption = re.sub(r'&quot;?', '', caption)
# &amp
caption = re.sub(r'&amp', '', 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))
+274
View File
@@ -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)
-31
View File
@@ -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)
-3
View File
@@ -1,3 +0,0 @@
from fastvideo2.cli import main
raise SystemExit(main())
-213
View File
@@ -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()
-92
View File
@@ -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())
-4
View File
@@ -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"]
-3
View File
@@ -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())
-269
View File
@@ -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()
-178
View File
@@ -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)
-18
View File
@@ -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.
@@ -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"
}
@@ -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"
}
@@ -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"
}
@@ -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
}
@@ -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
}
@@ -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.

Some files were not shown because too many files have changed in this diff Show More