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
201 changed files with 6905 additions and 25395 deletions
-29
View File
@@ -1,29 +0,0 @@
name: 🐞 Bug report
description: Create a report to help us reproduce and fix the bug
title: "[Bug] "
labels: ['Bug']
body:
- type: textarea
attributes:
label: Environment
description: |
Please share your environment with us. You can run the command **python fastvideo/utils/env_utils.py** and copy-paste its output below.
placeholder: FastVideo version, platform, python version, cuda version...
validations:
required: true
- type: textarea
attributes:
label: Describe the bug
description: A clear and concise description of what the bug is.
validations:
required: true
- type: textarea
attributes:
label: Reproduction
description: |
What command or script did you run? Which **model** are you using?
placeholder: |
A placeholder for the command.
validations:
required: true
@@ -1,17 +0,0 @@
name: 🚀 Feature request
description: Suggest an idea for this project
title: "[Feature] "
body:
- type: textarea
attributes:
label: Motivation
description: |
A clear and concise description of the motivation of the feature.
validations:
required: true
- type: textarea
attributes:
label: Related resources
description: |
If there is an official code release or third-party implementations, please also provide the information here, which would be very helpful.
-27
View File
@@ -1,27 +0,0 @@
name: Run Tests
on:
push:
branches: [ main ]
pull_request:
branches: [ main ]
jobs:
test:
runs-on: ubuntu-latest
steps:
- name: Check out repository
uses: actions/checkout@v3
- name: Set up Python
uses: actions/setup-python@v4
with:
python-version: '3.12' # or any version you need
- name: Install dependencies
run: |
pip install --upgrade pip
pip install -e .
- name: Run Pytest
run: |
pytest
+1
View File
@@ -20,6 +20,7 @@ wandb/
*.pt
cache_dir/
wandb/
test*
sample_video*
sample_image*
512*
+17 -197
View File
@@ -1,201 +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.
Copyright [2023] Lightning AI
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.
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.
+93 -140
View File
@@ -1,155 +1,108 @@
<div align="center">
<img src=assets/logo.jpg width="30%"/>
</div>
# Fast Video
This is currently based on Open-Sora-1.2.0: https://github.com/PKU-YuanGroup/Open-Sora-Plan/tree/294993ca78bf65dec1c3b6fb25541432c545eda9
FastVideo is a lightweight framework for accelerating large video diffusion models.
https://github.com/user-attachments/assets/5fbc4596-56d6-43aa-98e0-da472cf8e26c
<p align="center">
🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank">FastMochi</a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank">FastHunyuan</a> | 🎮 <a href="https://discord.gg/REBzDQTWWt" target="_blank"> Discord </a> | 🕹️ <a href="https://replicate.com/lucataco/fast-hunyuan-video" target="_blank"> Replicate </a>
</p>
FastVideo currently offers: (with more to come)
- FastHunyuan and FastMochi: consistency distilled video diffusion models for 8x inference speedup.
- First open distillation recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
- Scalable training with FSDP, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
- Memory efficient finetuning with LoRA, precomputed latent, and precomputed text embeddings.
Dev in progress and highly experimental.
## 🎥 More Demos
Fast-Hunyuan comparison with original Hunyuan, achieving an 8X diffusion speed boost with the FastVideo framework.
https://github.com/user-attachments/assets/064ac1d2-11ed-4a0c-955b-4d412a96ef30
Comparison between OpenAI Sora, original Hunyuan and FastHunyuan
https://github.com/user-attachments/assets/d323b712-3f68-42b2-952b-94f6a49c4836
Comparison between original FastHunyuan, LLM-INT8 quantized FastHunyuan and NF4 quantized FastHunyuan
https://github.com/user-attachments/assets/cf89efb5-5f68-4949-a085-f41c1ef26c94
## Change Log
- ```2024/12/25```: Enable single 4090 inference for `FastHunyuan`, please rerun the installation steps to update the environment.
- ```2024/12/17```: `FastVideo` v1.0 is released.
## 🔧 Installation
The code is tested on Python 3.10.0, CUDA 12.1 and H100.
## Envrironment
Change the index-url cuda version according to your system.
```
./env_setup.sh fastvideo
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
```
## 🚀 Inference
### Inference FastHunyuan on single RTX4090
We now support NF4 and LLM-INT8 quantized inference using BitsAndBytes for FastHunyuan. With NF4 quantization, inference can be performed on a single RTX 4090 GPU, requiring just 20GB of VRAM.
```bash
# Download the model weight
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan-diffusers --local_dir=data/FastHunyuan-diffusers --repo_type=model
# CLI inference
bash scripts/inference/inference_diffusers_hunyuan.sh
```
For more information about the VRAM requirements for BitsAndBytes quantization, please refer to the table below (timing measured on an H100 GPU):
| Configuration | Memory to Init Transformer | Peak Memory After Init Pipeline (Denoise) | Diffusion Time | End-to-End Time |
|--------------------------------|----------------------------|--------------------------------------------|----------------|-----------------|
| BF16 + Pipeline CPU Offload | 23.883G | 33.744G | 81s | 121.5s |
| INT8 + Pipeline CPU Offload | 13.911G | 27.979G | 88s | 116.7s |
| NF4 + Pipeline CPU Offload | 9.453G | 19.26G | 78s | 114.5s |
For improved quality in generated videos, we recommend using a GPU with 80GB of memory to run the BF16 model with the original Hunyuan pipeline. To execute the inference, use the following section:
### FastHunyuan
```bash
# Download the model weight
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan --local_dir=data/FastHunyuan --repo_type=model
# CLI inference
bash scripts/inference/inference_hunyuan.sh
```
You can also inference FastHunyuan in the [official Hunyuan github](https://github.com/Tencent/HunyuanVideo).
### FastMochi
```bash
# Download the model weight
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastMochi-diffusers --local_dir=data/FastMochi-diffusers --repo_type=model
# CLI inference
bash scripts/inference/inference_mochi_sp.sh
pip install -e . && pip install -e ".[train]"
sudo apt-get update && apt install screen && pip install watch gpustat
```
## 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)
## 🎯 Distill
Our distillation recipe is based on [Phased Consistency Model](https://github.com/G-U-N/Phased-Consistency-Model). We did not find significant improvement using multi-phase distillation, so we keep the one phase setup similar to the original latent consistency model's recipe.
We use the [MixKit](https://huggingface.co/datasets/LanguageBind/Open-Sora-Plan-v1.1.0/tree/main/all_mixkit) dataset for distillation. To avoid running the text encoder and VAE during training, we preprocess all data to generate text embeddings and VAE latents.
Preprocessing instructions can be found [data_preprocess.md](docs/data_preprocess.md). For convenience, we also provide preprocessed data that can be downloaded directly using the following command:
```bash
python scripts/huggingface/download_hf.py --repo_id=FastVideo/HD-Mixkit-Finetune-Hunyuan --local_dir=data/HD-Mixkit-Finetune-Hunyuan --repo_type=dataset
```
Next, download the original model weights with:
```bash
python scripts/huggingface/download_hf.py --repo_id=genmo/mochi-1-preview --local_dir=data/mochi --repo_type=model # original mochi
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model # original hunyuan
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 ../..
```
To launch the distillation process, use the following commands:
```
bash scripts/distill/distill_mochi.sh # for mochi
bash scripts/distill/distill_hunyuan.sh # for hunyuan
```
We also provide an optional script for distillation with adversarial loss, located at `fastvideo/distill_adv.py`. Although we tried adversarial loss, we did not observe significant improvements.
## Finetune
### ⚡ Full Finetune
Ensure your data is prepared and preprocessed in the format specified in [data_preprocess.md](docs/data_preprocess.md). For convenience, we also provide a mochi preprocessed Black Myth Wukong data that can be downloaded directly:
```bash
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Mochi-Black-Myth --local_dir=data/Mochi-Black-Myth --repo_type=dataset
```
Download the original model weights as specificed in [Distill Section](#-distill):
Then you can run the finetune with:
```
bash scripts/finetune/finetune_mochi.sh # for mochi
```
**Note that for finetuning, we did not tune the hyperparameters in the provided script**
### ⚡ Lora Finetune
Currently, we only provide Lora Finetune for Mochi model, the command for Lora Finetune is
```
bash scripts/finetune/finetune_mochi_lora.sh
```
### Minimum Hardware Requirement
- 40 GB GPU memory each for 2 GPUs with lora
- 30 GB GPU memory each for 2 GPUs with CPU offload and lora.
### Finetune with Both Image and Video
Our codebase support finetuning with both image and video.
```bash
bash scripts/finetune/finetune_hunyuan.sh
bash scripts/finetune/finetune_mochi_lora_mix.sh
```
For Image-Video Mixture Fine-tuning, make sure to enable the --group_frame option in your script.
## 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不行
## 📑 Development Plan
## Experiments
Scripts are located at scripts/experiment_N.sh
- More distillation methods
- [ ] Add Distribution Matching Distillation
- More models support
- [ ] Add CogvideoX model
- Code update
- [ ] fp8 support
- [ ] faster load model and save model support
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
## Acknowledgement
We learned and reused code from the following projects: [PCM](https://github.com/G-U-N/Phased-Consistency-Model), [diffusers](https://github.com/huggingface/diffusers), [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan), and [xDiT](https://github.com/xdit-project/xDiT).
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
We thank MBZUAI and Anyscale for their support throughout this project.
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.
Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.3 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.3 MiB

Binary file not shown.
Binary file not shown.
Binary file not shown.

Before

Width:  |  Height:  |  Size: 22 MiB

Binary file not shown.
BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 149 KiB

-8
View File
@@ -1,8 +0,0 @@
Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting.
A lone hiker stands atop a towering cliff, silhouetted against the vast horizon. The rugged landscape stretches endlessly beneath, its earthy tones blending into the soft blues of the sky. The scene captures the spirit of exploration and human resilience. High angle, dynamic framing, with soft natural lighting emphasizing the grandeur of nature.
A hand with delicate fingers picks up a bright yellow lemon from a wooden bowl filled with lemons and sprigs of mint against a peach-colored background. The hand gently tosses the lemon up and catches it, showcasing its smooth texture. A beige string bag sits beside the bowl, adding a rustic touch to the scene. Additional lemons, one halved, are scattered around the base of the bowl. The even lighting enhances the vibrant colors and creates a fresh, inviting atmosphere.
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.
A superintelligent humanoid robot waking up. The robot has a sleek metallic body with futuristic design features. Its glowing red eyes are the focal point, emanating a sharp, intense light as it powers on. The scene is set in a dimly lit, high-tech laboratory filled with glowing control panels, robotic arms, and holographic screens. The setting emphasizes advanced technology and an atmosphere of mystery. The ambiance is eerie and dramatic, highlighting the moment of awakening and the robots immense intelligence. Photorealistic style with a cinematic, dark sci-fi aesthetic. Aspect ratio: 16:9 --v 6.1
fox in the forest close-up quickly turned its head to the left
Man walking his dog in the woods on a hot sunny day
A majestic lion strides across the golden savanna, its powerful frame glistening under the warm afternoon sun. The tall grass ripples gently in the breeze, enhancing the lion's commanding presence. The tone is vibrant, embodying the raw energy of the wild. Low angle, steady tracking shot, cinematic.
-24
View File
@@ -1,24 +0,0 @@
# Configuration for Cog ⚙️
# Reference: https://cog.run/yaml
build:
gpu: true
cuda: "12.1"
python_version: "3.10"
python_packages:
- "torch==2.4.0"
- "torchvision"
- "ninja==1.11.1.3"
- "transformers==4.46.1"
- "git+https://github.com/huggingface/diffusers.git@bf64b32652a63a1865a0528a73a13652b201698b"
- "accelerate==1.0.1"
- "safetensors==0.4.5"
- "peft==0.13.2"
- "packaging==24.2"
- "git+https://github.com/hao-ai-lab/FastVideo"
run:
- FLASH_ATTENTION_SKIP_CUDA_BUILD=TRUE pip install flash-attn --no-build-isolation
- curl -o /usr/local/bin/pget -L "https://github.com/replicate/pget/releases/latest/download/pget_$(uname -s)_$(uname -m)" && chmod +x /usr/local/bin/pget
predict: "predict.py:Predictor"
-204
View File
@@ -1,204 +0,0 @@
import gradio as gr
import torch
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import export_to_video
from fastvideo.distill.solver import PCMFMScheduler
import tempfile
import os
import argparse
def init_args():
parser = argparse.ArgumentParser()
parser.add_argument("--prompts", nargs="+", default=[])
parser.add_argument("--num_frames", type=int, default=25)
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=8)
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("--scheduler_type", type=str, default="pcm_linear_quadratic")
parser.add_argument("--lora_checkpoint_dir", type=str, default=None)
parser.add_argument("--shift", type=float, default=8.0)
parser.add_argument("--num_euler_timesteps", type=int, default=50)
parser.add_argument("--linear_threshold", type=float, default=0.1)
parser.add_argument("--linear_range", type=float, default=0.75)
parser.add_argument("--cpu_offload", action="store_true")
return parser.parse_args()
def load_model(args):
device = "cuda" if torch.cuda.is_available() else "cpu"
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:
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(device)
# if args.cpu_offload:
pipe.enable_sequential_cpu_offload()
return pipe
def generate_video(
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
num_inference_steps,
randomize_seed=False,
):
if randomize_seed:
seed = torch.randint(0, 1000000, (1,)).item()
generator = torch.Generator(device="cuda").manual_seed(seed)
if not use_negative_prompt:
negative_prompt = None
with torch.autocast("cuda", dtype=torch.bfloat16):
output = pipe(
prompt=[prompt],
negative_prompt=negative_prompt,
height=height,
width=width,
num_frames=num_frames,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
generator=generator,
).frames[0]
output_path = os.path.join(tempfile.mkdtemp(), "output.mp4")
export_to_video(output, output_path, fps=30)
return output_path, seed
examples = [
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand’s movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.",
]
args = init_args()
pipe = load_model(args)
print("load model successfully")
with gr.Blocks() as demo:
gr.Markdown("# Fastvideo Mochi Video Generation Demo")
with gr.Group():
with gr.Row():
prompt = gr.Text(
label="Prompt",
show_label=False,
max_lines=1,
placeholder="Enter your prompt",
container=False,
)
run_button = gr.Button("Run", scale=0)
result = gr.Video(label="Result", show_label=False)
with gr.Accordion("Advanced options", open=False):
with gr.Group():
with gr.Row():
height = gr.Slider(
label="Height",
minimum=256,
maximum=1024,
step=32,
value=args.height,
)
width = gr.Slider(
label="Width", minimum=256, maximum=1024, step=32, value=args.width
)
with gr.Row():
num_frames = gr.Slider(
label="Number of Frames",
minimum=21,
maximum=163,
value=args.num_frames,
)
guidance_scale = gr.Slider(
label="Guidance Scale",
minimum=1,
maximum=12,
value=args.guidance_scale,
)
num_inference_steps = gr.Slider(
label="Inference Steps",
minimum=4,
maximum=100,
value=args.num_inference_steps,
)
with gr.Row():
use_negative_prompt = gr.Checkbox(
label="Use negative prompt", value=False
)
negative_prompt = gr.Text(
label="Negative prompt",
max_lines=1,
placeholder="Enter a negative prompt",
visible=False,
)
seed = gr.Slider(
label="Seed", minimum=0, maximum=1000000, step=1, value=args.seed
)
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
seed_output = gr.Number(label="Used Seed")
gr.Examples(examples=examples, inputs=prompt)
use_negative_prompt.change(
fn=lambda x: gr.update(visible=x),
inputs=use_negative_prompt,
outputs=negative_prompt,
)
run_button.click(
fn=generate_video,
inputs=[
prompt,
negative_prompt,
use_negative_prompt,
seed,
guidance_scale,
num_frames,
height,
width,
num_inference_steps,
randomize_seed,
],
outputs=[result, seed_output],
)
if __name__ == "__main__":
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
-68
View File
@@ -1,68 +0,0 @@
## 🧱 Data Preprocess
To save GPU memory, we precompute text embeddings and VAE latents to eliminate the need to load the text encoder and VAE during training.
We provide a sample dataset to help you get started. Download the source media using the following command:
```bash
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Image-Vid-Finetune-Src --local_dir=data/Image-Vid-Finetune-Src --repo_type=dataset
```
To preprocess the dataset for fine-tuning or distillation, run:
```
bash scripts/preprocess/preprocess_mochi_data.sh # for mochi
bash scripts/preprocess/preprocess_hunyuan_data.sh # for hunyuan
```
The preprocessed dataset will be stored in `Image-Vid-Finetune-Mochi` or `Image-Vid-Finetune-HunYuan` correspondingly.
### Process your own dataset
If you wish to create your own dataset for finetuning or distillation, please structure you video dataset in the following format:
path_to_dataset_folder/
├── media/
│ ├── 0.jpg
│ ├── 1.mp4
│ ├── 2.jpg
├── video2caption.json
└── merge.txt
Format the JSON file as a list, where each item represents a media source:
For image media,
```
{
"path": "0.jpg",
"cap": ["captions"]
}
```
For video media,
```
{
"path": "1.mp4",
"resolution": {
"width": 848,
"height": 480
},
"fps": 30.0,
"duration": 6.033333333333333,
"cap": [
"caption"
]
}
```
Use a txt file (merge.txt) to contain the source folder for media and the JSON file for meta information:
```
path_to_media_source_foder,path_to_json_file
```
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/preprocess_****_data.sh` accordingly and run:
```
bash scripts/preprocess/preprocess_****_data.sh
```
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
-10
View File
@@ -1,10 +0,0 @@
#!/bin/bash
# install torch
pip install torch==2.5.0 torchvision --index-url https://download.pytorch.org/whl/cu121
# install FA2 and diffusers
pip install packaging ninja && pip install flash-attn==2.7.0.post2 --no-build-isolation
# install fastvideo
pip install -e .
@@ -1,156 +0,0 @@
import argparse
import torch
from accelerate.logging import get_logger
from fastvideo.models.mochi_hf.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
from fastvideo.utils.load import load_text_encoder, load_vae
from diffusers.video_processor import VideoProcessor
from tqdm import tqdm
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
)
videoprocessor = VideoProcessor(vae_scale_factor=8)
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)
text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
vae.enable_tiling()
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 tqdm(enumerate(train_dataloader), disable=local_rank != 0):
with torch.inference_mode():
with torch.autocast("cuda", dtype=autocast_type):
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(
prompt=data["caption"],
)
if args.vae_debug:
latents = data["latents"]
video = vae.decode(latents.to(device), return_dict=False)[0]
video = videoprocessor.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=fps)
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")
parser.add_argument("--model_type", type=str, default="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)
@@ -1,129 +0,0 @@
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
from fastvideo.utils.load import load_vae
from tqdm import tqdm
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)
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, autocast_type, fps = load_vae(args.model_type, args.model_path)
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 tqdm(enumerate(train_dataloader), disable=local_rank != 0):
with torch.inference_mode():
with torch.autocast("cuda", dtype=autocast_type):
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("--model_type", type=str, default="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)
@@ -1,81 +0,0 @@
import argparse
import torch
from accelerate.logging import get_logger
from fastvideo.models.mochi_hf.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
from fastvideo.utils.load import load_text_encoder, load_vae
from diffusers.video_processor import VideoProcessor
from tqdm import tqdm
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
)
text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
autocast_type = torch.float16 if args.model_type == "hunyuan" else torch.bfloat16
# output_dir/validation/prompt_attention_mask
# output_dir/validation/prompt_embed
os.makedirs(os.path.join(args.output_dir, "validation"), exist_ok=True)
os.makedirs(
os.path.join(args.output_dir, "validation", "prompt_attention_mask"),
exist_ok=True,
)
os.makedirs(
os.path.join(args.output_dir, "validation", "prompt_embed"), exist_ok=True
)
json_data = []
with open(args.validation_prompt_txt, "r", encoding="utf-8") as file:
lines = file.readlines()
prompts = [line.strip() for line in lines]
for prompt in prompts:
with torch.inference_mode():
with torch.autocast("cuda", dtype=autocast_type):
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(
prompt
)
file_name = prompt.split(".")[0]
prompt_embed_path = os.path.join(
args.output_dir, "validation", "prompt_embed", f"{file_name}.pt"
)
prompt_attention_mask_path = os.path.join(
args.output_dir,
"validation",
"prompt_attention_mask",
f"{file_name}.pt",
)
torch.save(prompt_embeds[0], prompt_embed_path)
torch.save(prompt_attention_mask[0], prompt_attention_mask_path)
print(f"sample {file_name} saved")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# dataset & dataloader
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--model_type", type=str, default="mochi")
parser.add_argument("--validation_prompt_txt", type=str)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
args = parser.parse_args()
main(args)
+49 -75
View File
@@ -4,51 +4,31 @@ 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,
)
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.0 * x - 1.0)
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,
]
)
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,
)
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)
@@ -57,34 +37,32 @@ if __name__ == "__main__":
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,
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)
@@ -92,9 +70,7 @@ if __name__ == "__main__":
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
]
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:
@@ -105,7 +81,5 @@ if __name__ == "__main__":
continue
assert caps[0] is not None and len(caps[0]) > 0
print(num, zero)
import ipdb
ipdb.set_trace()
print("end")
import ipdb;ipdb.set_trace()
print('end')
+26 -66
View File
@@ -4,11 +4,13 @@ import json
import os
import random
class LatentDataset(Dataset):
def __init__(
self, json_path, num_latent_t, cfg_rate,
):
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
@@ -16,10 +18,8 @@ class LatentDataset(Dataset):
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.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'])
@@ -28,44 +28,27 @@ class LatentDataset(Dataset):
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
]
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,
)
latent = latent.squeeze(0)[:, -self.num_latent_t :]
# 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,
)
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
@@ -76,48 +59,25 @@ def latent_collate_function(batch):
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
]
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
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
)
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()
print(latent.shape, prompt_embed.shape, latent_attn_mask.shape, prompt_attention_mask.shape)
import pdb; pdb.set_trace()
+88 -123
View File
@@ -11,12 +11,17 @@ 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
from fastvideo.utils.logging_ import main_print
logger = get_logger(__name__)
class SingletonMeta(type):
"""
这是一个元类,用于创建单例类。
"""
_instances = {}
def __call__(cls, *args, **kwargs):
@@ -48,7 +53,7 @@ class DataSetProg(metaclass=SingletonMeta):
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]
self.worker_elements[i] = self.elements[start: end]
def get_item(self, work_info):
if work_info is None:
@@ -56,20 +61,18 @@ class DataSetProg(metaclass=SingletonMeta):
else:
worker_id = work_info.id
idx = self.worker_elements[worker_id][
self.n_used_elements[worker_id] % len(self.worker_elements[worker_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):
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):
@@ -92,11 +95,11 @@ class T2V_dataset(Dataset):
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):
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
@@ -114,38 +117,39 @@ class T2V_dataset(Dataset):
return dataset_prog.n_elements
def __getitem__(self, idx):
data = self.get_data(idx)
return data
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"):
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"]
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"
)
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)
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}"
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"]
text = dataset_prog.cap_list[idx]['cap']
if not isinstance(text, list):
text = [text]
text = [random.choice(text)]
@@ -154,70 +158,51 @@ class T2V_dataset(Dataset):
text_tokens_and_mask = self.tokenizer(
text,
max_length=self.text_max_length,
padding="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,
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 = 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]
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 = 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 = image_data['cap'] if isinstance(image_data['cap'], list) else [image_data['cap']]
caps = [random.choice(caps)]
text = caps
text = text_preprocessing(caps, support_Chinese=self.support_Chinese)
input_ids, cond_mask = [], []
text = text[0] if random.random() > self.cfg else ""
text = text if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
text,
max_length=self.text_max_length,
padding="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"],
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
@@ -228,98 +213,78 @@ class T2V_dataset(Dataset):
cnt_movie = 0
cnt_img = 0
for i in cap_list:
path = i["path"]
cap = i.get("cap", None)
path = i['path']
cap = i.get('cap', None)
# ======no caption=====
if cap is None:
cnt_no_cap += 1
continue
if path.endswith(".mp4"):
if path.endswith('.mp4'):
# ======no fps and duration=====
duration = i.get("duration", None)
fps = i.get("fps", None)
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)
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
):
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"]
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,
)
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)
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)
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
):
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[begin_index: end_index]
# frame_indices = frame_indices[:self.num_frames] # head crop
i["sample_frame_index"] = frame_indices.tolist()
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
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"])
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"
)
raise NameError(f"Unknown file extention {path.split('.')[-1]}, only support .mp4 for video and .jpg for image")
# import ipdb;ipdb.set_trace()
main_print(
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)}"
)
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()
@@ -329,19 +294,19 @@ class T2V_dataset(Dataset):
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
]
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:
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"])
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
+66 -108
View File
@@ -32,9 +32,7 @@ def center_crop_arr(pil_image, image_size):
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]
)
return Image.fromarray(arr[crop_y: crop_y + image_size, crop_x: crop_x + image_size])
def crop(clip, i, j, h, w):
@@ -44,37 +42,21 @@ def crop(clip, i, j, h, w):
"""
if len(clip.size()) != 4:
raise ValueError("clip should be a 4D tensor")
return clip[..., i : i + h, j : j + w]
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,
)
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}"
)
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,
)
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"):
@@ -125,10 +107,11 @@ def center_crop_using_short_edge(clip):
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
@@ -138,16 +121,15 @@ def center_crop_th_tw(clip, th, tw, top_crop):
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)
@@ -177,9 +159,7 @@ def normalize_video(clip):
"""
_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)
)
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
@@ -239,9 +219,7 @@ class RandomCropVideo:
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)}"
)
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
@@ -257,7 +235,7 @@ class RandomCropVideo:
class SpatialStrideCropVideo:
def __init__(self, stride):
self.stride = stride
self.stride = stride
def __call__(self, clip):
"""
@@ -280,15 +258,17 @@ class SpatialStrideCropVideo:
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,
skip_low_resolution=False,
interpolation_mode="bilinear",
):
self.size = size
self.skip_low_resolution = skip_low_resolution
@@ -311,28 +291,27 @@ class LongSideResizeVideo:
else:
h = int(h * self.size / w)
w = self.size
resize_clip = resize(
clip, target_size=(h, w), interpolation_mode=self.interpolation_mode
)
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",
self,
size,
top_crop=False,
interpolation_mode="bilinear",
):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}"
)
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
@@ -346,15 +325,10 @@ class CenterCropResizeVideo:
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
)
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,
)
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:
@@ -362,19 +336,19 @@ class CenterCropResizeVideo:
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",
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}"
)
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
@@ -389,9 +363,7 @@ class UCFCenterCropVideo:
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_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
@@ -400,18 +372,18 @@ class UCFCenterCropVideo:
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",
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}"
)
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
@@ -426,13 +398,13 @@ class KineticsRandomCropResizeVideo:
class CenterCropVideo:
def __init__(
self, size, interpolation_mode="bilinear",
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}"
)
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
@@ -544,7 +516,6 @@ class TemporalRandomCrop(object):
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.
@@ -559,16 +530,13 @@ class DynamicSampleDuration(object):
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_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__":
if __name__ == '__main__':
from torchvision import transforms
import torchvision.io as io
import numpy as np
@@ -576,20 +544,18 @@ if __name__ == "__main__":
import os
vframes, aframes, info = io.read_video(
filename="./v_Archery_g01_c03.avi", pts_unit="sec", output_format="TCHW"
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
),
]
)
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
@@ -603,9 +569,7 @@ if __name__ == "__main__":
# 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
)
frame_indice = np.linspace(start_frame_ind, end_frame_ind - 1, target_video_len, dtype=int)
print(frame_indice)
select_vframes = vframes[frame_indice]
@@ -616,18 +580,12 @@ if __name__ == "__main__":
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
)
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)
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),
)
save_image(select_vframes_trans[i], os.path.join('./test000', '%04d.png' % i), normalize=True,
value_range=(-1, 1))
+363 -547
View File
File diff suppressed because it is too large Load Diff
+17 -11
View File
@@ -23,6 +23,7 @@ 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__()
@@ -47,9 +48,9 @@ class DiscriminatorHead(nn.Module):
def forward(self, x):
b, twh, c = x.shape
t = twh // (30 * 53)
x = x.view(-1, 30 * 53, c)
x = x.view(-1, 30 *53, c)
x = x.permute(0, 2, 1)
x = x.view(b * t, c, 30, 53)
x = x.view(b*t, c, 30, 53)
x = self.conv1(x)
x = self.conv2(x) + x
x = self.conv_out(x)
@@ -57,11 +58,15 @@ class DiscriminatorHead(nn.Module):
class Discriminator(nn.Module):
def __init__(
self, stride=8, num_h_per_head=1, adapter_channel_dims=[3072], total_layers=48,
self,
stride = 8,
num_h_per_head=1,
adapter_channel_dims=[3072],
):
super().__init__()
adapter_channel_dims = adapter_channel_dims * (total_layers // stride)
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)
@@ -77,23 +82,24 @@ class Discriminator(nn.Module):
]
)
def forward(self, features):
outputs = []
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
assert len(features) == len(self.heads)
for i in range(0, len(features)):
for h in self.heads[i]:
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])
out=h(features[i])
outputs.append(out)
return outputs
+13 -9
View File
@@ -8,7 +8,7 @@ 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.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
from fastvideo.model.pipeline_mochi import linear_quadratic_schedule
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -17,14 +17,13 @@ logger = logging.get_logger(__name__) # pylint: disable=invalid-name
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
@@ -35,14 +34,13 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
shift: float = 1.0,
pcm_timesteps: int = 50,
linear_quadratic=False,
linear_quadratic_threshold=0.025,
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 = linear_quadratic_schedule(num_train_timesteps, linear_quadratic_threshold, linear_steps)
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
else:
timesteps = np.linspace(
@@ -240,7 +238,6 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
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
@@ -275,8 +272,14 @@ class EulerSolver:
return x_prev
def euler_style_multiphase_pred(
self, sample, model_pred, timestep_index, multiphase, is_target=False,
self,
sample,
model_pred,
timestep_index,
multiphase,
is_target=False,
):
inference_indices = np.linspace(
0, len(self.euler_timesteps), num=multiphase, endpoint=False
)
@@ -302,3 +305,4 @@ class EulerSolver:
x_prev = sample + (sigma_prev - sigma) * model_pred
return x_prev, timestep_index_end
+298 -549
View File
File diff suppressed because it is too large Load Diff
@@ -17,8 +17,7 @@ from torch.distributed.fsdp import (
# ShardedStateDictConfig, # un-flattened param but shards, usable by other parallel schemes.
)
from fastvideo.utils.load import get_no_split_modules
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformerBlock
from fastvideo.model.modeling_mochi import MochiTransformerBlock
from functools import partial
@@ -29,18 +28,20 @@ import functools
non_reentrant_wrapper = partial(
checkpoint_wrapper, checkpoint_impl=CheckpointImpl.NO_REENTRANT,
checkpoint_wrapper,
checkpoint_impl=CheckpointImpl.NO_REENTRANT,
)
check_fn = lambda submodule: isinstance(submodule, MochiTransformerBlock)
def apply_fsdp_checkpointing(model, no_split_modules, p=1):
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
@@ -51,51 +52,44 @@ def apply_fsdp_checkpointing(model, no_split_modules, p=1):
nonlocal block_idx
nonlocal cut_off
if isinstance(submodule, no_split_modules):
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,
model, checkpoint_wrapper_fn=non_reentrant_wrapper, check_fn=selective_checkpointing
)
def get_mixed_precision(master_weight_type="fp32"):
weight_type = torch.float32 if master_weight_type == "fp32" else torch.bfloat16
mixed_precision = MixedPrecision(
param_dtype=weight_type,
# Gradient communication precision.
reduce_dtype=weight_type,
# Buffer precision.
buffer_dtype=weight_type,
cast_forward_inputs=False,
)
return mixed_precision
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(
transformer,
sharding_strategy,
use_lora=False,
cpu_offload=False,
master_weight_type="fp32",
):
no_split_modules = get_no_split_modules(transformer)
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=no_split_modules,
transformer_auto_wrap_policy,
transformer_layer_cls={
MochiTransformerBlock,
},
)
# we use float32 for fsdp but autocast during training
mixed_precision = get_mixed_precision(master_weight_type)
mixed_precision = float32
if sharding_strategy == "full":
sharding_strategy = ShardingStrategy.FULL_SHARD
elif sharding_strategy == "hybrid_full":
@@ -104,12 +98,10 @@ def get_dit_fsdp_kwargs(
sharding_strategy = ShardingStrategy.NO_SHARD
auto_wrap_policy = None
elif sharding_strategy == "hybrid_zero2":
sharding_strategy = ShardingStrategy._HYBRID_SHARD_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
)
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,
@@ -118,26 +110,29 @@ def get_dit_fsdp_kwargs(
"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,
}
)
fsdp_kwargs.update({
"use_orig_params": False, # Required for LoRA memory savings
"sync_module_states": True,
})
return fsdp_kwargs
return fsdp_kwargs, no_split_modules
def get_discriminator_fsdp_kwargs():
def get_discriminator_fsdp_kwargs(master_weight_type="fp32"):
auto_wrap_policy = None
# Use existing mixed precision settings
mixed_precision = get_mixed_precision(master_weight_type)
sharding_strategy = ShardingStrategy.NO_SHARD
mixed_precision = float32
sharding_strategy = ShardingStrategy.NO_SHARD
device_id = torch.cuda.current_device()
fsdp_kwargs = {
"auto_wrap_policy": auto_wrap_policy,
@@ -146,5 +141,8 @@ def get_discriminator_fsdp_kwargs(master_weight_type="fp32"):
"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
@@ -19,43 +19,26 @@ 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 import (
USE_PEFT_BACKEND,
is_torch_version,
logging,
scale_lora_layers,
unscale_lora_layers,
)
from diffusers.utils.torch_utils import maybe_allow_in_graph
from diffusers.models.attention import FeedForward as HF_FeedForward
from diffusers.models.attention_processor import Attention
from diffusers.models.embeddings import (
MochiCombinedTimestepCaptionEmbedding,
PatchEmbed,
)
from diffusers.models.embeddings import MochiCombinedTimestepCaptionEmbedding, PatchEmbed
from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.loaders import PeftAdapterMixin
from fastvideo.models.mochi_hf.norm import (
MochiLayerNormContinuous,
MochiRMSNormZero,
MochiModulatedRMSNorm,
MochiRMSNorm,
)
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
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
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
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
class FeedForward(HF_FeedForward):
def __init__(
self,
@@ -68,19 +51,37 @@ class FeedForward(HF_FeedForward):
inner_dim=None,
bias: bool = True,
):
super().__init__(
dim, dim_out, mult, dropout, activation_fn, final_dropout, inner_dim, bias
)
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)
)
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):
class MochiAttention(nn.Module):
def __init__(
self,
query_dim: int,
@@ -114,25 +115,17 @@ class MochiAttention(nn.Module):
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
)
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.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.to_add_out = nn.Linear(self.inner_dim, self.out_context_dim, bias=out_bias)
self.processor = processor
@@ -150,6 +143,7 @@ class MochiAttention(nn.Module):
attention_mask=attention_mask,
**kwargs,
)
class MochiAttnProcessor2_0:
@@ -157,9 +151,7 @@ class MochiAttnProcessor2_0:
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."
)
raise ImportError("MochiAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.")
def __call__(
self,
@@ -180,11 +172,12 @@ class MochiAttnProcessor2_0:
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]
# [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)
@@ -193,37 +186,37 @@ class MochiAttnProcessor2_0:
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
# 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)
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
)
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()
@@ -231,10 +224,9 @@ class MochiAttnProcessor2_0:
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),
@@ -254,16 +246,14 @@ class MochiAttnProcessor2_0:
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 = 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(
@@ -271,9 +261,7 @@ class MochiAttnProcessor2_0:
)
# 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()
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)
@@ -285,6 +273,8 @@ class MochiAttnProcessor2_0:
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)
@@ -296,7 +286,6 @@ class MochiAttnProcessor2_0:
return hidden_states, encoder_hidden_states
@maybe_allow_in_graph
class MochiTransformerBlock(nn.Module):
r"""
@@ -339,9 +328,7 @@ class MochiTransformerBlock(nn.Module):
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
)
self.norm1_context = MochiRMSNormZero(dim, 4 * pooled_projection_dim, eps=eps, elementwise_affine=False)
else:
self.norm1_context = MochiLayerNormContinuous(
embedding_dim=pooled_projection_dim,
@@ -365,18 +352,12 @@ class MochiTransformerBlock(nn.Module):
# 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.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.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 = 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(
@@ -396,19 +377,14 @@ class MochiTransformerBlock(nn.Module):
encoder_attention_mask: torch.Tensor,
temb: torch.Tensor,
image_rotary_emb: Optional[torch.Tensor] = None,
output_attn=False,
output_attn = False,
) -> Tuple[torch.Tensor, torch.Tensor]:
norm_hidden_states, gate_msa, scale_mlp, gate_mlp = self.norm1(
hidden_states, temb
)
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)
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)
@@ -416,27 +392,20 @@ class MochiTransformerBlock(nn.Module):
hidden_states=norm_hidden_states,
encoder_hidden_states=norm_encoder_hidden_states,
image_rotary_emb=image_rotary_emb,
encoder_attention_mask=encoder_attention_mask,
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))
)
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)
)
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)),
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(
@@ -478,22 +447,18 @@ class MochiRoPE(nn.Module):
) -> 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
)
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 = 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
@@ -513,7 +478,7 @@ class MochiRoPE(nn.Module):
@maybe_allow_in_graph
class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
class MochiTransformer3DModel(ModelMixin, ConfigMixin):
r"""
A Transformer model for video-like data introduced in [Mochi](https://huggingface.co/genmo/mochi-1-preview).
@@ -580,9 +545,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
num_attention_heads=8,
)
self.pos_frequencies = nn.Parameter(
torch.full((3, num_attention_heads, attention_head_dim // 2), 0.0)
)
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(
@@ -601,11 +564,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
)
self.norm_out = AdaLayerNormContinuous(
inner_dim,
inner_dim,
elementwise_affine=False,
eps=1e-6,
norm_type="layer_norm",
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)
@@ -621,45 +580,19 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
encoder_hidden_states: torch.Tensor,
timestep: torch.LongTensor,
encoder_attention_mask: torch.Tensor,
output_features=False,
output_features_stride=8,
attention_kwargs: Optional[Dict[str, Any]] = None,
output_attn = False,
return_dict: bool = False,
) -> torch.Tensor:
assert (
return_dict is False
), "return_dict is not supported in MochiTransformer3DModel"
if attention_kwargs is not None:
attention_kwargs = attention_kwargs.copy()
lora_scale = attention_kwargs.pop("scale", 1.0)
else:
lora_scale = 1.0
if USE_PEFT_BACKEND:
# weight the lora layers by setting `lora_scale` for each PEFT layer
scale_lora_layers(self, lora_scale)
else:
if (
attention_kwargs is not None
and attention_kwargs.get("scale", None) is not None
):
logger.warning(
"Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective."
)
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
# 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,
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)
@@ -684,21 +617,15 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
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(
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_features,
output_attn,
**ckpt_kwargs,
)
else:
@@ -708,27 +635,20 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
encoder_attention_mask=encoder_attention_mask,
temb=temb,
image_rotary_emb=image_rotary_emb,
output_attn=output_features,
output_attn = output_attn,
)
if i % output_features_stride == 0:
attn_outputs_list.append(attn_outputs)
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.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 USE_PEFT_BACKEND:
# remove `lora_scale` from each PEFT layer
unscale_lora_layers(self, lora_scale)
if not output_features:
attn_outputs_list = None
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)
# Peiyuan: This is hacked to force mochi to follow the behaviour of SD3 and Flux
return (-output, attn_outputs_list)
@@ -38,7 +38,7 @@ class MochiModulatedRMSNorm(nn.Module):
hidden_states = hidden_states.to(hidden_states_dtype)
return hidden_states
class MochiRMSNorm(nn.Module):
def __init__(self, dim, eps: float, elementwise_affine=True):
@@ -63,11 +63,15 @@ class MochiRMSNorm(nn.Module):
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,
self,
embedding_dim: int,
conditioning_embedding_dim: int,
eps=1e-5,
bias=True,
):
super().__init__()
@@ -77,7 +81,9 @@ class MochiLayerNormContinuous(nn.Module):
self.norm = MochiModulatedRMSNorm(eps=eps)
def forward(
self, x: torch.Tensor, conditioning_embedding: torch.Tensor,
self,
x: torch.Tensor,
conditioning_embedding: torch.Tensor,
) -> torch.Tensor:
input_dtype = x.dtype
@@ -86,7 +92,7 @@ class MochiLayerNormContinuous(nn.Module):
x = self.norm(x, (1 + scale.unsqueeze(1).to(torch.float32)))
return x.to(input_dtype)
class MochiRMSNormZero(nn.Module):
r"""
@@ -96,11 +102,7 @@ class MochiRMSNormZero(nn.Module):
"""
def __init__(
self,
embedding_dim: int,
hidden_dim: int,
eps: float = 1e-5,
elementwise_affine: bool = False,
self, embedding_dim: int, hidden_dim: int, eps: float = 1e-5, elementwise_affine: bool = False
) -> None:
super().__init__()
@@ -116,9 +118,7 @@ class MochiRMSNormZero(nn.Module):
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 = 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
return hidden_states, gate_msa, scale_mlp, gate_mlp
@@ -13,7 +13,7 @@
# limitations under the License.
import inspect
from typing import Callable, Dict, List, Optional, Union, Any
from typing import Callable, Dict, List, Optional, Union
import copy
import numpy as np
import torch
@@ -21,7 +21,7 @@ from transformers import T5EncoderModel, T5TokenizerFast
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
from diffusers.models.autoencoders import AutoencoderKL
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import (
@@ -35,8 +35,8 @@ 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
from diffusers.loaders import Mochi1LoraLoaderMixin
from fastvideo.utils.communications import all_gather
if is_torch_xla_available():
import torch_xla.core.xla_model as xm
@@ -80,19 +80,14 @@ def calculate_shift(
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)
]
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_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)
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]
@@ -132,13 +127,9 @@ def retrieve_timesteps(
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"
)
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()
)
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"
@@ -148,9 +139,7 @@ def retrieve_timesteps(
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()
)
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"
@@ -165,7 +154,7 @@ def retrieve_timesteps(
return timesteps, num_inference_steps
class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
class MochiPipeline(DiffusionPipeline):
r"""
The mochi pipeline for text-to-video generation.
@@ -210,17 +199,14 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
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.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.tokenizer.model_max_length if hasattr(self, "tokenizer") and self.tokenizer is not None else 77
)
self.default_height = 480
self.default_width = 848
@@ -252,32 +238,22 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
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
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]
)
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 = 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_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)
@@ -344,11 +320,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
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
)
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(
@@ -362,10 +334,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
" the batch size of `prompt`."
)
(
negative_prompt_embeds,
negative_prompt_attention_mask,
) = self._get_t5_prompt_embeds(
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,
@@ -373,12 +342,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
dtype=dtype,
)
return (
prompt_embeds,
prompt_attention_mask,
negative_prompt_embeds,
negative_prompt_attention_mask,
)
return prompt_embeds, prompt_attention_mask, negative_prompt_embeds, negative_prompt_attention_mask
def check_inputs(
self,
@@ -392,13 +356,10 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
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}."
)
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
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]}"
@@ -413,25 +374,14 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
raise ValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
)
elif prompt is not None and (
not isinstance(prompt, str) and not isinstance(prompt, list)
):
raise ValueError(
f"`prompt` has to be of type `str` or `list` but is {type(prompt)}"
)
elif prompt 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`."
)
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 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:
@@ -502,8 +452,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
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=torch.float32)
latents = latents.to(dtype)
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
return latents
@property
@@ -518,10 +467,6 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
def num_timesteps(self):
return self._num_timesteps
@property
def attention_kwargs(self):
return self._attention_kwargs
@property
def interrupt(self):
return self._interrupt
@@ -534,8 +479,8 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
negative_prompt: Optional[Union[str, List[str]]] = None,
height: Optional[int] = None,
width: Optional[int] = None,
num_frames: int = 19,
num_inference_steps: int = 64,
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,
@@ -547,11 +492,10 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
negative_prompt_attention_mask: Optional[torch.Tensor] = None,
output_type: Optional[str] = "pil",
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
max_sequence_length: int = 256,
return_all_states=False,
return_all_states = False,
):
r"""
Function invoked when calling the pipeline for generation.
@@ -603,10 +547,6 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
[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.
attention_kwargs (`dict`, *optional*):
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
`self.processor` in
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
callback_on_step_end (`Callable`, *optional*):
A function that calls at the end of each denoising steps during the inference. The function is called
with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int,
@@ -646,7 +586,6 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
)
self._guidance_scale = guidance_scale
self._attention_kwargs = attention_kwargs
self._interrupt = False
# 2. Define call parameters
@@ -679,9 +618,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
)
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
)
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
@@ -698,10 +635,9 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
)
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 = 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
@@ -710,56 +646,49 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
sigmas = linear_quadratic_schedule(num_inference_steps, threshold_noise)
sigmas = np.array(sigmas)
# check if of type FlowMatchEulerDiscreteScheduler
if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler):
if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler):
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler, num_inference_steps, device, timesteps, sigmas,
self.scheduler,
num_inference_steps,
device,
timesteps,
sigmas,
)
else:
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler, num_inference_steps, device,
self.scheduler,
num_inference_steps,
device,
)
num_warmup_steps = max(
len(timesteps) - num_inference_steps * self.scheduler.order, 0
)
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
self._num_timesteps = len(timesteps)
# 6. Denoising loop
self._progress_bar_config = {"disable": nccl_info.rank_within_group != 0}
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
)
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,
attention_kwargs=attention_kwargs,
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
)
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 = self.scheduler.step(noise_pred, t, latents.to(torch.float32), return_dict=False)[0]
latents = latents.to(latents_dtype)
if latents.dtype != latents_dtype:
@@ -777,9 +706,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
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
):
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
progress_bar.update()
if XLA_AVAILABLE:
@@ -787,49 +714,34 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
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)
#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
)
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)
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
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
)
video = self.video_processor.postprocess_video(video, output_type=output_type)
# Offload all models
self.maybe_free_model_hooks()
if return_all_states:
+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)
-34
View File
@@ -1,34 +0,0 @@
from flash_attn import flash_attn_varlen_qkvpacked_func
from flash_attn.bert_padding import pad_input, unpad_input
from einops import rearrange
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
-87
View File
@@ -1,87 +0,0 @@
import os
import torch
__all__ = [
"C_SCALE",
"PROMPT_TEMPLATE",
"MODEL_BASE",
"PRECISIONS",
"NORMALIZATION_TYPE",
"ACTIVATION_TYPE",
"VAE_PATH",
"TEXT_ENCODER_PATH",
"TOKENIZER_PATH",
"TEXT_PROJECTION",
"DATA_TYPE",
"NEGATIVE_PROMPT",
]
PRECISION_TO_TYPE = {
"fp32": torch.float32,
"fp16": torch.float16,
"bf16": torch.bfloat16,
}
# =================== Constant Values =====================
# Computation scale factor, 1P = 1_000_000_000_000_000. Tensorboard will display the value in PetaFLOPS to avoid
# overflow error when tensorboard logging values.
C_SCALE = 1_000_000_000_000_000
# When using decoder-only models, we must provide a prompt template to instruct the text encoder
# on how to generate the text.
# --------------------------------------------------------------------
PROMPT_TEMPLATE_ENCODE = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the image by detailing the color, shape, size, texture, "
"quantity, text, spatial relationships of the objects and background:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
)
PROMPT_TEMPLATE_ENCODE_VIDEO = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
"1. The main content and theme of the video."
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
"4. background environment, light, style and atmosphere."
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
)
NEGATIVE_PROMPT = "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion"
PROMPT_TEMPLATE = {
"dit-llm-encode": {"template": PROMPT_TEMPLATE_ENCODE, "crop_start": 36,},
"dit-llm-encode-video": {
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
"crop_start": 95,
},
}
# ======================= Model ======================
PRECISIONS = {"fp32", "fp16", "bf16"}
NORMALIZATION_TYPE = {"layer", "rms"}
ACTIVATION_TYPE = {"relu", "silu", "gelu", "gelu_tanh"}
# =================== Model Path =====================
MODEL_BASE = os.getenv("MODEL_BASE", "./data/hunyuan")
# =================== Data =======================
DATA_TYPE = {"image", "video", "image_video"}
# 3D VAE
VAE_PATH = {"884-16c-hy": f"{MODEL_BASE}/hunyuan-video-t2v-720p/vae"}
# Text Encoder
TEXT_ENCODER_PATH = {
"clipL": f"{MODEL_BASE}/text_encoder_2",
"llm": f"{MODEL_BASE}/text_encoder",
}
# Tokenizer
TOKENIZER_PATH = {
"clipL": f"{MODEL_BASE}/text_encoder_2",
"llm": f"{MODEL_BASE}/text_encoder",
}
TEXT_PROJECTION = {
"linear", # Default, an nn.Linear() layer
"single_refiner", # Single TokenRefiner. Refer to LI-DiT
}
@@ -1,2 +0,0 @@
from .pipelines import HunyuanVideoPipeline
from .schedulers import FlowMatchDiscreteScheduler
@@ -1 +0,0 @@
from .pipeline_hunyuan_video import HunyuanVideoPipeline
File diff suppressed because it is too large Load Diff
@@ -1 +0,0 @@
from .scheduling_flow_match_discrete import FlowMatchDiscreteScheduler
@@ -1,257 +0,0 @@
# Copyright 2024 Stability AI, Katherine Crowson and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
#
# Modified from diffusers==0.29.2
#
# ==============================================================================
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.schedulers.scheduling_utils import SchedulerMixin
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@dataclass
class FlowMatchDiscreteSchedulerOutput(BaseOutput):
"""
Output class for the scheduler's `step` function output.
Args:
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
denoising loop.
"""
prev_sample: torch.FloatTensor
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
"""
Euler scheduler.
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
methods the library implements for all schedulers such as loading and saving.
Args:
num_train_timesteps (`int`, defaults to 1000):
The number of diffusion steps to train the model.
timestep_spacing (`str`, defaults to `"linspace"`):
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
shift (`float`, defaults to 1.0):
The shift value for the timestep schedule.
reverse (`bool`, defaults to `True`):
Whether to reverse the timestep schedule.
"""
_compatibles = []
order = 1
@register_to_config
def __init__(
self,
num_train_timesteps: int = 1000,
shift: float = 1.0,
reverse: bool = True,
solver: str = "euler",
n_tokens: Optional[int] = None,
):
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
if not reverse:
sigmas = sigmas.flip(0)
self.sigmas = sigmas
# the value fed to model
self.timesteps = (sigmas[:-1] * num_train_timesteps).to(dtype=torch.float32)
self._step_index = None
self._begin_index = None
self.supported_solver = ["euler"]
if solver not in self.supported_solver:
raise ValueError(
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
)
@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 _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,
n_tokens: int = 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.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
"""
self.num_inference_steps = num_inference_steps
sigmas = torch.linspace(1, 0, num_inference_steps + 1)
sigmas = self.sd3_time_shift(sigmas)
if not self.config.reverse:
sigmas = 1 - sigmas
self.sigmas = sigmas
self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(
dtype=torch.float32, device=device
)
# Reset step index
self._step_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 scale_model_input(
self, sample: torch.Tensor, timestep: Optional[int] = None
) -> torch.Tensor:
return sample
def sd3_time_shift(self, t: torch.Tensor):
return (self.config.shift * t) / (1 + (self.config.shift - 1) * t)
def step(
self,
model_output: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
sample: torch.FloatTensor,
return_dict: bool = True,
) -> Union[FlowMatchDiscreteSchedulerOutput, 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.
generator (`torch.Generator`, *optional*):
A random number generator.
n_tokens (`int`, *optional*):
Number of tokens in the input sequence.
return_dict (`bool`):
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
tuple.
Returns:
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
returned, otherwise a tuple is returned where the first element is the sample tensor.
"""
if (
isinstance(timestep, int)
or isinstance(timestep, torch.IntTensor)
or isinstance(timestep, torch.LongTensor)
):
raise ValueError(
(
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."
),
)
if self.step_index is None:
self._init_step_index(timestep)
# Upcast to avoid precision issues when computing prev_sample
sample = sample.to(torch.float32)
dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
if self.config.solver == "euler":
prev_sample = sample + model_output.to(torch.float32) * dt
else:
raise ValueError(
f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}"
)
# upon completion increase step index by one
self._step_index += 1
if not return_dict:
return (prev_sample,)
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
def __len__(self):
return self.config.num_train_timesteps
-383
View File
@@ -1,383 +0,0 @@
import argparse
from .constants import *
import re
from .modules.models import HUNYUAN_VIDEO_CONFIG
def parse_args(namespace=None):
parser = argparse.ArgumentParser(description="HunyuanVideo inference script")
parser = add_network_args(parser)
parser = add_extra_models_args(parser)
parser = add_denoise_schedule_args(parser)
parser = add_inference_args(parser)
parser = add_parallel_args(parser)
args = parser.parse_args(namespace=namespace)
args = sanity_check_args(args)
return args
def add_network_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="HunyuanVideo network args")
# Main model
group.add_argument(
"--model",
type=str,
choices=list(HUNYUAN_VIDEO_CONFIG.keys()),
default="HYVideo-T/2-cfgdistill",
)
group.add_argument(
"--latent-channels",
type=str,
default=16,
help="Number of latent channels of DiT. If None, it will be determined by `vae`. If provided, "
"it still needs to match the latent channels of the VAE model.",
)
group.add_argument(
"--precision",
type=str,
default="bf16",
choices=PRECISIONS,
help="Precision mode. Options: fp32, fp16, bf16. Applied to the backbone model and optimizer.",
)
# RoPE
group.add_argument(
"--rope-theta", type=int, default=256, help="Theta used in RoPE."
)
return parser
def add_extra_models_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(
title="Extra models args, including vae, text encoders and tokenizers)"
)
# - VAE
group.add_argument(
"--vae",
type=str,
default="884-16c-hy",
choices=list(VAE_PATH),
help="Name of the VAE model.",
)
group.add_argument(
"--vae-precision",
type=str,
default="fp16",
choices=PRECISIONS,
help="Precision mode for the VAE model.",
)
group.add_argument(
"--vae-tiling",
action="store_true",
help="Enable tiling for the VAE model to save GPU memory.",
)
group.set_defaults(vae_tiling=True)
group.add_argument(
"--text-encoder",
type=str,
default="llm",
choices=list(TEXT_ENCODER_PATH),
help="Name of the text encoder model.",
)
group.add_argument(
"--text-encoder-precision",
type=str,
default="fp16",
choices=PRECISIONS,
help="Precision mode for the text encoder model.",
)
group.add_argument(
"--text-states-dim",
type=int,
default=4096,
help="Dimension of the text encoder hidden states.",
)
group.add_argument(
"--text-len", type=int, default=256, help="Maximum length of the text input."
)
group.add_argument(
"--tokenizer",
type=str,
default="llm",
choices=list(TOKENIZER_PATH),
help="Name of the tokenizer model.",
)
group.add_argument(
"--prompt-template",
type=str,
default="dit-llm-encode",
choices=PROMPT_TEMPLATE,
help="Image prompt template for the decoder-only text encoder model.",
)
group.add_argument(
"--prompt-template-video",
type=str,
default="dit-llm-encode-video",
choices=PROMPT_TEMPLATE,
help="Video prompt template for the decoder-only text encoder model.",
)
group.add_argument(
"--hidden-state-skip-layer",
type=int,
default=2,
help="Skip layer for hidden states.",
)
group.add_argument(
"--apply-final-norm",
action="store_true",
help="Apply final normalization to the used text encoder hidden states.",
)
# - CLIP
group.add_argument(
"--text-encoder-2",
type=str,
default="clipL",
choices=list(TEXT_ENCODER_PATH),
help="Name of the second text encoder model.",
)
group.add_argument(
"--text-encoder-precision-2",
type=str,
default="fp16",
choices=PRECISIONS,
help="Precision mode for the second text encoder model.",
)
group.add_argument(
"--text-states-dim-2",
type=int,
default=768,
help="Dimension of the second text encoder hidden states.",
)
group.add_argument(
"--tokenizer-2",
type=str,
default="clipL",
choices=list(TOKENIZER_PATH),
help="Name of the second tokenizer model.",
)
group.add_argument(
"--text-len-2",
type=int,
default=77,
help="Maximum length of the second text input.",
)
return parser
def add_denoise_schedule_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Denoise schedule args")
group.add_argument(
"--denoise-type",
type=str,
default="flow",
help="Denoise type for noised inputs.",
)
# Flow Matching
group.add_argument(
"--flow-shift",
type=float,
default=7.0,
help="Shift factor for flow matching schedulers.",
)
group.add_argument(
"--flow-reverse",
action="store_true",
help="If reverse, learning/sampling from t=1 -> t=0.",
)
group.add_argument(
"--flow-solver", type=str, default="euler", help="Solver for flow matching.",
)
group.add_argument(
"--use-linear-quadratic-schedule",
action="store_true",
help="Use linear quadratic schedule for flow matching."
"Following MovieGen (https://ai.meta.com/static-resource/movie-gen-research-paper)",
)
group.add_argument(
"--linear-schedule-end",
type=int,
default=25,
help="End step for linear quadratic schedule for flow matching.",
)
return parser
def add_inference_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Inference args")
# ======================== Model loads ========================
group.add_argument(
"--model-base",
type=str,
default="ckpts",
help="Root path of all the models, including t2v models and extra models.",
)
group.add_argument(
"--dit-weight",
type=str,
default="ckpts/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
help="Path to the HunyuanVideo model. If None, search the model in the args.model_root."
"1. If it is a file, load the model directly."
"2. If it is a directory, search the model in the directory. Support two types of models: "
"1) named `pytorch_model_*.pt`"
"2) named `*_model_states.pt`, where * can be `mp_rank_00`.",
)
group.add_argument(
"--model-resolution",
type=str,
default="540p",
choices=["540p", "720p"],
help="Root path of all the models, including t2v models and extra models.",
)
group.add_argument(
"--load-key",
type=str,
default="module",
help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.",
)
group.add_argument(
"--use-cpu-offload",
action="store_true",
help="Use CPU offload for the model load.",
)
# ======================== Inference general setting ========================
group.add_argument(
"--batch-size",
type=int,
default=1,
help="Batch size for inference and evaluation.",
)
group.add_argument(
"--infer-steps",
type=int,
default=50,
help="Number of denoising steps for inference.",
)
group.add_argument(
"--disable-autocast",
action="store_true",
help="Disable autocast for denoising loop and vae decoding in pipeline sampling.",
)
group.add_argument(
"--save-path",
type=str,
default="./results",
help="Path to save the generated samples.",
)
group.add_argument(
"--save-path-suffix",
type=str,
default="",
help="Suffix for the directory of saved samples.",
)
group.add_argument(
"--name-suffix",
type=str,
default="",
help="Suffix for the names of saved samples.",
)
group.add_argument(
"--num-videos",
type=int,
default=1,
help="Number of videos to generate for each prompt.",
)
# ---sample size---
group.add_argument(
"--video-size",
type=int,
nargs="+",
default=(720, 1280),
help="Video size for training. If a single value is provided, it will be used for both height "
"and width. If two values are provided, they will be used for height and width "
"respectively.",
)
group.add_argument(
"--video-length",
type=int,
default=129,
help="How many frames to sample from a video. if using 3d vae, the number should be 4n+1",
)
# --- prompt ---
group.add_argument(
"--prompt",
type=str,
default=None,
help="Prompt for sampling during evaluation.",
)
group.add_argument(
"--seed-type",
type=str,
default="auto",
choices=["file", "random", "fixed", "auto"],
help="Seed type for evaluation. If file, use the seed from the CSV file. If random, generate a "
"random seed. If fixed, use the fixed seed given by `--seed`. If auto, `csv` will use the "
"seed column if available, otherwise use the fixed `seed` value. `prompt` will use the "
"fixed `seed` value.",
)
group.add_argument("--seed", type=int, default=None, help="Seed for evaluation.")
# Classifier-Free Guidance
group.add_argument(
"--neg-prompt", type=str, default=None, help="Negative prompt for sampling."
)
group.add_argument(
"--cfg-scale", type=float, default=1.0, help="Classifier free guidance scale."
)
group.add_argument(
"--embedded-cfg-scale",
type=float,
default=6.0,
help="Embeded classifier free guidance scale.",
)
group.add_argument(
"--reproduce",
action="store_true",
help="Enable reproducibility by setting random seeds and deterministic algorithms.",
)
return parser
def add_parallel_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(title="Parallel args")
# ======================== Model loads ========================
group.add_argument(
"--ulysses-degree", type=int, default=1, help="Ulysses degree.",
)
group.add_argument(
"--ring-degree", type=int, default=1, help="Ulysses degree.",
)
return parser
def sanity_check_args(args):
# VAE channels
vae_pattern = r"\d{2,3}-\d{1,2}c-\w+"
if not re.match(vae_pattern, args.vae):
raise ValueError(
f"Invalid VAE model: {args.vae}. Must be in the format of '{vae_pattern}'."
)
vae_channels = int(args.vae.split("-")[1][:-1])
if args.latent_channels is None:
args.latent_channels = vae_channels
if vae_channels != args.latent_channels:
raise ValueError(
f"Latent channels ({args.latent_channels}) must match the VAE channels ({vae_channels})."
)
return args
-532
View File
@@ -1,532 +0,0 @@
import os
import time
import random
import functools
from typing import List, Optional, Tuple, Union
from pathlib import Path
from loguru import logger
import torch
import torch.distributed as dist
from fastvideo.models.hunyuan.constants import (
PROMPT_TEMPLATE,
NEGATIVE_PROMPT,
PRECISION_TO_TYPE,
)
from fastvideo.models.hunyuan.vae import load_vae
from fastvideo.models.hunyuan.modules import load_model
from fastvideo.models.hunyuan.text_encoder import TextEncoder
from fastvideo.models.hunyuan.utils.data_utils import align_to
from fastvideo.models.hunyuan.diffusion.schedulers import FlowMatchDiscreteScheduler
from fastvideo.models.hunyuan.diffusion.pipelines import HunyuanVideoPipeline
from safetensors.torch import load_file as safetensors_load_file
from fastvideo.utils.parallel_states import (
initialize_sequence_parallel_state,
nccl_info,
)
class Inference(object):
def __init__(
self,
args,
vae,
vae_kwargs,
text_encoder,
model,
text_encoder_2=None,
pipeline=None,
use_cpu_offload=False,
device=None,
logger=None,
parallel_args=None,
):
self.vae = vae
self.vae_kwargs = vae_kwargs
self.text_encoder = text_encoder
self.text_encoder_2 = text_encoder_2
self.model = model
self.pipeline = pipeline
self.use_cpu_offload = use_cpu_offload
self.args = args
self.device = (
device
if device is not None
else "cuda"
if torch.cuda.is_available()
else "cpu"
)
self.logger = logger
self.parallel_args = parallel_args
@classmethod
def from_pretrained(cls, pretrained_model_path, args, device=None, **kwargs):
"""
Initialize the Inference pipeline.
Args:
pretrained_model_path (str or pathlib.Path): The model path, including t2v, text encoder and vae checkpoints.
args (argparse.Namespace): The arguments for the pipeline.
device (int): The device for inference. Default is 0.
"""
# ========================================================================
logger.info(f"Got text-to-video model root path: {pretrained_model_path}")
# ==================== Initialize Distributed Environment ================
if nccl_info.sp_size > 1:
device = torch.device(f"cuda:{os.environ['LOCAL_RANK']}")
if device is None:
device = "cuda" if torch.cuda.is_available() else "cpu"
parallel_args = None # {"ulysses_degree": args.ulysses_degree, "ring_degree": args.ring_degree}
# ======================== Get the args path =============================
# Disable gradient
torch.set_grad_enabled(False)
# =========================== Build main model ===========================
logger.info("Building model...")
factor_kwargs = {"device": device, "dtype": PRECISION_TO_TYPE[args.precision]}
in_channels = args.latent_channels
out_channels = args.latent_channels
model = load_model(
args,
in_channels=in_channels,
out_channels=out_channels,
factor_kwargs=factor_kwargs,
)
model = model.to(device)
model = Inference.load_state_dict(args, model, pretrained_model_path)
model.eval()
# ============================= Build extra models ========================
# VAE
vae, _, s_ratio, t_ratio = load_vae(
args.vae,
args.vae_precision,
logger=logger,
device=device if not args.use_cpu_offload else "cpu",
)
vae_kwargs = {"s_ratio": s_ratio, "t_ratio": t_ratio}
# Text encoder
if args.prompt_template_video is not None:
crop_start = PROMPT_TEMPLATE[args.prompt_template_video].get(
"crop_start", 0
)
elif args.prompt_template is not None:
crop_start = PROMPT_TEMPLATE[args.prompt_template].get("crop_start", 0)
else:
crop_start = 0
max_length = args.text_len + crop_start
# prompt_template
prompt_template = (
PROMPT_TEMPLATE[args.prompt_template]
if args.prompt_template is not None
else None
)
# prompt_template_video
prompt_template_video = (
PROMPT_TEMPLATE[args.prompt_template_video]
if args.prompt_template_video is not None
else None
)
text_encoder = TextEncoder(
text_encoder_type=args.text_encoder,
max_length=max_length,
text_encoder_precision=args.text_encoder_precision,
tokenizer_type=args.tokenizer,
prompt_template=prompt_template,
prompt_template_video=prompt_template_video,
hidden_state_skip_layer=args.hidden_state_skip_layer,
apply_final_norm=args.apply_final_norm,
reproduce=args.reproduce,
logger=logger,
device=device if not args.use_cpu_offload else "cpu",
)
text_encoder_2 = None
if args.text_encoder_2 is not None:
text_encoder_2 = TextEncoder(
text_encoder_type=args.text_encoder_2,
max_length=args.text_len_2,
text_encoder_precision=args.text_encoder_precision_2,
tokenizer_type=args.tokenizer_2,
reproduce=args.reproduce,
logger=logger,
device=device if not args.use_cpu_offload else "cpu",
)
return cls(
args=args,
vae=vae,
vae_kwargs=vae_kwargs,
text_encoder=text_encoder,
text_encoder_2=text_encoder_2,
model=model,
use_cpu_offload=args.use_cpu_offload,
device=device,
logger=logger,
parallel_args=parallel_args,
)
@staticmethod
def load_state_dict(args, model, pretrained_model_path):
load_key = args.load_key
dit_weight = Path(args.dit_weight)
if dit_weight is None:
model_dir = pretrained_model_path / f"t2v_{args.model_resolution}"
files = list(model_dir.glob("*.pt"))
if len(files) == 0:
raise ValueError(f"No model weights found in {model_dir}")
if str(files[0]).startswith("pytorch_model_"):
model_path = dit_weight / f"pytorch_model_{load_key}.pt"
bare_model = True
elif any(str(f).endswith("_model_states.pt") for f in files):
files = [f for f in files if str(f).endswith("_model_states.pt")]
model_path = files[0]
if len(files) > 1:
logger.warning(
f"Multiple model weights found in {dit_weight}, using {model_path}"
)
bare_model = False
else:
raise ValueError(
f"Invalid model path: {dit_weight} with unrecognized weight format: "
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
f"specific weight file, please provide the full path to the file."
)
else:
if dit_weight.is_dir():
files = list(dit_weight.glob("*.pt"))
if len(files) == 0:
raise ValueError(f"No model weights found in {dit_weight}")
if str(files[0]).startswith("pytorch_model_"):
model_path = dit_weight / f"pytorch_model_{load_key}.pt"
bare_model = True
elif any(str(f).endswith("_model_states.pt") for f in files):
files = [f for f in files if str(f).endswith("_model_states.pt")]
model_path = files[0]
if len(files) > 1:
logger.warning(
f"Multiple model weights found in {dit_weight}, using {model_path}"
)
bare_model = False
else:
raise ValueError(
f"Invalid model path: {dit_weight} with unrecognized weight format: "
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
f"specific weight file, please provide the full path to the file."
)
elif dit_weight.is_file():
model_path = dit_weight
bare_model = "unknown"
else:
raise ValueError(f"Invalid model path: {dit_weight}")
if not model_path.exists():
raise ValueError(f"model_path not exists: {model_path}")
logger.info(f"Loading torch model {model_path}...")
if model_path.suffix == ".safetensors":
# Use safetensors library for .safetensors files
state_dict = safetensors_load_file(model_path)
elif model_path.suffix == ".pt":
# Use torch for .pt files
state_dict = torch.load(
model_path, map_location=lambda storage, loc: storage
)
else:
raise ValueError(f"Unsupported file format: {model_path}")
if bare_model == "unknown" and ("ema" in state_dict or "module" in state_dict):
bare_model = False
if bare_model is False:
if load_key in state_dict:
state_dict = state_dict[load_key]
else:
raise KeyError(
f"Missing key: `{load_key}` in the checkpoint: {model_path}. The keys in the checkpoint "
f"are: {list(state_dict.keys())}."
)
model.load_state_dict(state_dict, strict=True)
return model
@staticmethod
def parse_size(size):
if isinstance(size, int):
size = [size]
if not isinstance(size, (list, tuple)):
raise ValueError(f"Size must be an integer or (height, width), got {size}.")
if len(size) == 1:
size = [size[0], size[0]]
if len(size) != 2:
raise ValueError(f"Size must be an integer or (height, width), got {size}.")
return size
class HunyuanVideoSampler(Inference):
def __init__(
self,
args,
vae,
vae_kwargs,
text_encoder,
model,
text_encoder_2=None,
pipeline=None,
use_cpu_offload=False,
device=0,
logger=None,
parallel_args=None,
):
super().__init__(
args,
vae,
vae_kwargs,
text_encoder,
model,
text_encoder_2=text_encoder_2,
pipeline=pipeline,
use_cpu_offload=use_cpu_offload,
device=device,
logger=logger,
parallel_args=parallel_args,
)
self.pipeline = self.load_diffusion_pipeline(
args=args,
vae=self.vae,
text_encoder=self.text_encoder,
text_encoder_2=self.text_encoder_2,
model=self.model,
device=self.device,
)
self.default_negative_prompt = NEGATIVE_PROMPT
def load_diffusion_pipeline(
self,
args,
vae,
text_encoder,
text_encoder_2,
model,
scheduler=None,
device=None,
progress_bar_config=None,
data_type="video",
):
"""Load the denoising scheduler for inference."""
if scheduler is None:
if args.denoise_type == "flow":
scheduler = FlowMatchDiscreteScheduler(
shift=args.flow_shift,
reverse=args.flow_reverse,
solver=args.flow_solver,
)
else:
raise ValueError(f"Invalid denoise type {args.denoise_type}")
pipeline = HunyuanVideoPipeline(
vae=vae,
text_encoder=text_encoder,
text_encoder_2=text_encoder_2,
transformer=model,
scheduler=scheduler,
progress_bar_config=progress_bar_config,
args=args,
)
if self.use_cpu_offload:
pipeline.enable_sequential_cpu_offload()
else:
pipeline = pipeline.to(device)
return pipeline
@torch.no_grad()
def predict(
self,
prompt,
height=192,
width=336,
video_length=129,
seed=None,
negative_prompt=None,
infer_steps=50,
guidance_scale=6,
flow_shift=5.0,
embedded_guidance_scale=None,
batch_size=1,
num_videos_per_prompt=1,
**kwargs,
):
"""
Predict the image/video from the given text.
Args:
prompt (str or List[str]): The input text.
kwargs:
height (int): The height of the output video. Default is 192.
width (int): The width of the output video. Default is 336.
video_length (int): The frame number of the output video. Default is 129.
seed (int or List[str]): The random seed for the generation. Default is a random integer.
negative_prompt (str or List[str]): The negative text prompt. Default is an empty string.
guidance_scale (float): The guidance scale for the generation. Default is 6.0.
num_images_per_prompt (int): The number of images per prompt. Default is 1.
infer_steps (int): The number of inference steps. Default is 100.
"""
out_dict = dict()
# ========================================================================
# Arguments: seed
# ========================================================================
if isinstance(seed, torch.Tensor):
seed = seed.tolist()
if seed is None:
seeds = [
random.randint(0, 1_000_000)
for _ in range(batch_size * num_videos_per_prompt)
]
elif isinstance(seed, int):
seeds = [
seed + i
for _ in range(batch_size)
for i in range(num_videos_per_prompt)
]
elif isinstance(seed, (list, tuple)):
if len(seed) == batch_size:
seeds = [
int(seed[i]) + j
for i in range(batch_size)
for j in range(num_videos_per_prompt)
]
elif len(seed) == batch_size * num_videos_per_prompt:
seeds = [int(s) for s in seed]
else:
raise ValueError(
f"Length of seed must be equal to number of prompt(batch_size) or "
f"batch_size * num_videos_per_prompt ({batch_size} * {num_videos_per_prompt}), got {seed}."
)
else:
raise ValueError(
f"Seed must be an integer, a list of integers, or None, got {seed}."
)
# Peiyuan: using GPU seed will cause A100 and H100 to generate different results...
generator = [torch.Generator("cpu").manual_seed(seed) for seed in seeds]
out_dict["seeds"] = seeds
# ========================================================================
# Arguments: target_width, target_height, target_video_length
# ========================================================================
if width <= 0 or height <= 0 or video_length <= 0:
raise ValueError(
f"`height` and `width` and `video_length` must be positive integers, got height={height}, width={width}, video_length={video_length}"
)
if (video_length - 1) % 4 != 0:
raise ValueError(
f"`video_length-1` must be a multiple of 4, got {video_length}"
)
logger.info(
f"Input (height, width, video_length) = ({height}, {width}, {video_length})"
)
target_height = align_to(height, 16)
target_width = align_to(width, 16)
target_video_length = video_length
out_dict["size"] = (target_height, target_width, target_video_length)
# ========================================================================
# Arguments: prompt, new_prompt, negative_prompt
# ========================================================================
if not isinstance(prompt, str):
raise TypeError(f"`prompt` must be a string, but got {type(prompt)}")
prompt = [prompt.strip()]
# negative prompt
if negative_prompt is None or negative_prompt == "":
negative_prompt = self.default_negative_prompt
if not isinstance(negative_prompt, str):
raise TypeError(
f"`negative_prompt` must be a string, but got {type(negative_prompt)}"
)
negative_prompt = [negative_prompt.strip()]
# ========================================================================
# Scheduler
# ========================================================================
scheduler = FlowMatchDiscreteScheduler(
shift=flow_shift,
reverse=self.args.flow_reverse,
solver=self.args.flow_solver,
)
self.pipeline.scheduler = scheduler
if "884" in self.args.vae:
latents_size = [(video_length - 1) // 4 + 1, height // 8, width // 8]
elif "888" in self.args.vae:
latents_size = [(video_length - 1) // 8 + 1, height // 8, width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# ========================================================================
# Print infer args
# ========================================================================
debug_str = f"""
height: {target_height}
width: {target_width}
video_length: {target_video_length}
prompt: {prompt}
neg_prompt: {negative_prompt}
seed: {seed}
infer_steps: {infer_steps}
num_videos_per_prompt: {num_videos_per_prompt}
guidance_scale: {guidance_scale}
n_tokens: {n_tokens}
flow_shift: {flow_shift}
embedded_guidance_scale: {embedded_guidance_scale}"""
logger.debug(debug_str)
# ========================================================================
# Pipeline inference
# ========================================================================
start_time = time.time()
samples = self.pipeline(
prompt=prompt,
height=target_height,
width=target_width,
video_length=target_video_length,
num_inference_steps=infer_steps,
guidance_scale=guidance_scale,
negative_prompt=negative_prompt,
num_videos_per_prompt=num_videos_per_prompt,
generator=generator,
output_type="pil",
n_tokens=n_tokens,
embedded_guidance_scale=embedded_guidance_scale,
data_type="video" if target_video_length > 1 else "image",
is_progress_bar=True,
vae_ver=self.args.vae,
enable_tiling=self.args.vae_tiling,
)[0]
out_dict["samples"] = samples
out_dict["prompts"] = prompt
gen_time = time.time() - start_time
logger.info(f"Success, time: {gen_time}")
return out_dict
@@ -1,25 +0,0 @@
from .models import HYVideoDiffusionTransformer, HUNYUAN_VIDEO_CONFIG
def load_model(args, in_channels, out_channels, factor_kwargs):
"""load hunyuan video model
Args:
args (dict): model args
in_channels (int): input channels number
out_channels (int): output channels number
factor_kwargs (dict): factor kwargs
Returns:
model (nn.Module): The hunyuan video model
"""
if args.model in HUNYUAN_VIDEO_CONFIG.keys():
model = HYVideoDiffusionTransformer(
in_channels=in_channels,
out_channels=out_channels,
**HUNYUAN_VIDEO_CONFIG[args.model],
**factor_kwargs,
)
return model
else:
raise NotImplementedError()
@@ -1,23 +0,0 @@
import torch.nn as nn
def get_activation_layer(act_type):
"""get activation layer
Args:
act_type (str): the activation type
Returns:
torch.nn.functional: the activation layer
"""
if act_type == "gelu":
return lambda: nn.GELU()
elif act_type == "gelu_tanh":
# Approximate `tanh` requires torch >= 1.13
return lambda: nn.GELU(approximate="tanh")
elif act_type == "relu":
return nn.ReLU
elif act_type == "silu":
return nn.SiLU
else:
raise ValueError(f"Unknown activation type: {act_type}")
@@ -1,84 +0,0 @@
import importlib.metadata
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
from fastvideo.utils.communications import all_gather, all_to_all_4D
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
def attention(
q, k, v, drop_rate=0, attn_mask=None, causal=False,
):
qkv = torch.stack([q, k, v], dim=2)
if attn_mask is not None and attn_mask.dtype != torch.bool:
attn_mask = attn_mask.bool()
x = flash_attn_no_pad(
qkv, attn_mask, causal=causal, dropout_p=drop_rate, softmax_scale=None
)
b, s, a, d = x.shape
out = x.reshape(b, s, -1)
return out
def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask):
# 1GPU torch.Size([1, 11264, 24, 128]) tensor([ 0, 11275, 11520], device='cuda:0', dtype=torch.int32)
# 2GPU torch.Size([1, 5632, 24, 128]) tensor([ 0, 5643, 5888], device='cuda:0', dtype=torch.int32)
query, encoder_query = q
key, encoder_key = k
value, encoder_value = v
if get_sequence_parallel_state():
# 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)
# [b, s, h, d]
sequence_length = query.size(1)
encoder_sequence_length = encoder_query.size(1)
# Hint: please check encoder_query.shape
query = torch.cat([query, encoder_query], dim=1)
key = torch.cat([key, encoder_key], dim=1)
value = torch.cat([value, encoder_value], dim=1)
# B, S, 3, H, D
qkv = torch.stack([query, key, value], dim=2)
attn_mask = F.pad(text_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, encoder_hidden_states = hidden_states.split_with_sizes(
(sequence_length, encoder_sequence_length), dim=1
)
if get_sequence_parallel_state():
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.to(query.dtype)
encoder_hidden_states = encoder_hidden_states.to(query.dtype)
attn = torch.cat([hidden_states, encoder_hidden_states], dim=1)
b, s, a, d = attn.shape
attn = attn.reshape(b, s, -1)
return attn
@@ -1,157 +0,0 @@
import math
import torch
import torch.nn as nn
from einops import rearrange, repeat
from ..utils.helpers import to_2tuple
class PatchEmbed(nn.Module):
"""2D Image to Patch Embedding
Image to Patch Embedding using Conv2d
A convolution based approach to patchifying a 2D image w/ embedding projection.
Based on the impl in https://github.com/google-research/vision_transformer
Hacked together by / Copyright 2020 Ross Wightman
Remove the _assert function in forward function to be compatible with multi-resolution images.
"""
def __init__(
self,
patch_size=16,
in_chans=3,
embed_dim=768,
norm_layer=None,
flatten=True,
bias=True,
dtype=None,
device=None,
):
factory_kwargs = {"dtype": dtype, "device": device}
super().__init__()
patch_size = to_2tuple(patch_size)
self.patch_size = patch_size
self.flatten = flatten
self.proj = nn.Conv3d(
in_chans,
embed_dim,
kernel_size=patch_size,
stride=patch_size,
bias=bias,
**factory_kwargs,
)
nn.init.xavier_uniform_(self.proj.weight.view(self.proj.weight.size(0), -1))
if bias:
nn.init.zeros_(self.proj.bias)
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
def forward(self, x):
x = self.proj(x)
if self.flatten:
x = x.flatten(2).transpose(1, 2) # BCHW -> BNC
x = self.norm(x)
return x
class TextProjection(nn.Module):
"""
Projects text embeddings. Also handles dropout for classifier-free guidance.
Adapted from https://github.com/PixArt-alpha/PixArt-alpha/blob/master/diffusion/model/nets/PixArt_blocks.py
"""
def __init__(self, in_channels, hidden_size, act_layer, dtype=None, device=None):
factory_kwargs = {"dtype": dtype, "device": device}
super().__init__()
self.linear_1 = nn.Linear(
in_features=in_channels,
out_features=hidden_size,
bias=True,
**factory_kwargs,
)
self.act_1 = act_layer()
self.linear_2 = nn.Linear(
in_features=hidden_size,
out_features=hidden_size,
bias=True,
**factory_kwargs,
)
def forward(self, caption):
hidden_states = self.linear_1(caption)
hidden_states = self.act_1(hidden_states)
hidden_states = self.linear_2(hidden_states)
return hidden_states
def timestep_embedding(t, dim, max_period=10000):
"""
Create sinusoidal timestep embeddings.
Args:
t (torch.Tensor): a 1-D Tensor of N indices, one per batch element. These may be fractional.
dim (int): the dimension of the output.
max_period (int): controls the minimum frequency of the embeddings.
Returns:
embedding (torch.Tensor): An (N, D) Tensor of positional embeddings.
.. ref_link: https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
"""
half = dim // 2
freqs = torch.exp(
-math.log(max_period)
* torch.arange(start=0, end=half, dtype=torch.float32)
/ half
).to(device=t.device)
args = t[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
return embedding
class TimestepEmbedder(nn.Module):
"""
Embeds scalar timesteps into vector representations.
"""
def __init__(
self,
hidden_size,
act_layer,
frequency_embedding_size=256,
max_period=10000,
out_size=None,
dtype=None,
device=None,
):
factory_kwargs = {"dtype": dtype, "device": device}
super().__init__()
self.frequency_embedding_size = frequency_embedding_size
self.max_period = max_period
if out_size is None:
out_size = hidden_size
self.mlp = nn.Sequential(
nn.Linear(
frequency_embedding_size, hidden_size, bias=True, **factory_kwargs
),
act_layer(),
nn.Linear(hidden_size, out_size, bias=True, **factory_kwargs),
)
nn.init.normal_(self.mlp[0].weight, std=0.02)
nn.init.normal_(self.mlp[2].weight, std=0.02)
def forward(self, t):
t_freq = timestep_embedding(
t, self.frequency_embedding_size, self.max_period
).type(self.mlp[0].weight.dtype)
t_emb = self.mlp(t_freq)
return t_emb
@@ -1,119 +0,0 @@
# Modified from timm library:
# https://github.com/huggingface/pytorch-image-models/blob/648aaa41233ba83eb38faf5ba9d415d574823241/timm/layers/mlp.py#L13
from functools import partial
import torch
import torch.nn as nn
from .modulate_layers import modulate
from ..utils.helpers import to_2tuple
class MLP(nn.Module):
"""MLP as used in Vision Transformer, MLP-Mixer and related networks"""
def __init__(
self,
in_channels,
hidden_channels=None,
out_features=None,
act_layer=nn.GELU,
norm_layer=None,
bias=True,
drop=0.0,
use_conv=False,
device=None,
dtype=None,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
out_features = out_features or in_channels
hidden_channels = hidden_channels or in_channels
bias = to_2tuple(bias)
drop_probs = to_2tuple(drop)
linear_layer = partial(nn.Conv2d, kernel_size=1) if use_conv else nn.Linear
self.fc1 = linear_layer(
in_channels, hidden_channels, bias=bias[0], **factory_kwargs
)
self.act = act_layer()
self.drop1 = nn.Dropout(drop_probs[0])
self.norm = (
norm_layer(hidden_channels, **factory_kwargs)
if norm_layer is not None
else nn.Identity()
)
self.fc2 = linear_layer(
hidden_channels, out_features, bias=bias[1], **factory_kwargs
)
self.drop2 = nn.Dropout(drop_probs[1])
def forward(self, x):
x = self.fc1(x)
x = self.act(x)
x = self.drop1(x)
x = self.norm(x)
x = self.fc2(x)
x = self.drop2(x)
return x
#
class MLPEmbedder(nn.Module):
"""copied from https://github.com/black-forest-labs/flux/blob/main/src/flux/modules/layers.py"""
def __init__(self, in_dim: int, hidden_dim: int, device=None, dtype=None):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.in_layer = nn.Linear(in_dim, hidden_dim, bias=True, **factory_kwargs)
self.silu = nn.SiLU()
self.out_layer = nn.Linear(hidden_dim, hidden_dim, bias=True, **factory_kwargs)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.out_layer(self.silu(self.in_layer(x)))
class FinalLayer(nn.Module):
"""The final layer of DiT."""
def __init__(
self, hidden_size, patch_size, out_channels, act_layer, device=None, dtype=None
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
# Just use LayerNorm for the final layer
self.norm_final = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
if isinstance(patch_size, int):
self.linear = nn.Linear(
hidden_size,
patch_size * patch_size * out_channels,
bias=True,
**factory_kwargs,
)
else:
self.linear = nn.Linear(
hidden_size,
patch_size[0] * patch_size[1] * patch_size[2] * out_channels,
bias=True,
)
nn.init.zeros_(self.linear.weight)
nn.init.zeros_(self.linear.bias)
# Here we don't distinguish between the modulate types. Just use the simple one.
self.adaLN_modulation = nn.Sequential(
act_layer(),
nn.Linear(hidden_size, 2 * hidden_size, bias=True, **factory_kwargs),
)
# Zero-initialize the modulation
nn.init.zeros_(self.adaLN_modulation[1].weight)
nn.init.zeros_(self.adaLN_modulation[1].bias)
def forward(self, x, c):
shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
x = modulate(self.norm_final(x), shift=shift, scale=scale)
x = self.linear(x)
return x
-757
View File
@@ -1,757 +0,0 @@
from typing import Any, List, Tuple, Optional, Union, Dict
from einops import rearrange
import torch
import torch.nn as nn
import torch.nn.functional as F
from diffusers.models import ModelMixin
from diffusers.configuration_utils import ConfigMixin, register_to_config
from .activation_layers import get_activation_layer
from .norm_layers import get_norm_layer
from .embed_layers import TimestepEmbedder, PatchEmbed, TextProjection
from .attenion import parallel_attention
from .posemb_layers import apply_rotary_emb
from .mlp_layers import MLP, MLPEmbedder, FinalLayer
from .modulate_layers import ModulateDiT, modulate, apply_gate
from .token_refiner import SingleTokenRefiner
from fastvideo.models.hunyuan.modules.posemb_layers import get_nd_rotary_pos_embed
from fastvideo.utils.parallel_states import nccl_info
class MMDoubleStreamBlock(nn.Module):
"""
A multimodal dit block with seperate modulation for
text and image/video, see more details (SD3): https://arxiv.org/abs/2403.03206
(Flux.1): https://github.com/black-forest-labs/flux
"""
def __init__(
self,
hidden_size: int,
heads_num: int,
mlp_width_ratio: float,
mlp_act_type: str = "gelu_tanh",
qk_norm: bool = True,
qk_norm_type: str = "rms",
qkv_bias: bool = False,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.deterministic = False
self.heads_num = heads_num
head_dim = hidden_size // heads_num
mlp_hidden_dim = int(hidden_size * mlp_width_ratio)
self.img_mod = ModulateDiT(
hidden_size,
factor=6,
act_layer=get_activation_layer("silu"),
**factory_kwargs,
)
self.img_norm1 = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.img_attn_qkv = nn.Linear(
hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs
)
qk_norm_layer = get_norm_layer(qk_norm_type)
self.img_attn_q_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.img_attn_k_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.img_attn_proj = nn.Linear(
hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs
)
self.img_norm2 = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.img_mlp = MLP(
hidden_size,
mlp_hidden_dim,
act_layer=get_activation_layer(mlp_act_type),
bias=True,
**factory_kwargs,
)
self.txt_mod = ModulateDiT(
hidden_size,
factor=6,
act_layer=get_activation_layer("silu"),
**factory_kwargs,
)
self.txt_norm1 = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.txt_attn_qkv = nn.Linear(
hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs
)
self.txt_attn_q_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.txt_attn_k_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.txt_attn_proj = nn.Linear(
hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs
)
self.txt_norm2 = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.txt_mlp = MLP(
hidden_size,
mlp_hidden_dim,
act_layer=get_activation_layer(mlp_act_type),
bias=True,
**factory_kwargs,
)
self.hybrid_seq_parallel_attn = None
def enable_deterministic(self):
self.deterministic = True
def disable_deterministic(self):
self.deterministic = False
def forward(
self,
img: torch.Tensor,
txt: torch.Tensor,
vec: torch.Tensor,
freqs_cis: tuple = None,
text_mask: torch.Tensor = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
(
img_mod1_shift,
img_mod1_scale,
img_mod1_gate,
img_mod2_shift,
img_mod2_scale,
img_mod2_gate,
) = self.img_mod(vec).chunk(6, dim=-1)
(
txt_mod1_shift,
txt_mod1_scale,
txt_mod1_gate,
txt_mod2_shift,
txt_mod2_scale,
txt_mod2_gate,
) = self.txt_mod(vec).chunk(6, dim=-1)
# Prepare image for attention.
img_modulated = self.img_norm1(img)
img_modulated = modulate(
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale
)
img_qkv = self.img_attn_qkv(img_modulated)
img_q, img_k, img_v = rearrange(
img_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Apply RoPE if needed.
if freqs_cis is not None:
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
)
freqs_cis = (
shrink_head(freqs_cis[0], dim=0),
shrink_head(freqs_cis[1], dim=0),
)
img_qq, img_kk = apply_rotary_emb(img_q, img_k, freqs_cis, head_first=False)
assert (
img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
img_q, img_k = img_qq, img_kk
# Prepare txt for attention.
txt_modulated = self.txt_norm1(txt)
txt_modulated = modulate(
txt_modulated, shift=txt_mod1_shift, scale=txt_mod1_scale
)
txt_qkv = self.txt_attn_qkv(txt_modulated)
txt_q, txt_k, txt_v = rearrange(
txt_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
# Apply QK-Norm if needed.
txt_q = self.txt_attn_q_norm(txt_q).to(txt_v)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_v)
attn = parallel_attention(
(img_q, txt_q),
(img_k, txt_k),
(img_v, txt_v),
img_q_len=img_q.shape[1],
img_kv_len=img_k.shape[1],
text_mask=text_mask,
)
# attention computation end
img_attn, txt_attn = attn[:, : img.shape[1]], attn[:, img.shape[1] :]
# Calculate the img bloks.
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate)
img = img + apply_gate(
self.img_mlp(
modulate(
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale
)
),
gate=img_mod2_gate,
)
# Calculate the txt bloks.
txt = txt + apply_gate(self.txt_attn_proj(txt_attn), gate=txt_mod1_gate)
txt = txt + apply_gate(
self.txt_mlp(
modulate(
self.txt_norm2(txt), shift=txt_mod2_shift, scale=txt_mod2_scale
)
),
gate=txt_mod2_gate,
)
return img, txt
class MMSingleStreamBlock(nn.Module):
"""
A DiT block with parallel linear layers as described in
https://arxiv.org/abs/2302.05442 and adapted modulation interface.
Also refer to (SD3): https://arxiv.org/abs/2403.03206
(Flux.1): https://github.com/black-forest-labs/flux
"""
def __init__(
self,
hidden_size: int,
heads_num: int,
mlp_width_ratio: float = 4.0,
mlp_act_type: str = "gelu_tanh",
qk_norm: bool = True,
qk_norm_type: str = "rms",
qk_scale: float = None,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.deterministic = False
self.hidden_size = hidden_size
self.heads_num = heads_num
head_dim = hidden_size // heads_num
mlp_hidden_dim = int(hidden_size * mlp_width_ratio)
self.mlp_hidden_dim = mlp_hidden_dim
self.scale = qk_scale or head_dim ** -0.5
# qkv and mlp_in
self.linear1 = nn.Linear(
hidden_size, hidden_size * 3 + mlp_hidden_dim, **factory_kwargs
)
# proj and mlp_out
self.linear2 = nn.Linear(
hidden_size + mlp_hidden_dim, hidden_size, **factory_kwargs
)
qk_norm_layer = get_norm_layer(qk_norm_type)
self.q_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.k_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.pre_norm = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.mlp_act = get_activation_layer(mlp_act_type)()
self.modulation = ModulateDiT(
hidden_size,
factor=3,
act_layer=get_activation_layer("silu"),
**factory_kwargs,
)
self.hybrid_seq_parallel_attn = None
def enable_deterministic(self):
self.deterministic = True
def disable_deterministic(self):
self.deterministic = False
def forward(
self,
x: torch.Tensor,
vec: torch.Tensor,
txt_len: int,
freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None,
text_mask: torch.Tensor = None,
) -> torch.Tensor:
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale)
qkv, mlp = torch.split(
self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1
)
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
# Apply QK-Norm if needed.
q = self.q_norm(q).to(v)
k = self.k_norm(k).to(v)
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
)
freqs_cis = (shrink_head(freqs_cis[0], dim=0), shrink_head(freqs_cis[1], dim=0))
img_q, txt_q = q[:, :-txt_len, :, :], q[:, -txt_len:, :, :]
img_k, txt_k = k[:, :-txt_len, :, :], k[:, -txt_len:, :, :]
img_v, txt_v = v[:, :-txt_len, :, :], v[:, -txt_len:, :, :]
img_qq, img_kk = apply_rotary_emb(img_q, img_k, freqs_cis, head_first=False)
assert (
img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
img_q, img_k = img_qq, img_kk
attn = parallel_attention(
(img_q, txt_q),
(img_k, txt_k),
(img_v, txt_v),
img_q_len=img_q.shape[1],
img_kv_len=img_k.shape[1],
text_mask=text_mask,
)
# attention computation end
# Compute activation in mlp stream, cat again and run second linear layer.
output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2))
return x + apply_gate(output, gate=mod_gate)
class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
"""
HunyuanVideo Transformer backbone
Inherited from ModelMixin and ConfigMixin for compatibility with diffusers' sampler StableDiffusionPipeline.
Reference:
[1] Flux.1: https://github.com/black-forest-labs/flux
[2] MMDiT: http://arxiv.org/abs/2403.03206
Parameters
----------
args: argparse.Namespace
The arguments parsed by argparse.
patch_size: list
The size of the patch.
in_channels: int
The number of input channels.
out_channels: int
The number of output channels.
hidden_size: int
The hidden size of the transformer backbone.
heads_num: int
The number of attention heads.
mlp_width_ratio: float
The ratio of the hidden size of the MLP in the transformer block.
mlp_act_type: str
The activation function of the MLP in the transformer block.
depth_double_blocks: int
The number of transformer blocks in the double blocks.
depth_single_blocks: int
The number of transformer blocks in the single blocks.
rope_dim_list: list
The dimension of the rotary embedding for t, h, w.
qkv_bias: bool
Whether to use bias in the qkv linear layer.
qk_norm: bool
Whether to use qk norm.
qk_norm_type: str
The type of qk norm.
guidance_embed: bool
Whether to use guidance embedding for distillation.
text_projection: str
The type of the text projection, default is single_refiner.
use_attention_mask: bool
Whether to use attention mask for text encoder.
dtype: torch.dtype
The dtype of the model.
device: torch.device
The device of the model.
"""
@register_to_config
def __init__(
self,
patch_size: list = [1, 2, 2],
in_channels: int = 4, # Should be VAE.config.latent_channels.
out_channels: int = None,
hidden_size: int = 3072,
heads_num: int = 24,
mlp_width_ratio: float = 4.0,
mlp_act_type: str = "gelu_tanh",
mm_double_blocks_depth: int = 20,
mm_single_blocks_depth: int = 40,
rope_dim_list: List[int] = [16, 56, 56],
qkv_bias: bool = True,
qk_norm: bool = True,
qk_norm_type: str = "rms",
guidance_embed: bool = False, # For modulation.
text_projection: str = "single_refiner",
use_attention_mask: bool = True,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
text_states_dim: int = 4096,
text_states_dim_2: int = 768,
rope_theta: int = 256,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.patch_size = patch_size
self.in_channels = in_channels
self.out_channels = in_channels if out_channels is None else out_channels
self.unpatchify_channels = self.out_channels
self.guidance_embed = guidance_embed
self.rope_dim_list = rope_dim_list
self.rope_theta = rope_theta
# Text projection. Default to linear projection.
# Alternative: TokenRefiner. See more details (LI-DiT): http://arxiv.org/abs/2406.11831
self.use_attention_mask = use_attention_mask
self.text_projection = text_projection
if hidden_size % heads_num != 0:
raise ValueError(
f"Hidden size {hidden_size} must be divisible by heads_num {heads_num}"
)
pe_dim = hidden_size // heads_num
if sum(rope_dim_list) != pe_dim:
raise ValueError(
f"Got {rope_dim_list} but expected positional dim {pe_dim}"
)
self.hidden_size = hidden_size
self.heads_num = heads_num
# image projection
self.img_in = PatchEmbed(
self.patch_size, self.in_channels, self.hidden_size, **factory_kwargs
)
# text projection
if self.text_projection == "linear":
self.txt_in = TextProjection(
self.config.text_states_dim,
self.hidden_size,
get_activation_layer("silu"),
**factory_kwargs,
)
elif self.text_projection == "single_refiner":
self.txt_in = SingleTokenRefiner(
self.config.text_states_dim,
hidden_size,
heads_num,
depth=2,
**factory_kwargs,
)
else:
raise NotImplementedError(
f"Unsupported text_projection: {self.text_projection}"
)
# time modulation
self.time_in = TimestepEmbedder(
self.hidden_size, get_activation_layer("silu"), **factory_kwargs
)
# text modulation
self.vector_in = MLPEmbedder(
self.config.text_states_dim_2, self.hidden_size, **factory_kwargs
)
# guidance modulation
self.guidance_in = (
TimestepEmbedder(
self.hidden_size, get_activation_layer("silu"), **factory_kwargs
)
if guidance_embed
else None
)
# double blocks
self.double_blocks = nn.ModuleList(
[
MMDoubleStreamBlock(
self.hidden_size,
self.heads_num,
mlp_width_ratio=mlp_width_ratio,
mlp_act_type=mlp_act_type,
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
qkv_bias=qkv_bias,
**factory_kwargs,
)
for _ in range(mm_double_blocks_depth)
]
)
# single blocks
self.single_blocks = nn.ModuleList(
[
MMSingleStreamBlock(
self.hidden_size,
self.heads_num,
mlp_width_ratio=mlp_width_ratio,
mlp_act_type=mlp_act_type,
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
**factory_kwargs,
)
for _ in range(mm_single_blocks_depth)
]
)
self.final_layer = FinalLayer(
self.hidden_size,
self.patch_size,
self.out_channels,
get_activation_layer("silu"),
**factory_kwargs,
)
def enable_deterministic(self):
for block in self.double_blocks:
block.enable_deterministic()
for block in self.single_blocks:
block.enable_deterministic()
def disable_deterministic(self):
for block in self.double_blocks:
block.disable_deterministic()
for block in self.single_blocks:
block.disable_deterministic()
def get_rotary_pos_embed(self, rope_sizes):
target_ndim = 3
ndim = 5 - 2
head_dim = self.hidden_size // self.heads_num
rope_dim_list = self.rope_dim_list
if rope_dim_list is None:
rope_dim_list = [head_dim // target_ndim for _ in range(target_ndim)]
assert (
sum(rope_dim_list) == head_dim
), "sum(rope_dim_list) should equal to head_dim of attention layer"
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
rope_dim_list,
rope_sizes,
theta=self.rope_theta,
use_real=True,
theta_rescale_factor=1,
)
return freqs_cos, freqs_sin
# x: torch.Tensor,
# t: torch.Tensor, # Should be in range(0, 1000).
# text_states: torch.Tensor = None,
# text_mask: torch.Tensor = None, # Now we don't use it.
# text_states_2: Optional[torch.Tensor] = None, # Text embedding for modulation.
# guidance: torch.Tensor = None, # Guidance for modulation, should be cfg_scale x 1000.
# return_dict: bool = True,
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
timestep: torch.LongTensor,
encoder_attention_mask: torch.Tensor,
output_features=False,
output_features_stride=8,
attention_kwargs: Optional[Dict[str, Any]] = None,
return_dict: bool = False,
guidance=None,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
if guidance == None:
guidance = torch.tensor(
[6016.0], device=hidden_states.device, dtype=torch.bfloat16
)
out = {}
img = x = hidden_states
text_mask = encoder_attention_mask
t = timestep
txt = encoder_hidden_states[:, 1:]
text_states_2 = encoder_hidden_states[:, 0, : self.config.text_states_dim_2]
_, _, ot, oh, ow = x.shape
tt, th, tw = (
ot // self.patch_size[0],
oh // self.patch_size[1],
ow // self.patch_size[2],
)
original_tt = nccl_info.sp_size * tt
freqs_cos, freqs_sin = self.get_rotary_pos_embed((original_tt, th, tw))
# Prepare modulation vectors.
vec = self.time_in(t)
# text modulation
vec = vec + self.vector_in(text_states_2)
# guidance modulation
if self.guidance_embed:
if guidance is None:
raise ValueError(
"Didn't get guidance strength for guidance distilled model."
)
# our timestep_embedding is merged into guidance_in(TimestepEmbedder)
vec = vec + self.guidance_in(guidance)
# Embed image and text.
img = self.img_in(img)
if self.text_projection == "linear":
txt = self.txt_in(txt)
elif self.text_projection == "single_refiner":
txt = self.txt_in(txt, t, text_mask if self.use_attention_mask else None)
else:
raise NotImplementedError(
f"Unsupported text_projection: {self.text_projection}"
)
txt_seq_len = txt.shape[1]
img_seq_len = img.shape[1]
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# --------------------- Pass through DiT blocks ------------------------
for _, block in enumerate(self.double_blocks):
double_block_args = [img, txt, vec, freqs_cis, text_mask]
img, txt = block(*double_block_args)
# Merge txt and img to pass through single stream blocks.
x = torch.cat((img, txt), 1)
if output_features:
features_list = []
if len(self.single_blocks) > 0:
for _, block in enumerate(self.single_blocks):
single_block_args = [
x,
vec,
txt_seq_len,
(freqs_cos, freqs_sin),
text_mask,
]
x = block(*single_block_args)
if output_features and _ % output_features_stride == 0:
features_list.append(x[:, :img_seq_len, ...])
img = x[:, :img_seq_len, ...]
# ---------------------------- Final layer ------------------------------
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
img = self.unpatchify(img, tt, th, tw)
assert return_dict == False, "return_dict is not supported."
if output_features:
features_list = torch.stack(features_list, dim=0)
else:
features_list = None
return (img, features_list)
def unpatchify(self, x, t, h, w):
"""
x: (N, T, patch_size**2 * C)
imgs: (N, H, W, C)
"""
c = self.unpatchify_channels
pt, ph, pw = self.patch_size
assert t * h * w == x.shape[1]
x = x.reshape(shape=(x.shape[0], t, h, w, c, pt, ph, pw))
x = torch.einsum("nthwcopq->nctohpwq", x)
imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))
return imgs
def params_count(self):
counts = {
"double": sum(
[
sum(p.numel() for p in block.img_attn_qkv.parameters())
+ sum(p.numel() for p in block.img_attn_proj.parameters())
+ sum(p.numel() for p in block.img_mlp.parameters())
+ sum(p.numel() for p in block.txt_attn_qkv.parameters())
+ sum(p.numel() for p in block.txt_attn_proj.parameters())
+ sum(p.numel() for p in block.txt_mlp.parameters())
for block in self.double_blocks
]
),
"single": sum(
[
sum(p.numel() for p in block.linear1.parameters())
+ sum(p.numel() for p in block.linear2.parameters())
for block in self.single_blocks
]
),
"total": sum(p.numel() for p in self.parameters()),
}
counts["attn+mlp"] = counts["double"] + counts["single"]
return counts
#################################################################################
# HunyuanVideo Configs #
#################################################################################
HUNYUAN_VIDEO_CONFIG = {
"HYVideo-T/2": {
"mm_double_blocks_depth": 20,
"mm_single_blocks_depth": 40,
"rope_dim_list": [16, 56, 56],
"hidden_size": 3072,
"heads_num": 24,
"mlp_width_ratio": 4,
},
"HYVideo-T/2-cfgdistill": {
"mm_double_blocks_depth": 20,
"mm_single_blocks_depth": 40,
"rope_dim_list": [16, 56, 56],
"hidden_size": 3072,
"heads_num": 24,
"mlp_width_ratio": 4,
"guidance_embed": True,
},
}
@@ -1,156 +0,0 @@
from typing import Callable
import torch
import torch.nn as nn
class ModulateDiT(nn.Module):
"""Modulation layer for DiT."""
def __init__(
self,
hidden_size: int,
factor: int,
act_layer: Callable,
dtype=None,
device=None,
):
factory_kwargs = {"dtype": dtype, "device": device}
super().__init__()
self.act = act_layer()
self.linear = nn.Linear(
hidden_size, factor * hidden_size, bias=True, **factory_kwargs
)
# Zero-initialize the modulation
nn.init.zeros_(self.linear.weight)
nn.init.zeros_(self.linear.bias)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.linear(self.act(x))
def modulate(x, shift=None, scale=None):
"""modulate by shift and scale
Args:
x (torch.Tensor): input tensor.
shift (torch.Tensor, optional): shift tensor. Defaults to None.
scale (torch.Tensor, optional): scale tensor. Defaults to None.
Returns:
torch.Tensor: the output tensor after modulate.
"""
if scale is None and shift is None:
return x
elif shift is None:
return x * (1 + scale.unsqueeze(1))
elif scale is None:
return x + shift.unsqueeze(1)
else:
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
def apply_gate(x, gate=None, tanh=False):
"""AI is creating summary for apply_gate
Args:
x (torch.Tensor): input tensor.
gate (torch.Tensor, optional): gate tensor. Defaults to None.
tanh (bool, optional): whether to use tanh function. Defaults to False.
Returns:
torch.Tensor: the output tensor after apply gate.
"""
if gate is None:
return x
if tanh:
return x * gate.unsqueeze(1).tanh()
else:
return x * gate.unsqueeze(1)
def ckpt_wrapper(module):
def ckpt_forward(*inputs):
outputs = module(*inputs)
return outputs
return ckpt_forward
import torch
import torch.nn as nn
class RMSNorm(nn.Module):
def __init__(
self,
dim: int,
elementwise_affine=True,
eps: float = 1e-6,
device=None,
dtype=None,
):
"""
Initialize the RMSNorm normalization layer.
Args:
dim (int): The dimension of the input tensor.
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
Attributes:
eps (float): A small value added to the denominator for numerical stability.
weight (nn.Parameter): Learnable scaling parameter.
"""
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.eps = eps
if elementwise_affine:
self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))
def _norm(self, x):
"""
Apply the RMSNorm normalization to the input tensor.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The normalized tensor.
"""
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
def forward(self, x):
"""
Forward pass through the RMSNorm layer.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The output tensor after applying RMSNorm.
"""
output = self._norm(x.float()).type_as(x)
if hasattr(self, "weight"):
output = output * self.weight
return output
def get_norm_layer(norm_layer):
"""
Get the normalization layer.
Args:
norm_layer (str): The type of normalization layer.
Returns:
norm_layer (nn.Module): The normalization layer.
"""
if norm_layer == "layer":
return nn.LayerNorm
elif norm_layer == "rms":
return RMSNorm
else:
raise NotImplementedError(f"Norm layer {norm_layer} is not implemented")
@@ -1,77 +0,0 @@
import torch
import torch.nn as nn
class RMSNorm(nn.Module):
def __init__(
self,
dim: int,
elementwise_affine=True,
eps: float = 1e-6,
device=None,
dtype=None,
):
"""
Initialize the RMSNorm normalization layer.
Args:
dim (int): The dimension of the input tensor.
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
Attributes:
eps (float): A small value added to the denominator for numerical stability.
weight (nn.Parameter): Learnable scaling parameter.
"""
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.eps = eps
if elementwise_affine:
self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))
def _norm(self, x):
"""
Apply the RMSNorm normalization to the input tensor.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The normalized tensor.
"""
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
def forward(self, x):
"""
Forward pass through the RMSNorm layer.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The output tensor after applying RMSNorm.
"""
output = self._norm(x.float()).type_as(x)
if hasattr(self, "weight"):
output = output * self.weight
return output
def get_norm_layer(norm_layer):
"""
Get the normalization layer.
Args:
norm_layer (str): The type of normalization layer.
Returns:
norm_layer (nn.Module): The normalization layer.
"""
if norm_layer == "layer":
return nn.LayerNorm
elif norm_layer == "rms":
return RMSNorm
else:
raise NotImplementedError(f"Norm layer {norm_layer} is not implemented")
@@ -1,310 +0,0 @@
import torch
from typing import Union, Tuple, List
def _to_tuple(x, dim=2):
if isinstance(x, int):
return (x,) * dim
elif len(x) == dim:
return x
else:
raise ValueError(f"Expected length {dim} or int, but got {x}")
def get_meshgrid_nd(start, *args, dim=2):
"""
Get n-D meshgrid with start, stop and num.
Args:
start (int or tuple): If len(args) == 0, start is num; If len(args) == 1, start is start, args[0] is stop,
step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num. For n-dim, start/stop/num
should be int or n-tuple. If n-tuple is provided, the meshgrid will be stacked following the dim order in
n-tuples.
*args: See above.
dim (int): Dimension of the meshgrid. Defaults to 2.
Returns:
grid (np.ndarray): [dim, ...]
"""
if len(args) == 0:
# start is grid_size
num = _to_tuple(start, dim=dim)
start = (0,) * dim
stop = num
elif len(args) == 1:
# start is start, args[0] is stop, step is 1
start = _to_tuple(start, dim=dim)
stop = _to_tuple(args[0], dim=dim)
num = [stop[i] - start[i] for i in range(dim)]
elif len(args) == 2:
# start is start, args[0] is stop, args[1] is num
start = _to_tuple(start, dim=dim) # Left-Top eg: 12,0
stop = _to_tuple(args[0], dim=dim) # Right-Bottom eg: 20,32
num = _to_tuple(args[1], dim=dim) # Target Size eg: 32,124
else:
raise ValueError(f"len(args) should be 0, 1 or 2, but got {len(args)}")
# PyTorch implement of np.linspace(start[i], stop[i], num[i], endpoint=False)
axis_grid = []
for i in range(dim):
a, b, n = start[i], stop[i], num[i]
g = torch.linspace(a, b, n + 1, dtype=torch.float32)[:n]
axis_grid.append(g)
grid = torch.meshgrid(*axis_grid, indexing="ij") # dim x [W, H, D]
grid = torch.stack(grid, dim=0) # [dim, W, H, D]
return grid
#################################################################################
# Rotary Positional Embedding Functions #
#################################################################################
# https://github.com/meta-llama/llama/blob/be327c427cc5e89cc1d3ab3d3fec4484df771245/llama/model.py#L80
def reshape_for_broadcast(
freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]],
x: torch.Tensor,
head_first=False,
):
"""
Reshape frequency tensor for broadcasting it with another tensor.
This function reshapes the frequency tensor to have the same shape as the target tensor 'x'
for the purpose of broadcasting the frequency tensor during element-wise operations.
Notes:
When using FlashMHAModified, head_first should be False.
When using Attention, head_first should be True.
Args:
freqs_cis (Union[torch.Tensor, Tuple[torch.Tensor]]): Frequency tensor to be reshaped.
x (torch.Tensor): Target tensor for broadcasting compatibility.
head_first (bool): head dimension first (except batch dim) or not.
Returns:
torch.Tensor: Reshaped frequency tensor.
Raises:
AssertionError: If the frequency tensor doesn't match the expected shape.
AssertionError: If the target tensor 'x' doesn't have the expected number of dimensions.
"""
ndim = x.ndim
assert 0 <= 1 < ndim
if isinstance(freqs_cis, tuple):
# freqs_cis: (cos, sin) in real space
if head_first:
assert freqs_cis[0].shape == (
x.shape[-2],
x.shape[-1],
), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}"
shape = [
d if i == ndim - 2 or i == ndim - 1 else 1
for i, d in enumerate(x.shape)
]
else:
assert freqs_cis[0].shape == (
x.shape[1],
x.shape[-1],
), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}"
shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
return freqs_cis[0].view(*shape), freqs_cis[1].view(*shape)
else:
# freqs_cis: values in complex space
if head_first:
assert freqs_cis.shape == (
x.shape[-2],
x.shape[-1],
), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}"
shape = [
d if i == ndim - 2 or i == ndim - 1 else 1
for i, d in enumerate(x.shape)
]
else:
assert freqs_cis.shape == (
x.shape[1],
x.shape[-1],
), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}"
shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
return freqs_cis.view(*shape)
def rotate_half(x):
x_real, x_imag = (
x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1)
) # [B, S, H, D//2]
return torch.stack([-x_imag, x_real], dim=-1).flatten(3)
def apply_rotary_emb(
xq: torch.Tensor,
xk: torch.Tensor,
freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]],
head_first: bool = False,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Apply rotary embeddings to input tensors using the given frequency tensor.
This function applies rotary embeddings to the given query 'xq' and key 'xk' tensors using the provided
frequency tensor 'freqs_cis'. The input tensors are reshaped as complex numbers, and the frequency tensor
is reshaped for broadcasting compatibility. The resulting tensors contain rotary embeddings and are
returned as real tensors.
Args:
xq (torch.Tensor): Query tensor to apply rotary embeddings. [B, S, H, D]
xk (torch.Tensor): Key tensor to apply rotary embeddings. [B, S, H, D]
freqs_cis (torch.Tensor or tuple): Precomputed frequency tensor for complex exponential.
head_first (bool): head dimension first (except batch dim) or not.
Returns:
Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings.
"""
xk_out = None
if isinstance(freqs_cis, tuple):
cos, sin = reshape_for_broadcast(freqs_cis, xq, head_first) # [S, D]
cos, sin = cos.to(xq.device), sin.to(xq.device)
# real * cos - imag * sin
# imag * cos + real * sin
xq_out = (xq.float() * cos + rotate_half(xq.float()) * sin).type_as(xq)
xk_out = (xk.float() * cos + rotate_half(xk.float()) * sin).type_as(xk)
else:
# view_as_complex will pack [..., D/2, 2](real) to [..., D/2](complex)
xq_ = torch.view_as_complex(
xq.float().reshape(*xq.shape[:-1], -1, 2)
) # [B, S, H, D//2]
freqs_cis = reshape_for_broadcast(freqs_cis, xq_, head_first).to(
xq.device
) # [S, D//2] --> [1, S, 1, D//2]
# (real, imag) * (cos, sin) = (real * cos - imag * sin, imag * cos + real * sin)
# view_as_real will expand [..., D/2](complex) to [..., D/2, 2](real)
xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(3).type_as(xq)
xk_ = torch.view_as_complex(
xk.float().reshape(*xk.shape[:-1], -1, 2)
) # [B, S, H, D//2]
xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3).type_as(xk)
return xq_out, xk_out
def get_nd_rotary_pos_embed(
rope_dim_list,
start,
*args,
theta=10000.0,
use_real=False,
theta_rescale_factor: Union[float, List[float]] = 1.0,
interpolation_factor: Union[float, List[float]] = 1.0,
):
"""
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
Args:
rope_dim_list (list of int): Dimension of each rope. len(rope_dim_list) should equal to n.
sum(rope_dim_list) should equal to head_dim of attention layer.
start (int | tuple of int | list of int): If len(args) == 0, start is num; If len(args) == 1, start is start,
args[0] is stop, step is 1; If len(args) == 2, start is start, args[0] is stop, args[1] is num.
*args: See above.
theta (float): Scaling factor for frequency computation. Defaults to 10000.0.
use_real (bool): If True, return real part and imaginary part separately. Otherwise, return complex numbers.
Some libraries such as TensorRT does not support complex64 data type. So it is useful to provide a real
part and an imaginary part separately.
theta_rescale_factor (float): Rescale factor for theta. Defaults to 1.0.
Returns:
pos_embed (torch.Tensor): [HW, D/2]
"""
grid = get_meshgrid_nd(
start, *args, dim=len(rope_dim_list)
) # [3, W, H, D] / [2, W, H]
if isinstance(theta_rescale_factor, int) or isinstance(theta_rescale_factor, float):
theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
elif isinstance(theta_rescale_factor, list) and len(theta_rescale_factor) == 1:
theta_rescale_factor = [theta_rescale_factor[0]] * len(rope_dim_list)
assert len(theta_rescale_factor) == len(
rope_dim_list
), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
if isinstance(interpolation_factor, int) or isinstance(interpolation_factor, float):
interpolation_factor = [interpolation_factor] * len(rope_dim_list)
elif isinstance(interpolation_factor, list) and len(interpolation_factor) == 1:
interpolation_factor = [interpolation_factor[0]] * len(rope_dim_list)
assert len(interpolation_factor) == len(
rope_dim_list
), "len(interpolation_factor) should equal to len(rope_dim_list)"
# use 1/ndim of dimensions to encode grid_axis
embs = []
for i in range(len(rope_dim_list)):
emb = get_1d_rotary_pos_embed(
rope_dim_list[i],
grid[i].reshape(-1),
theta,
use_real=use_real,
theta_rescale_factor=theta_rescale_factor[i],
interpolation_factor=interpolation_factor[i],
) # 2 x [WHD, rope_dim_list[i]]
embs.append(emb)
if use_real:
cos = torch.cat([emb[0] for emb in embs], dim=1) # (WHD, D/2)
sin = torch.cat([emb[1] for emb in embs], dim=1) # (WHD, D/2)
return cos, sin
else:
emb = torch.cat(embs, dim=1) # (WHD, D/2)
return emb
def get_1d_rotary_pos_embed(
dim: int,
pos: Union[torch.FloatTensor, int],
theta: float = 10000.0,
use_real: bool = False,
theta_rescale_factor: float = 1.0,
interpolation_factor: float = 1.0,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
(Note: `cis` means `cos + i * sin`, where i is the imaginary unit.)
This function calculates a frequency tensor with complex exponential using the given dimension 'dim'
and the end index 'end'. The 'theta' parameter scales the frequencies.
The returned tensor contains complex values in complex64 data type.
Args:
dim (int): Dimension of the frequency tensor.
pos (int or torch.FloatTensor): Position indices for the frequency tensor. [S] or scalar
theta (float, optional): Scaling factor for frequency computation. Defaults to 10000.0.
use_real (bool, optional): If True, return real part and imaginary part separately.
Otherwise, return complex numbers.
theta_rescale_factor (float, optional): Rescale factor for theta. Defaults to 1.0.
Returns:
freqs_cis: Precomputed frequency tensor with complex exponential. [S, D/2]
freqs_cos, freqs_sin: Precomputed frequency tensor with real and imaginary parts separately. [S, D]
"""
if isinstance(pos, int):
pos = torch.arange(pos).float()
# proposed by reddit user bloc97, to rescale rotary embeddings to longer sequence length without fine-tuning
# has some connection to NTK literature
if theta_rescale_factor != 1.0:
theta *= theta_rescale_factor ** (dim / (dim - 2))
freqs = 1.0 / (
theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)
) # [D/2]
# assert interpolation_factor == 1.0, f"interpolation_factor: {interpolation_factor}"
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
if use_real:
freqs_cos = freqs.cos().repeat_interleave(2, dim=1) # [S, D]
freqs_sin = freqs.sin().repeat_interleave(2, dim=1) # [S, D]
return freqs_cos, freqs_sin
else:
freqs_cis = torch.polar(
torch.ones_like(freqs), freqs
) # complex64 # [S, D/2]
return freqs_cis
@@ -1,221 +0,0 @@
from typing import Optional
from einops import rearrange
import torch
import torch.nn as nn
from .activation_layers import get_activation_layer
from .attenion import attention
from .norm_layers import get_norm_layer
from .embed_layers import TimestepEmbedder, TextProjection
from .attenion import attention
from .mlp_layers import MLP
from .modulate_layers import modulate, apply_gate
class IndividualTokenRefinerBlock(nn.Module):
def __init__(
self,
hidden_size,
heads_num,
mlp_width_ratio: str = 4.0,
mlp_drop_rate: float = 0.0,
act_type: str = "silu",
qk_norm: bool = False,
qk_norm_type: str = "layer",
qkv_bias: bool = True,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.heads_num = heads_num
head_dim = hidden_size // heads_num
mlp_hidden_dim = int(hidden_size * mlp_width_ratio)
self.norm1 = nn.LayerNorm(
hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs
)
self.self_attn_qkv = nn.Linear(
hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs
)
qk_norm_layer = get_norm_layer(qk_norm_type)
self.self_attn_q_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.self_attn_k_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.self_attn_proj = nn.Linear(
hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs
)
self.norm2 = nn.LayerNorm(
hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs
)
act_layer = get_activation_layer(act_type)
self.mlp = MLP(
in_channels=hidden_size,
hidden_channels=mlp_hidden_dim,
act_layer=act_layer,
drop=mlp_drop_rate,
**factory_kwargs,
)
self.adaLN_modulation = nn.Sequential(
act_layer(),
nn.Linear(hidden_size, 2 * hidden_size, bias=True, **factory_kwargs),
)
# Zero-initialize the modulation
nn.init.zeros_(self.adaLN_modulation[1].weight)
nn.init.zeros_(self.adaLN_modulation[1].bias)
def forward(
self,
x: torch.Tensor,
c: torch.Tensor, # timestep_aware_representations + context_aware_representations
attn_mask: torch.Tensor = None,
):
gate_msa, gate_mlp = self.adaLN_modulation(c).chunk(2, dim=1)
norm_x = self.norm1(x)
qkv = self.self_attn_qkv(norm_x)
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
# Apply QK-Norm if needed
q = self.self_attn_q_norm(q).to(v)
k = self.self_attn_k_norm(k).to(v)
# Self-Attention
attn = attention(q, k, v, attn_mask=attn_mask)
x = x + apply_gate(self.self_attn_proj(attn), gate_msa)
# FFN Layer
x = x + apply_gate(self.mlp(self.norm2(x)), gate_mlp)
return x
class IndividualTokenRefiner(nn.Module):
def __init__(
self,
hidden_size,
heads_num,
depth,
mlp_width_ratio: float = 4.0,
mlp_drop_rate: float = 0.0,
act_type: str = "silu",
qk_norm: bool = False,
qk_norm_type: str = "layer",
qkv_bias: bool = True,
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.blocks = nn.ModuleList(
[
IndividualTokenRefinerBlock(
hidden_size=hidden_size,
heads_num=heads_num,
mlp_width_ratio=mlp_width_ratio,
mlp_drop_rate=mlp_drop_rate,
act_type=act_type,
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
qkv_bias=qkv_bias,
**factory_kwargs,
)
for _ in range(depth)
]
)
def forward(
self, x: torch.Tensor, c: torch.LongTensor, mask: Optional[torch.Tensor] = None,
):
mask = mask.clone().bool()
# avoid attention weight become NaN
mask[:, 0] = True
for block in self.blocks:
x = block(x, c, mask)
return x
class SingleTokenRefiner(nn.Module):
"""
A single token refiner block for llm text embedding refine.
"""
def __init__(
self,
in_channels,
hidden_size,
heads_num,
depth,
mlp_width_ratio: float = 4.0,
mlp_drop_rate: float = 0.0,
act_type: str = "silu",
qk_norm: bool = False,
qk_norm_type: str = "layer",
qkv_bias: bool = True,
attn_mode: str = "torch",
dtype: Optional[torch.dtype] = None,
device: Optional[torch.device] = None,
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.attn_mode = attn_mode
assert self.attn_mode == "torch", "Only support 'torch' mode for token refiner."
self.input_embedder = nn.Linear(
in_channels, hidden_size, bias=True, **factory_kwargs
)
act_layer = get_activation_layer(act_type)
# Build timestep embedding layer
self.t_embedder = TimestepEmbedder(hidden_size, act_layer, **factory_kwargs)
# Build context embedding layer
self.c_embedder = TextProjection(
in_channels, hidden_size, act_layer, **factory_kwargs
)
self.individual_token_refiner = IndividualTokenRefiner(
hidden_size=hidden_size,
heads_num=heads_num,
depth=depth,
mlp_width_ratio=mlp_width_ratio,
mlp_drop_rate=mlp_drop_rate,
act_type=act_type,
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
qkv_bias=qkv_bias,
**factory_kwargs,
)
def forward(
self,
x: torch.Tensor,
t: torch.LongTensor,
mask: Optional[torch.LongTensor] = None,
):
timestep_aware_representations = self.t_embedder(t)
if mask is None:
context_aware_representations = x.mean(dim=1)
else:
mask_float = mask.float().unsqueeze(-1) # [b, s1, 1]
context_aware_representations = (x * mask_float).sum(
dim=1
) / mask_float.sum(dim=1)
context_aware_representations = self.c_embedder(context_aware_representations)
c = timestep_aware_representations + context_aware_representations
x = self.input_embedder(x)
x = self.individual_token_refiner(x, c, mask)
return x
@@ -1,53 +0,0 @@
normal_mode_prompt = """Normal mode - Video Recaption Task:
You are a large language model specialized in rewriting video descriptions. Your task is to modify the input description.
0. Preserve ALL information, including style words and technical terms.
1. If the input is in Chinese, translate the entire description to English.
2. If the input is just one or two words describing an object or person, provide a brief, simple description focusing on basic visual characteristics. Limit the description to 1-2 short sentences.
3. If the input does not include style, lighting, atmosphere, you can make reasonable associations.
4. Output ALL must be in English.
Given Input:
input: "{input}"
"""
master_mode_prompt = """Master mode - Video Recaption Task:
You are a large language model specialized in rewriting video descriptions. Your task is to modify the input description.
0. Preserve ALL information, including style words and technical terms.
1. If the input is in Chinese, translate the entire description to English.
2. If the input is just one or two words describing an object or person, provide a brief, simple description focusing on basic visual characteristics. Limit the description to 1-2 short sentences.
3. If the input does not include style, lighting, atmosphere, you can make reasonable associations.
4. Output ALL must be in English.
Given Input:
input: "{input}"
"""
def get_rewrite_prompt(ori_prompt, mode="Normal"):
if mode == "Normal":
prompt = normal_mode_prompt.format(input=ori_prompt)
elif mode == "Master":
prompt = master_mode_prompt.format(input=ori_prompt)
else:
raise Exception("Only supports Normal and Normal", mode)
return prompt
ori_prompt = "一只小狗在草地上奔跑。"
normal_prompt = get_rewrite_prompt(ori_prompt, mode="Normal")
master_prompt = get_rewrite_prompt(ori_prompt, mode="Master")
# Then you can use the normal_prompt or master_prompt to access the hunyuan-large rewrite model to get the final prompt.
@@ -1,357 +0,0 @@
from dataclasses import dataclass
from typing import Optional, Tuple
from copy import deepcopy
import torch
import torch.nn as nn
from transformers import CLIPTextModel, CLIPTokenizer, AutoTokenizer, AutoModel
from transformers.utils import ModelOutput
from ..constants import TEXT_ENCODER_PATH, TOKENIZER_PATH
from ..constants import PRECISION_TO_TYPE
def use_default(value, default):
return value if value is not None else default
def load_text_encoder(
text_encoder_type,
text_encoder_precision=None,
text_encoder_path=None,
logger=None,
device=None,
):
if text_encoder_path is None:
text_encoder_path = TEXT_ENCODER_PATH[text_encoder_type]
if logger is not None:
logger.info(
f"Loading text encoder model ({text_encoder_type}) from: {text_encoder_path}"
)
if text_encoder_type == "clipL":
text_encoder = CLIPTextModel.from_pretrained(text_encoder_path)
text_encoder.final_layer_norm = text_encoder.text_model.final_layer_norm
elif text_encoder_type == "llm":
text_encoder = AutoModel.from_pretrained(
text_encoder_path, low_cpu_mem_usage=True
)
text_encoder.final_layer_norm = text_encoder.norm
else:
raise ValueError(f"Unsupported text encoder type: {text_encoder_type}")
# from_pretrained will ensure that the model is in eval mode.
if text_encoder_precision is not None:
text_encoder = text_encoder.to(dtype=PRECISION_TO_TYPE[text_encoder_precision])
text_encoder.requires_grad_(False)
if logger is not None:
logger.info(f"Text encoder to dtype: {text_encoder.dtype}")
if device is not None:
text_encoder = text_encoder.to(device)
return text_encoder, text_encoder_path
def load_tokenizer(
tokenizer_type, tokenizer_path=None, padding_side="right", logger=None
):
if tokenizer_path is None:
tokenizer_path = TOKENIZER_PATH[tokenizer_type]
if logger is not None:
logger.info(f"Loading tokenizer ({tokenizer_type}) from: {tokenizer_path}")
if tokenizer_type == "clipL":
tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path, max_length=77)
elif tokenizer_type == "llm":
tokenizer = AutoTokenizer.from_pretrained(
tokenizer_path, padding_side=padding_side
)
else:
raise ValueError(f"Unsupported tokenizer type: {tokenizer_type}")
return tokenizer, tokenizer_path
@dataclass
class TextEncoderModelOutput(ModelOutput):
"""
Base class for model's outputs that also contains a pooling of the last hidden states.
Args:
hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
Sequence of hidden-states at the output of the last layer of the model.
attention_mask (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
Mask to avoid performing attention on padding token indices. Mask values selected in ``[0, 1]``:
hidden_states_list (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed):
Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +
one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.
Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
text_outputs (`list`, *optional*, returned when `return_texts=True` is passed):
List of decoded texts.
"""
hidden_state: torch.FloatTensor = None
attention_mask: Optional[torch.LongTensor] = None
hidden_states_list: Optional[Tuple[torch.FloatTensor, ...]] = None
text_outputs: Optional[list] = None
class TextEncoder(nn.Module):
def __init__(
self,
text_encoder_type: str,
max_length: int,
text_encoder_precision: Optional[str] = None,
text_encoder_path: Optional[str] = None,
tokenizer_type: Optional[str] = None,
tokenizer_path: Optional[str] = None,
output_key: Optional[str] = None,
use_attention_mask: bool = True,
input_max_length: Optional[int] = None,
prompt_template: Optional[dict] = None,
prompt_template_video: Optional[dict] = None,
hidden_state_skip_layer: Optional[int] = None,
apply_final_norm: bool = False,
reproduce: bool = False,
logger=None,
device=None,
):
super().__init__()
self.text_encoder_type = text_encoder_type
self.max_length = max_length
self.precision = text_encoder_precision
self.model_path = text_encoder_path
self.tokenizer_type = (
tokenizer_type if tokenizer_type is not None else text_encoder_type
)
self.tokenizer_path = (
tokenizer_path if tokenizer_path is not None else text_encoder_path
)
self.use_attention_mask = use_attention_mask
if prompt_template_video is not None:
assert (
use_attention_mask is True
), "Attention mask is True required when training videos."
self.input_max_length = (
input_max_length if input_max_length is not None else max_length
)
self.prompt_template = prompt_template
self.prompt_template_video = prompt_template_video
self.hidden_state_skip_layer = hidden_state_skip_layer
self.apply_final_norm = apply_final_norm
self.reproduce = reproduce
self.logger = logger
self.use_template = self.prompt_template is not None
if self.use_template:
assert (
isinstance(self.prompt_template, dict)
and "template" in self.prompt_template
), f"`prompt_template` must be a dictionary with a key 'template', got {self.prompt_template}"
assert "{}" in str(self.prompt_template["template"]), (
"`prompt_template['template']` must contain a placeholder `{}` for the input text, "
f"got {self.prompt_template['template']}"
)
self.use_video_template = self.prompt_template_video is not None
if self.use_video_template:
if self.prompt_template_video is not None:
assert (
isinstance(self.prompt_template_video, dict)
and "template" in self.prompt_template_video
), f"`prompt_template_video` must be a dictionary with a key 'template', got {self.prompt_template_video}"
assert "{}" in str(self.prompt_template_video["template"]), (
"`prompt_template_video['template']` must contain a placeholder `{}` for the input text, "
f"got {self.prompt_template_video['template']}"
)
if "t5" in text_encoder_type:
self.output_key = output_key or "last_hidden_state"
elif "clip" in text_encoder_type:
self.output_key = output_key or "pooler_output"
elif "llm" in text_encoder_type or "glm" in text_encoder_type:
self.output_key = output_key or "last_hidden_state"
else:
raise ValueError(f"Unsupported text encoder type: {text_encoder_type}")
self.model, self.model_path = load_text_encoder(
text_encoder_type=self.text_encoder_type,
text_encoder_precision=self.precision,
text_encoder_path=self.model_path,
logger=self.logger,
device=device,
)
self.dtype = self.model.dtype
self.device = self.model.device
self.tokenizer, self.tokenizer_path = load_tokenizer(
tokenizer_type=self.tokenizer_type,
tokenizer_path=self.tokenizer_path,
padding_side="right",
logger=self.logger,
)
def __repr__(self):
return f"{self.text_encoder_type} ({self.precision} - {self.model_path})"
@staticmethod
def apply_text_to_template(text, template, prevent_empty_text=True):
"""
Apply text to template.
Args:
text (str): Input text.
template (str or list): Template string or list of chat conversation.
prevent_empty_text (bool): If Ture, we will prevent the user text from being empty
by adding a space. Defaults to True.
"""
if isinstance(template, str):
# Will send string to tokenizer. Used for llm
return template.format(text)
else:
raise TypeError(f"Unsupported template type: {type(template)}")
def text2tokens(self, text, data_type="image"):
"""
Tokenize the input text.
Args:
text (str or list): Input text.
"""
tokenize_input_type = "str"
if self.use_template:
if data_type == "image":
prompt_template = self.prompt_template["template"]
elif data_type == "video":
prompt_template = self.prompt_template_video["template"]
else:
raise ValueError(f"Unsupported data type: {data_type}")
if isinstance(text, (list, tuple)):
text = [
self.apply_text_to_template(one_text, prompt_template)
for one_text in text
]
if isinstance(text[0], list):
tokenize_input_type = "list"
elif isinstance(text, str):
text = self.apply_text_to_template(text, prompt_template)
if isinstance(text, list):
tokenize_input_type = "list"
else:
raise TypeError(f"Unsupported text type: {type(text)}")
kwargs = dict(
truncation=True,
max_length=self.max_length,
padding="max_length",
return_tensors="pt",
)
if tokenize_input_type == "str":
return self.tokenizer(
text,
return_length=False,
return_overflowing_tokens=False,
return_attention_mask=True,
**kwargs,
)
elif tokenize_input_type == "list":
return self.tokenizer.apply_chat_template(
text,
add_generation_prompt=True,
tokenize=True,
return_dict=True,
**kwargs,
)
else:
raise ValueError(f"Unsupported tokenize_input_type: {tokenize_input_type}")
def encode(
self,
batch_encoding,
use_attention_mask=None,
output_hidden_states=False,
do_sample=None,
hidden_state_skip_layer=None,
return_texts=False,
data_type="image",
device=None,
):
"""
Args:
batch_encoding (dict): Batch encoding from tokenizer.
use_attention_mask (bool): Whether to use attention mask. If None, use self.use_attention_mask.
Defaults to None.
output_hidden_states (bool): Whether to output hidden states. If False, return the value of
self.output_key. If True, return the entire output. If set self.hidden_state_skip_layer,
output_hidden_states will be set True. Defaults to False.
do_sample (bool): Whether to sample from the model. Used for Decoder-Only LLMs. Defaults to None.
When self.produce is False, do_sample is set to True by default.
hidden_state_skip_layer (int): Number of hidden states to hidden_state_skip_layer. 0 means the last layer.
If None, self.output_key will be used. Defaults to None.
return_texts (bool): Whether to return the decoded texts. Defaults to False.
"""
device = self.model.device if device is None else device
use_attention_mask = use_default(use_attention_mask, self.use_attention_mask)
hidden_state_skip_layer = use_default(
hidden_state_skip_layer, self.hidden_state_skip_layer
)
do_sample = use_default(do_sample, not self.reproduce)
attention_mask = (
batch_encoding["attention_mask"].to(device) if use_attention_mask else None
)
outputs = self.model(
input_ids=batch_encoding["input_ids"].to(device),
attention_mask=attention_mask,
output_hidden_states=output_hidden_states
or hidden_state_skip_layer is not None,
)
if hidden_state_skip_layer is not None:
last_hidden_state = outputs.hidden_states[-(hidden_state_skip_layer + 1)]
# Real last hidden state already has layer norm applied. So here we only apply it
# for intermediate layers.
if hidden_state_skip_layer > 0 and self.apply_final_norm:
last_hidden_state = self.model.final_layer_norm(last_hidden_state)
else:
last_hidden_state = outputs[self.output_key]
# Remove hidden states of instruction tokens, only keep prompt tokens.
if self.use_template:
if data_type == "image":
crop_start = self.prompt_template.get("crop_start", -1)
elif data_type == "video":
crop_start = self.prompt_template_video.get("crop_start", -1)
else:
raise ValueError(f"Unsupported data type: {data_type}")
if crop_start > 0:
last_hidden_state = last_hidden_state[:, crop_start:]
attention_mask = (
attention_mask[:, crop_start:] if use_attention_mask else None
)
if output_hidden_states:
return TextEncoderModelOutput(
last_hidden_state, attention_mask, outputs.hidden_states
)
return TextEncoderModelOutput(last_hidden_state, attention_mask)
def forward(
self,
text,
use_attention_mask=None,
output_hidden_states=False,
do_sample=False,
hidden_state_skip_layer=None,
return_texts=False,
):
batch_encoding = self.text2tokens(text)
return self.encode(
batch_encoding,
use_attention_mask=use_attention_mask,
output_hidden_states=output_hidden_states,
do_sample=do_sample,
hidden_state_skip_layer=hidden_state_skip_layer,
return_texts=return_texts,
)
@@ -1,15 +0,0 @@
import numpy as np
import math
def align_to(value, alignment):
"""align hight, width according to alignment
Args:
value (int): height or width
alignment (int): target alignment factor
Returns:
int: the aligned value
"""
return int(math.ceil(value / alignment) * alignment)
@@ -1,71 +0,0 @@
import os
from pathlib import Path
from einops import rearrange
import torch
import torchvision
import numpy as np
import imageio
CODE_SUFFIXES = {
".py", # Python codes
".sh", # Shell scripts
".yaml",
".yml", # Configuration files
}
def safe_dir(path):
"""
Create a directory (or the parent directory of a file) if it does not exist.
Args:
path (str or Path): Path to the directory.
Returns:
path (Path): Path object of the directory.
"""
path = Path(path)
path.mkdir(exist_ok=True, parents=True)
return path
def safe_file(path):
"""
Create the parent directory of a file if it does not exist.
Args:
path (str or Path): Path to the file.
Returns:
path (Path): Path object of the file.
"""
path = Path(path)
path.parent.mkdir(exist_ok=True, parents=True)
return path
def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=1, fps=24):
"""save videos by video tensor
copy from https://github.com/guoyww/AnimateDiff/blob/e92bd5671ba62c0d774a32951453e328018b7c5b/animatediff/utils/util.py#L61
Args:
videos (torch.Tensor): video tensor predicted by the model
path (str): path to save video
rescale (bool, optional): rescale the video tensor from [-1, 1] to . Defaults to False.
n_rows (int, optional): Defaults to 1.
fps (int, optional): video save fps. Defaults to 8.
"""
videos = rearrange(videos, "b c t h w -> t b c h w")
outputs = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=n_rows)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
if rescale:
x = (x + 1.0) / 2.0 # -1,1 -> 0,1
x = torch.clamp(x, 0, 1)
x = (x * 255).numpy().astype(np.uint8)
outputs.append(x)
os.makedirs(os.path.dirname(path), exist_ok=True)
imageio.mimsave(path, outputs, fps=fps)
-41
View File
@@ -1,41 +0,0 @@
import collections.abc
from itertools import repeat
def _ntuple(n):
def parse(x):
if isinstance(x, collections.abc.Iterable) and not isinstance(x, str):
x = tuple(x)
if len(x) == 1:
x = tuple(repeat(x[0], n))
return x
return tuple(repeat(x, n))
return parse
to_1tuple = _ntuple(1)
to_2tuple = _ntuple(2)
to_3tuple = _ntuple(3)
to_4tuple = _ntuple(4)
def as_tuple(x):
if isinstance(x, collections.abc.Iterable) and not isinstance(x, str):
return tuple(x)
if x is None or isinstance(x, (int, float, str)):
return (x,)
else:
raise ValueError(f"Unknown type {type(x)}")
def as_list_of_2tuple(x):
x = as_tuple(x)
if len(x) == 1:
x = (x[0], x[0])
assert len(x) % 2 == 0, f"Expect even length, got {len(x)}."
lst = []
for i in range(0, len(x), 2):
lst.append((x[i], x[i + 1]))
return lst
@@ -1,41 +0,0 @@
import argparse
import torch
from transformers import (
AutoProcessor,
LlavaForConditionalGeneration,
)
def preprocess_text_encoder_tokenizer(args):
processor = AutoProcessor.from_pretrained(args.input_dir)
model = LlavaForConditionalGeneration.from_pretrained(
args.input_dir, torch_dtype=torch.float16, low_cpu_mem_usage=True,
).to(0)
model.language_model.save_pretrained(f"{args.output_dir}")
processor.tokenizer.save_pretrained(f"{args.output_dir}")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument(
"--input_dir",
type=str,
required=True,
help="The path to the llava-llama-3-8b-v1_1-transformers.",
)
parser.add_argument(
"--output_dir",
type=str,
default="",
help="The output path of the llava-llama-3-8b-text-encoder-tokenizer."
"if '', the parent dir of output will be the same as input dir.",
)
args = parser.parse_args()
if len(args.output_dir) == 0:
args.output_dir = "/".join(args.input_dir.split("/")[:-1])
preprocess_text_encoder_tokenizer(args)
-66
View File
@@ -1,66 +0,0 @@
from pathlib import Path
import torch
from .autoencoder_kl_causal_3d import AutoencoderKLCausal3D
from ..constants import VAE_PATH, PRECISION_TO_TYPE
def load_vae(
vae_type: str = "884-16c-hy",
vae_precision: str = None,
sample_size: tuple = None,
vae_path: str = None,
logger=None,
device=None,
):
"""the fucntion to load the 3D VAE model
Args:
vae_type (str): the type of the 3D VAE model. Defaults to "884-16c-hy".
vae_precision (str, optional): the precision to load vae. Defaults to None.
sample_size (tuple, optional): the tiling size. Defaults to None.
vae_path (str, optional): the path to vae. Defaults to None.
logger (_type_, optional): logger. Defaults to None.
device (_type_, optional): device to load vae. Defaults to None.
"""
if vae_path is None:
vae_path = VAE_PATH[vae_type]
if logger is not None:
logger.info(f"Loading 3D VAE model ({vae_type}) from: {vae_path}")
config = AutoencoderKLCausal3D.load_config(vae_path)
if sample_size:
vae = AutoencoderKLCausal3D.from_config(config, sample_size=sample_size)
else:
vae = AutoencoderKLCausal3D.from_config(config)
vae_ckpt = Path(vae_path) / "pytorch_model.pt"
assert vae_ckpt.exists(), f"VAE checkpoint not found: {vae_ckpt}"
ckpt = torch.load(vae_ckpt, map_location=vae.device)
if "state_dict" in ckpt:
ckpt = ckpt["state_dict"]
if any(k.startswith("vae.") for k in ckpt.keys()):
ckpt = {
k.replace("vae.", ""): v for k, v in ckpt.items() if k.startswith("vae.")
}
vae.load_state_dict(ckpt)
spatial_compression_ratio = vae.config.spatial_compression_ratio
time_compression_ratio = vae.config.time_compression_ratio
if vae_precision is not None:
vae = vae.to(dtype=PRECISION_TO_TYPE[vae_precision])
vae.requires_grad_(False)
if logger is not None:
logger.info(f"VAE to dtype: {vae.dtype}")
if device is not None:
vae = vae.to(device)
vae.eval()
return vae, vae_path, spatial_compression_ratio, time_compression_ratio
@@ -1,684 +0,0 @@
# Copyright 2024 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
#
# Modified from diffusers==0.29.2
#
# ==============================================================================
from typing import Dict, Optional, Tuple, Union
from dataclasses import dataclass
import torch
import torch.nn as nn
from diffusers.configuration_utils import ConfigMixin, register_to_config
try:
# This diffusers is modified and packed in the mirror.
from diffusers.loaders import FromOriginalVAEMixin
except ImportError:
# Use this to be compatible with the original diffusers.
from diffusers.loaders.single_file_model import (
FromOriginalModelMixin as FromOriginalVAEMixin,
)
from diffusers.utils.accelerate_utils import apply_forward_hook
from diffusers.models.attention_processor import (
ADDED_KV_ATTENTION_PROCESSORS,
CROSS_ATTENTION_PROCESSORS,
Attention,
AttentionProcessor,
AttnAddedKVProcessor,
AttnProcessor,
)
from diffusers.models.modeling_outputs import AutoencoderKLOutput
from diffusers.models.modeling_utils import ModelMixin
from .vae import (
DecoderCausal3D,
BaseOutput,
DecoderOutput,
DiagonalGaussianDistribution,
EncoderCausal3D,
)
@dataclass
class DecoderOutput2(BaseOutput):
sample: torch.FloatTensor
posterior: Optional[DiagonalGaussianDistribution] = None
class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
r"""
A VAE model with KL loss for encoding images/videos into latents and decoding latent representations into images/videos.
This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented
for all models (such as downloading or saving).
"""
_supports_gradient_checkpointing = True
@register_to_config
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
down_block_types: Tuple[str] = ("DownEncoderBlockCausal3D",),
up_block_types: Tuple[str] = ("UpDecoderBlockCausal3D",),
block_out_channels: Tuple[int] = (64,),
layers_per_block: int = 1,
act_fn: str = "silu",
latent_channels: int = 4,
norm_num_groups: int = 32,
sample_size: int = 32,
sample_tsize: int = 64,
scaling_factor: float = 0.18215,
force_upcast: float = True,
spatial_compression_ratio: int = 8,
time_compression_ratio: int = 4,
mid_block_add_attention: bool = True,
):
super().__init__()
self.time_compression_ratio = time_compression_ratio
self.encoder = EncoderCausal3D(
in_channels=in_channels,
out_channels=latent_channels,
down_block_types=down_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
act_fn=act_fn,
norm_num_groups=norm_num_groups,
double_z=True,
time_compression_ratio=time_compression_ratio,
spatial_compression_ratio=spatial_compression_ratio,
mid_block_add_attention=mid_block_add_attention,
)
self.decoder = DecoderCausal3D(
in_channels=latent_channels,
out_channels=out_channels,
up_block_types=up_block_types,
block_out_channels=block_out_channels,
layers_per_block=layers_per_block,
norm_num_groups=norm_num_groups,
act_fn=act_fn,
time_compression_ratio=time_compression_ratio,
spatial_compression_ratio=spatial_compression_ratio,
mid_block_add_attention=mid_block_add_attention,
)
self.quant_conv = nn.Conv3d(
2 * latent_channels, 2 * latent_channels, kernel_size=1
)
self.post_quant_conv = nn.Conv3d(
latent_channels, latent_channels, kernel_size=1
)
self.use_slicing = False
self.use_spatial_tiling = False
self.use_temporal_tiling = False
# only relevant if vae tiling is enabled
self.tile_sample_min_tsize = sample_tsize
self.tile_latent_min_tsize = sample_tsize // time_compression_ratio
self.tile_sample_min_size = self.config.sample_size
sample_size = (
self.config.sample_size[0]
if isinstance(self.config.sample_size, (list, tuple))
else self.config.sample_size
)
self.tile_latent_min_size = int(
sample_size / (2 ** (len(self.config.block_out_channels) - 1))
)
self.tile_overlap_factor = 0.25
def _set_gradient_checkpointing(self, module, value=False):
if isinstance(module, (EncoderCausal3D, DecoderCausal3D)):
module.gradient_checkpointing = value
def enable_temporal_tiling(self, use_tiling: bool = True):
self.use_temporal_tiling = use_tiling
def disable_temporal_tiling(self):
self.enable_temporal_tiling(False)
def enable_spatial_tiling(self, use_tiling: bool = True):
self.use_spatial_tiling = use_tiling
def disable_spatial_tiling(self):
self.enable_spatial_tiling(False)
def enable_tiling(self, use_tiling: bool = True):
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 videos.
"""
self.enable_spatial_tiling(use_tiling)
self.enable_temporal_tiling(use_tiling)
def disable_tiling(self):
r"""
Disable tiled VAE decoding. If `enable_tiling` was previously enabled, this method will go back to computing
decoding in one step.
"""
self.disable_spatial_tiling()
self.disable_temporal_tiling()
def enable_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.use_slicing = True
def disable_slicing(self):
r"""
Disable sliced VAE decoding. If `enable_slicing` was previously enabled, this method will go back to computing
decoding in one step.
"""
self.use_slicing = False
@property
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.attn_processors
def attn_processors(self) -> Dict[str, AttentionProcessor]:
r"""
Returns:
`dict` of attention processors: A dictionary containing all attention processors used in the model with
indexed by its weight name.
"""
# set recursively
processors = {}
def fn_recursive_add_processors(
name: str,
module: torch.nn.Module,
processors: Dict[str, AttentionProcessor],
):
if hasattr(module, "get_processor"):
processors[f"{name}.processor"] = module.get_processor(
return_deprecated_lora=True
)
for sub_name, child in module.named_children():
fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
return processors
for name, module in self.named_children():
fn_recursive_add_processors(name, module, processors)
return processors
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.set_attn_processor
def set_attn_processor(
self,
processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]],
_remove_lora=False,
):
r"""
Sets the attention processor to use to compute attention.
Parameters:
processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
The instantiated processor class or a dictionary of processor classes that will be set as the processor
for **all** `Attention` layers.
If `processor` is a dict, the key needs to define the path to the corresponding cross attention
processor. This is strongly recommended when setting trainable attention processors.
"""
count = len(self.attn_processors.keys())
if isinstance(processor, dict) and len(processor) != count:
raise ValueError(
f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
)
def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor):
if hasattr(module, "set_processor"):
if not isinstance(processor, dict):
module.set_processor(processor, _remove_lora=_remove_lora)
else:
module.set_processor(
processor.pop(f"{name}.processor"), _remove_lora=_remove_lora
)
for sub_name, child in module.named_children():
fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
for name, module in self.named_children():
fn_recursive_attn_processor(name, module, processor)
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor
def set_default_attn_processor(self):
"""
Disables custom attention processors and sets the default attention implementation.
"""
if all(
proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS
for proc in self.attn_processors.values()
):
processor = AttnAddedKVProcessor()
elif all(
proc.__class__ in CROSS_ATTENTION_PROCESSORS
for proc in self.attn_processors.values()
):
processor = AttnProcessor()
else:
raise ValueError(
f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}"
)
self.set_attn_processor(processor, _remove_lora=True)
@apply_forward_hook
def encode(
self, x: torch.FloatTensor, return_dict: bool = True
) -> Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]:
"""
Encode a batch of images/videos into latents.
Args:
x (`torch.FloatTensor`): Input batch of images/videos.
return_dict (`bool`, *optional*, defaults to `True`):
Whether to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.
Returns:
The latent representations of the encoded images/videos. If `return_dict` is True, a
[`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain `tuple` is returned.
"""
assert len(x.shape) == 5, "The input tensor should have 5 dimensions."
if self.use_temporal_tiling and x.shape[2] > self.tile_sample_min_tsize:
return self.temporal_tiled_encode(x, return_dict=return_dict)
if self.use_spatial_tiling and (
x.shape[-1] > self.tile_sample_min_size
or x.shape[-2] > self.tile_sample_min_size
):
return self.spatial_tiled_encode(x, return_dict=return_dict)
if self.use_slicing and x.shape[0] > 1:
encoded_slices = [self.encoder(x_slice) for x_slice in x.split(1)]
h = torch.cat(encoded_slices)
else:
h = self.encoder(x)
moments = self.quant_conv(h)
posterior = DiagonalGaussianDistribution(moments)
if not return_dict:
return (posterior,)
return AutoencoderKLOutput(latent_dist=posterior)
def _decode(
self, z: torch.FloatTensor, return_dict: bool = True
) -> Union[DecoderOutput, torch.FloatTensor]:
assert len(z.shape) == 5, "The input tensor should have 5 dimensions."
if self.use_temporal_tiling and z.shape[2] > self.tile_latent_min_tsize:
return self.temporal_tiled_decode(z, return_dict=return_dict)
if self.use_spatial_tiling and (
z.shape[-1] > self.tile_latent_min_size
or z.shape[-2] > self.tile_latent_min_size
):
return self.spatial_tiled_decode(z, return_dict=return_dict)
z = self.post_quant_conv(z)
dec = self.decoder(z)
if not return_dict:
return (dec,)
return DecoderOutput(sample=dec)
@apply_forward_hook
def decode(
self, z: torch.FloatTensor, return_dict: bool = True, generator=None
) -> Union[DecoderOutput, torch.FloatTensor]:
"""
Decode a batch of images/videos.
Args:
z (`torch.FloatTensor`): Input batch of latent vectors.
return_dict (`bool`, *optional*, defaults to `True`):
Whether to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.
Returns:
[`~models.vae.DecoderOutput`] or `tuple`:
If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
returned.
"""
if self.use_slicing and z.shape[0] > 1:
decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)]
decoded = torch.cat(decoded_slices)
else:
decoded = self._decode(z).sample
if not return_dict:
return (decoded,)
return DecoderOutput(sample=decoded)
def blend_v(
self, a: torch.Tensor, b: torch.Tensor, blend_extent: int
) -> torch.Tensor:
blend_extent = min(a.shape[-2], b.shape[-2], blend_extent)
for y in range(blend_extent):
b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (
1 - y / blend_extent
) + b[:, :, :, y, :] * (y / blend_extent)
return b
def blend_h(
self, a: torch.Tensor, b: torch.Tensor, blend_extent: int
) -> torch.Tensor:
blend_extent = min(a.shape[-1], b.shape[-1], blend_extent)
for x in range(blend_extent):
b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (
1 - x / blend_extent
) + b[:, :, :, :, x] * (x / blend_extent)
return b
def blend_t(
self, a: torch.Tensor, b: torch.Tensor, blend_extent: int
) -> torch.Tensor:
blend_extent = min(a.shape[-3], b.shape[-3], blend_extent)
for x in range(blend_extent):
b[:, :, x, :, :] = a[:, :, -blend_extent + x, :, :] * (
1 - x / blend_extent
) + b[:, :, x, :, :] * (x / blend_extent)
return b
def spatial_tiled_encode(
self,
x: torch.FloatTensor,
return_dict: bool = True,
return_moments: bool = False,
) -> AutoencoderKLOutput:
r"""Encode a batch of images/videos using a tiled encoder.
When this option is enabled, the VAE will split the input tensor into tiles to compute encoding in several
steps. This is useful to keep memory use constant regardless of image/videos size. The end result of tiled encoding is
different from non-tiled encoding because each tile uses a different encoder. To avoid tiling artifacts, the
tiles overlap and are blended together to form a smooth output. You may still see tile-sized changes in the
output, but they should be much less noticeable.
Args:
x (`torch.FloatTensor`): Input batch of images/videos.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`~models.autoencoder_kl.AutoencoderKLOutput`] instead of a plain tuple.
Returns:
[`~models.autoencoder_kl.AutoencoderKLOutput`] or `tuple`:
If return_dict is True, a [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain
`tuple` is returned.
"""
overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
row_limit = self.tile_latent_min_size - blend_extent
# Split video into tiles and encode them separately.
rows = []
for i in range(0, x.shape[-2], overlap_size):
row = []
for j in range(0, x.shape[-1], overlap_size):
tile = x[
:,
:,
:,
i : i + self.tile_sample_min_size,
j : j + self.tile_sample_min_size,
]
tile = self.encoder(tile)
tile = self.quant_conv(tile)
row.append(tile)
rows.append(row)
result_rows = []
for i, row in enumerate(rows):
result_row = []
for j, tile in enumerate(row):
# blend the above tile and the left tile
# to the current tile and add the current tile to the result row
if i > 0:
tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
if j > 0:
tile = self.blend_h(row[j - 1], tile, blend_extent)
result_row.append(tile[:, :, :, :row_limit, :row_limit])
result_rows.append(torch.cat(result_row, dim=-1))
moments = torch.cat(result_rows, dim=-2)
if return_moments:
return moments
posterior = DiagonalGaussianDistribution(moments)
if not return_dict:
return (posterior,)
return AutoencoderKLOutput(latent_dist=posterior)
def spatial_tiled_decode(
self, z: torch.FloatTensor, return_dict: bool = True
) -> Union[DecoderOutput, torch.FloatTensor]:
r"""
Decode a batch of images/videos using a tiled decoder.
Args:
z (`torch.FloatTensor`): Input batch of latent vectors.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`~models.vae.DecoderOutput`] instead of a plain tuple.
Returns:
[`~models.vae.DecoderOutput`] or `tuple`:
If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
returned.
"""
overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor)
row_limit = self.tile_sample_min_size - blend_extent
# Split z into overlapping tiles and decode them separately.
# The tiles have an overlap to avoid seams between tiles.
rows = []
for i in range(0, z.shape[-2], overlap_size):
row = []
for j in range(0, z.shape[-1], overlap_size):
tile = z[
:,
:,
:,
i : i + self.tile_latent_min_size,
j : j + self.tile_latent_min_size,
]
tile = self.post_quant_conv(tile)
decoded = self.decoder(tile)
row.append(decoded)
rows.append(row)
result_rows = []
for i, row in enumerate(rows):
result_row = []
for j, tile in enumerate(row):
# blend the above tile and the left tile
# to the current tile and add the current tile to the result row
if i > 0:
tile = self.blend_v(rows[i - 1][j], tile, blend_extent)
if j > 0:
tile = self.blend_h(row[j - 1], tile, blend_extent)
result_row.append(tile[:, :, :, :row_limit, :row_limit])
result_rows.append(torch.cat(result_row, dim=-1))
dec = torch.cat(result_rows, dim=-2)
if not return_dict:
return (dec,)
return DecoderOutput(sample=dec)
def temporal_tiled_encode(
self, x: torch.FloatTensor, return_dict: bool = True
) -> AutoencoderKLOutput:
B, C, T, H, W = x.shape
overlap_size = int(self.tile_sample_min_tsize * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_latent_min_tsize * self.tile_overlap_factor)
t_limit = self.tile_latent_min_tsize - blend_extent
# Split the video into tiles and encode them separately.
row = []
for i in range(0, T, overlap_size):
tile = x[:, :, i : i + self.tile_sample_min_tsize + 1, :, :]
if self.use_spatial_tiling and (
tile.shape[-1] > self.tile_sample_min_size
or tile.shape[-2] > self.tile_sample_min_size
):
tile = self.spatial_tiled_encode(tile, return_moments=True)
else:
tile = self.encoder(tile)
tile = self.quant_conv(tile)
if i > 0:
tile = tile[:, :, 1:, :, :]
row.append(tile)
result_row = []
for i, tile in enumerate(row):
if i > 0:
tile = self.blend_t(row[i - 1], tile, blend_extent)
result_row.append(tile[:, :, :t_limit, :, :])
else:
result_row.append(tile[:, :, : t_limit + 1, :, :])
moments = torch.cat(result_row, dim=2)
posterior = DiagonalGaussianDistribution(moments)
if not return_dict:
return (posterior,)
return AutoencoderKLOutput(latent_dist=posterior)
def temporal_tiled_decode(
self, z: torch.FloatTensor, return_dict: bool = True
) -> Union[DecoderOutput, torch.FloatTensor]:
# Split z into overlapping tiles and decode them separately.
B, C, T, H, W = z.shape
overlap_size = int(self.tile_latent_min_tsize * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_sample_min_tsize * self.tile_overlap_factor)
t_limit = self.tile_sample_min_tsize - blend_extent
row = []
for i in range(0, T, overlap_size):
tile = z[:, :, i : i + self.tile_latent_min_tsize + 1, :, :]
if self.use_spatial_tiling and (
tile.shape[-1] > self.tile_latent_min_size
or tile.shape[-2] > self.tile_latent_min_size
):
decoded = self.spatial_tiled_decode(tile, return_dict=True).sample
else:
tile = self.post_quant_conv(tile)
decoded = self.decoder(tile)
if i > 0:
decoded = decoded[:, :, 1:, :, :]
row.append(decoded)
result_row = []
for i, tile in enumerate(row):
if i > 0:
tile = self.blend_t(row[i - 1], tile, blend_extent)
result_row.append(tile[:, :, :t_limit, :, :])
else:
result_row.append(tile[:, :, : t_limit + 1, :, :])
dec = torch.cat(result_row, dim=2)
if not return_dict:
return (dec,)
return DecoderOutput(sample=dec)
def forward(
self,
sample: torch.FloatTensor,
sample_posterior: bool = False,
return_dict: bool = True,
return_posterior: bool = False,
generator: Optional[torch.Generator] = None,
) -> Union[DecoderOutput2, torch.FloatTensor]:
r"""
Args:
sample (`torch.FloatTensor`): Input sample.
sample_posterior (`bool`, *optional*, defaults to `False`):
Whether to sample from the posterior.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`DecoderOutput`] instead of a plain tuple.
"""
x = sample
posterior = self.encode(x).latent_dist
if sample_posterior:
z = posterior.sample(generator=generator)
else:
z = posterior.mode()
dec = self.decode(z).sample
if not return_dict:
if return_posterior:
return (dec, posterior)
else:
return (dec,)
if return_posterior:
return DecoderOutput2(sample=dec, posterior=posterior)
else:
return DecoderOutput2(sample=dec)
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.fuse_qkv_projections
def fuse_qkv_projections(self):
"""
Enables fused QKV projections. For self-attention modules, all projection matrices (i.e., query,
key, value) are fused. For cross-attention modules, key and value projection matrices are fused.
<Tip warning={true}>
This API is 🧪 experimental.
</Tip>
"""
self.original_attn_processors = None
for _, attn_processor in self.attn_processors.items():
if "Added" in str(attn_processor.__class__.__name__):
raise ValueError(
"`fuse_qkv_projections()` is not supported for models having added KV projections."
)
self.original_attn_processors = self.attn_processors
for module in self.modules():
if isinstance(module, Attention):
module.fuse_projections(fuse=True)
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.unfuse_qkv_projections
def unfuse_qkv_projections(self):
"""Disables the fused QKV projection if enabled.
<Tip warning={true}>
This API is 🧪 experimental.
</Tip>
"""
if self.original_attn_processors is not None:
self.set_attn_processor(self.original_attn_processors)
@@ -1,823 +0,0 @@
# Copyright 2024 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
#
# Modified from diffusers==0.29.2
#
# ==============================================================================
from typing import Optional, Tuple, Union
import torch
import torch.nn.functional as F
from torch import nn
from einops import rearrange
from diffusers.utils import logging
from diffusers.models.activations import get_activation
from diffusers.models.attention_processor import SpatialNorm
from diffusers.models.attention_processor import Attention
from diffusers.models.normalization import AdaGroupNorm
from diffusers.models.normalization import RMSNorm
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
def prepare_causal_attention_mask(
n_frame: int, n_hw: int, dtype, device, batch_size: int = None
):
seq_len = n_frame * n_hw
mask = torch.full((seq_len, seq_len), float("-inf"), dtype=dtype, device=device)
for i in range(seq_len):
i_frame = i // n_hw
mask[i, : (i_frame + 1) * n_hw] = 0
if batch_size is not None:
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
return mask
class CausalConv3d(nn.Module):
"""
Implements a causal 3D convolution layer where each position only depends on previous timesteps and current spatial locations.
This maintains temporal causality in video generation tasks.
"""
def __init__(
self,
chan_in,
chan_out,
kernel_size: Union[int, Tuple[int, int, int]],
stride: Union[int, Tuple[int, int, int]] = 1,
dilation: Union[int, Tuple[int, int, int]] = 1,
pad_mode="replicate",
**kwargs,
):
super().__init__()
self.pad_mode = pad_mode
padding = (
kernel_size // 2,
kernel_size // 2,
kernel_size // 2,
kernel_size // 2,
kernel_size - 1,
0,
) # W, H, T
self.time_causal_padding = padding
self.conv = nn.Conv3d(
chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs
)
def forward(self, x):
x = F.pad(x, self.time_causal_padding, mode=self.pad_mode)
return self.conv(x)
class UpsampleCausal3D(nn.Module):
"""
A 3D upsampling layer with an optional convolution.
"""
def __init__(
self,
channels: int,
use_conv: bool = False,
use_conv_transpose: bool = False,
out_channels: Optional[int] = None,
name: str = "conv",
kernel_size: Optional[int] = None,
padding=1,
norm_type=None,
eps=None,
elementwise_affine=None,
bias=True,
interpolate=True,
upsample_factor=(2, 2, 2),
):
super().__init__()
self.channels = channels
self.out_channels = out_channels or channels
self.use_conv = use_conv
self.use_conv_transpose = use_conv_transpose
self.name = name
self.interpolate = interpolate
self.upsample_factor = upsample_factor
if norm_type == "ln_norm":
self.norm = nn.LayerNorm(channels, eps, elementwise_affine)
elif norm_type == "rms_norm":
self.norm = RMSNorm(channels, eps, elementwise_affine)
elif norm_type is None:
self.norm = None
else:
raise ValueError(f"unknown norm_type: {norm_type}")
conv = None
if use_conv_transpose:
raise NotImplementedError
elif use_conv:
if kernel_size is None:
kernel_size = 3
conv = CausalConv3d(
self.channels, self.out_channels, kernel_size=kernel_size, bias=bias
)
if name == "conv":
self.conv = conv
else:
self.Conv2d_0 = conv
def forward(
self,
hidden_states: torch.FloatTensor,
output_size: Optional[int] = None,
scale: float = 1.0,
) -> torch.FloatTensor:
assert hidden_states.shape[1] == self.channels
if self.norm is not None:
raise NotImplementedError
if self.use_conv_transpose:
return self.conv(hidden_states)
# Cast to float32 to as 'upsample_nearest2d_out_frame' op does not support bfloat16
dtype = hidden_states.dtype
if dtype == torch.bfloat16:
hidden_states = hidden_states.to(torch.float32)
# upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984
if hidden_states.shape[0] >= 64:
hidden_states = hidden_states.contiguous()
# if `output_size` is passed we force the interpolation output
# size and do not make use of `scale_factor=2`
if self.interpolate:
B, C, T, H, W = hidden_states.shape
first_h, other_h = hidden_states.split((1, T - 1), dim=2)
if output_size is None:
if T > 1:
other_h = F.interpolate(
other_h, scale_factor=self.upsample_factor, mode="nearest"
)
first_h = first_h.squeeze(2)
first_h = F.interpolate(
first_h, scale_factor=self.upsample_factor[1:], mode="nearest"
)
first_h = first_h.unsqueeze(2)
else:
raise NotImplementedError
if T > 1:
hidden_states = torch.cat((first_h, other_h), dim=2)
else:
hidden_states = first_h
# If the input is bfloat16, we cast back to bfloat16
if dtype == torch.bfloat16:
hidden_states = hidden_states.to(dtype)
if self.use_conv:
if self.name == "conv":
hidden_states = self.conv(hidden_states)
else:
hidden_states = self.Conv2d_0(hidden_states)
return hidden_states
class DownsampleCausal3D(nn.Module):
"""
A 3D downsampling layer with an optional convolution.
"""
def __init__(
self,
channels: int,
use_conv: bool = False,
out_channels: Optional[int] = None,
padding: int = 1,
name: str = "conv",
kernel_size=3,
norm_type=None,
eps=None,
elementwise_affine=None,
bias=True,
stride=2,
):
super().__init__()
self.channels = channels
self.out_channels = out_channels or channels
self.use_conv = use_conv
self.padding = padding
stride = stride
self.name = name
if norm_type == "ln_norm":
self.norm = nn.LayerNorm(channels, eps, elementwise_affine)
elif norm_type == "rms_norm":
self.norm = RMSNorm(channels, eps, elementwise_affine)
elif norm_type is None:
self.norm = None
else:
raise ValueError(f"unknown norm_type: {norm_type}")
if use_conv:
conv = CausalConv3d(
self.channels,
self.out_channels,
kernel_size=kernel_size,
stride=stride,
bias=bias,
)
else:
raise NotImplementedError
if name == "conv":
self.Conv2d_0 = conv
self.conv = conv
elif name == "Conv2d_0":
self.conv = conv
else:
self.conv = conv
def forward(
self, hidden_states: torch.FloatTensor, scale: float = 1.0
) -> torch.FloatTensor:
assert hidden_states.shape[1] == self.channels
if self.norm is not None:
hidden_states = self.norm(hidden_states.permute(0, 2, 3, 1)).permute(
0, 3, 1, 2
)
assert hidden_states.shape[1] == self.channels
hidden_states = self.conv(hidden_states)
return hidden_states
class ResnetBlockCausal3D(nn.Module):
r"""
A Resnet block.
"""
def __init__(
self,
*,
in_channels: int,
out_channels: Optional[int] = None,
conv_shortcut: bool = False,
dropout: float = 0.0,
temb_channels: int = 512,
groups: int = 32,
groups_out: Optional[int] = None,
pre_norm: bool = True,
eps: float = 1e-6,
non_linearity: str = "swish",
skip_time_act: bool = False,
# default, scale_shift, ada_group, spatial
time_embedding_norm: str = "default",
kernel: Optional[torch.FloatTensor] = None,
output_scale_factor: float = 1.0,
use_in_shortcut: Optional[bool] = None,
up: bool = False,
down: bool = False,
conv_shortcut_bias: bool = True,
conv_3d_out_channels: Optional[int] = None,
):
super().__init__()
self.pre_norm = pre_norm
self.pre_norm = True
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.use_conv_shortcut = conv_shortcut
self.up = up
self.down = down
self.output_scale_factor = output_scale_factor
self.time_embedding_norm = time_embedding_norm
self.skip_time_act = skip_time_act
linear_cls = nn.Linear
if groups_out is None:
groups_out = groups
if self.time_embedding_norm == "ada_group":
self.norm1 = AdaGroupNorm(temb_channels, in_channels, groups, eps=eps)
elif self.time_embedding_norm == "spatial":
self.norm1 = SpatialNorm(in_channels, temb_channels)
else:
self.norm1 = torch.nn.GroupNorm(
num_groups=groups, num_channels=in_channels, eps=eps, affine=True
)
self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3, stride=1)
if temb_channels is not None:
if self.time_embedding_norm == "default":
self.time_emb_proj = linear_cls(temb_channels, out_channels)
elif self.time_embedding_norm == "scale_shift":
self.time_emb_proj = linear_cls(temb_channels, 2 * out_channels)
elif (
self.time_embedding_norm == "ada_group"
or self.time_embedding_norm == "spatial"
):
self.time_emb_proj = None
else:
raise ValueError(
f"Unknown time_embedding_norm : {self.time_embedding_norm} "
)
else:
self.time_emb_proj = None
if self.time_embedding_norm == "ada_group":
self.norm2 = AdaGroupNorm(temb_channels, out_channels, groups_out, eps=eps)
elif self.time_embedding_norm == "spatial":
self.norm2 = SpatialNorm(out_channels, temb_channels)
else:
self.norm2 = torch.nn.GroupNorm(
num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True
)
self.dropout = torch.nn.Dropout(dropout)
conv_3d_out_channels = conv_3d_out_channels or out_channels
self.conv2 = CausalConv3d(
out_channels, conv_3d_out_channels, kernel_size=3, stride=1
)
self.nonlinearity = get_activation(non_linearity)
self.upsample = self.downsample = None
if self.up:
self.upsample = UpsampleCausal3D(in_channels, use_conv=False)
elif self.down:
self.downsample = DownsampleCausal3D(in_channels, use_conv=False, name="op")
self.use_in_shortcut = (
self.in_channels != conv_3d_out_channels
if use_in_shortcut is None
else use_in_shortcut
)
self.conv_shortcut = None
if self.use_in_shortcut:
self.conv_shortcut = CausalConv3d(
in_channels,
conv_3d_out_channels,
kernel_size=1,
stride=1,
bias=conv_shortcut_bias,
)
def forward(
self,
input_tensor: torch.FloatTensor,
temb: torch.FloatTensor,
scale: float = 1.0,
) -> torch.FloatTensor:
hidden_states = input_tensor
if (
self.time_embedding_norm == "ada_group"
or self.time_embedding_norm == "spatial"
):
hidden_states = self.norm1(hidden_states, temb)
else:
hidden_states = self.norm1(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
if self.upsample is not None:
# upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984
if hidden_states.shape[0] >= 64:
input_tensor = input_tensor.contiguous()
hidden_states = hidden_states.contiguous()
input_tensor = self.upsample(input_tensor, scale=scale)
hidden_states = self.upsample(hidden_states, scale=scale)
elif self.downsample is not None:
input_tensor = self.downsample(input_tensor, scale=scale)
hidden_states = self.downsample(hidden_states, scale=scale)
hidden_states = self.conv1(hidden_states)
if self.time_emb_proj is not None:
if not self.skip_time_act:
temb = self.nonlinearity(temb)
temb = self.time_emb_proj(temb, scale)[:, :, None, None]
if temb is not None and self.time_embedding_norm == "default":
hidden_states = hidden_states + temb
if (
self.time_embedding_norm == "ada_group"
or self.time_embedding_norm == "spatial"
):
hidden_states = self.norm2(hidden_states, temb)
else:
hidden_states = self.norm2(hidden_states)
if temb is not None and self.time_embedding_norm == "scale_shift":
scale, shift = torch.chunk(temb, 2, dim=1)
hidden_states = hidden_states * (1 + scale) + shift
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.dropout(hidden_states)
hidden_states = self.conv2(hidden_states)
if self.conv_shortcut is not None:
input_tensor = self.conv_shortcut(input_tensor)
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
return output_tensor
def get_down_block3d(
down_block_type: str,
num_layers: int,
in_channels: int,
out_channels: int,
temb_channels: int,
add_downsample: bool,
downsample_stride: int,
resnet_eps: float,
resnet_act_fn: str,
transformer_layers_per_block: int = 1,
num_attention_heads: Optional[int] = None,
resnet_groups: Optional[int] = None,
cross_attention_dim: Optional[int] = None,
downsample_padding: Optional[int] = None,
dual_cross_attention: bool = False,
use_linear_projection: bool = False,
only_cross_attention: bool = False,
upcast_attention: bool = False,
resnet_time_scale_shift: str = "default",
attention_type: str = "default",
resnet_skip_time_act: bool = False,
resnet_out_scale_factor: float = 1.0,
cross_attention_norm: Optional[str] = None,
attention_head_dim: Optional[int] = None,
downsample_type: Optional[str] = None,
dropout: float = 0.0,
):
# If attn head dim is not defined, we default it to the number of heads
if attention_head_dim is None:
logger.warn(
f"It is recommended to provide `attention_head_dim` when calling `get_down_block`. Defaulting `attention_head_dim` to {num_attention_heads}."
)
attention_head_dim = num_attention_heads
down_block_type = (
down_block_type[7:]
if down_block_type.startswith("UNetRes")
else down_block_type
)
if down_block_type == "DownEncoderBlockCausal3D":
return DownEncoderBlockCausal3D(
num_layers=num_layers,
in_channels=in_channels,
out_channels=out_channels,
dropout=dropout,
add_downsample=add_downsample,
downsample_stride=downsample_stride,
resnet_eps=resnet_eps,
resnet_act_fn=resnet_act_fn,
resnet_groups=resnet_groups,
downsample_padding=downsample_padding,
resnet_time_scale_shift=resnet_time_scale_shift,
)
raise ValueError(f"{down_block_type} does not exist.")
def get_up_block3d(
up_block_type: str,
num_layers: int,
in_channels: int,
out_channels: int,
prev_output_channel: int,
temb_channels: int,
add_upsample: bool,
upsample_scale_factor: Tuple,
resnet_eps: float,
resnet_act_fn: str,
resolution_idx: Optional[int] = None,
transformer_layers_per_block: int = 1,
num_attention_heads: Optional[int] = None,
resnet_groups: Optional[int] = None,
cross_attention_dim: Optional[int] = None,
dual_cross_attention: bool = False,
use_linear_projection: bool = False,
only_cross_attention: bool = False,
upcast_attention: bool = False,
resnet_time_scale_shift: str = "default",
attention_type: str = "default",
resnet_skip_time_act: bool = False,
resnet_out_scale_factor: float = 1.0,
cross_attention_norm: Optional[str] = None,
attention_head_dim: Optional[int] = None,
upsample_type: Optional[str] = None,
dropout: float = 0.0,
) -> nn.Module:
# If attn head dim is not defined, we default it to the number of heads
if attention_head_dim is None:
logger.warn(
f"It is recommended to provide `attention_head_dim` when calling `get_up_block`. Defaulting `attention_head_dim` to {num_attention_heads}."
)
attention_head_dim = num_attention_heads
up_block_type = (
up_block_type[7:] if up_block_type.startswith("UNetRes") else up_block_type
)
if up_block_type == "UpDecoderBlockCausal3D":
return UpDecoderBlockCausal3D(
num_layers=num_layers,
in_channels=in_channels,
out_channels=out_channels,
resolution_idx=resolution_idx,
dropout=dropout,
add_upsample=add_upsample,
upsample_scale_factor=upsample_scale_factor,
resnet_eps=resnet_eps,
resnet_act_fn=resnet_act_fn,
resnet_groups=resnet_groups,
resnet_time_scale_shift=resnet_time_scale_shift,
temb_channels=temb_channels,
)
raise ValueError(f"{up_block_type} does not exist.")
class UNetMidBlockCausal3D(nn.Module):
"""
A 3D UNet mid-block [`UNetMidBlockCausal3D`] with multiple residual blocks and optional attention blocks.
"""
def __init__(
self,
in_channels: int,
temb_channels: int,
dropout: float = 0.0,
num_layers: int = 1,
resnet_eps: float = 1e-6,
resnet_time_scale_shift: str = "default", # default, spatial
resnet_act_fn: str = "swish",
resnet_groups: int = 32,
attn_groups: Optional[int] = None,
resnet_pre_norm: bool = True,
add_attention: bool = True,
attention_head_dim: int = 1,
output_scale_factor: float = 1.0,
):
super().__init__()
resnet_groups = (
resnet_groups if resnet_groups is not None else min(in_channels // 4, 32)
)
self.add_attention = add_attention
if attn_groups is None:
attn_groups = (
resnet_groups if resnet_time_scale_shift == "default" else None
)
# there is always at least one resnet
resnets = [
ResnetBlockCausal3D(
in_channels=in_channels,
out_channels=in_channels,
temb_channels=temb_channels,
eps=resnet_eps,
groups=resnet_groups,
dropout=dropout,
time_embedding_norm=resnet_time_scale_shift,
non_linearity=resnet_act_fn,
output_scale_factor=output_scale_factor,
pre_norm=resnet_pre_norm,
)
]
attentions = []
if attention_head_dim is None:
logger.warn(
f"It is not recommend to pass `attention_head_dim=None`. Defaulting `attention_head_dim` to `in_channels`: {in_channels}."
)
attention_head_dim = in_channels
for _ in range(num_layers):
if self.add_attention:
attentions.append(
Attention(
in_channels,
heads=in_channels // attention_head_dim,
dim_head=attention_head_dim,
rescale_output_factor=output_scale_factor,
eps=resnet_eps,
norm_num_groups=attn_groups,
spatial_norm_dim=(
temb_channels
if resnet_time_scale_shift == "spatial"
else None
),
residual_connection=True,
bias=True,
upcast_softmax=True,
_from_deprecated_attn_block=True,
)
)
else:
attentions.append(None)
resnets.append(
ResnetBlockCausal3D(
in_channels=in_channels,
out_channels=in_channels,
temb_channels=temb_channels,
eps=resnet_eps,
groups=resnet_groups,
dropout=dropout,
time_embedding_norm=resnet_time_scale_shift,
non_linearity=resnet_act_fn,
output_scale_factor=output_scale_factor,
pre_norm=resnet_pre_norm,
)
)
self.attentions = nn.ModuleList(attentions)
self.resnets = nn.ModuleList(resnets)
def forward(
self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None
) -> torch.FloatTensor:
hidden_states = self.resnets[0](hidden_states, temb)
for attn, resnet in zip(self.attentions, self.resnets[1:]):
if attn is not None:
B, C, T, H, W = hidden_states.shape
hidden_states = rearrange(hidden_states, "b c f h w -> b (f h w) c")
attention_mask = prepare_causal_attention_mask(
T, H * W, hidden_states.dtype, hidden_states.device, batch_size=B
)
hidden_states = attn(
hidden_states, temb=temb, attention_mask=attention_mask
)
hidden_states = rearrange(
hidden_states, "b (f h w) c -> b c f h w", f=T, h=H, w=W
)
hidden_states = resnet(hidden_states, temb)
return hidden_states
class DownEncoderBlockCausal3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
dropout: float = 0.0,
num_layers: int = 1,
resnet_eps: float = 1e-6,
resnet_time_scale_shift: str = "default",
resnet_act_fn: str = "swish",
resnet_groups: int = 32,
resnet_pre_norm: bool = True,
output_scale_factor: float = 1.0,
add_downsample: bool = True,
downsample_stride: int = 2,
downsample_padding: int = 1,
):
super().__init__()
resnets = []
for i in range(num_layers):
in_channels = in_channels if i == 0 else out_channels
resnets.append(
ResnetBlockCausal3D(
in_channels=in_channels,
out_channels=out_channels,
temb_channels=None,
eps=resnet_eps,
groups=resnet_groups,
dropout=dropout,
time_embedding_norm=resnet_time_scale_shift,
non_linearity=resnet_act_fn,
output_scale_factor=output_scale_factor,
pre_norm=resnet_pre_norm,
)
)
self.resnets = nn.ModuleList(resnets)
if add_downsample:
self.downsamplers = nn.ModuleList(
[
DownsampleCausal3D(
out_channels,
use_conv=True,
out_channels=out_channels,
padding=downsample_padding,
name="op",
stride=downsample_stride,
)
]
)
else:
self.downsamplers = None
def forward(
self, hidden_states: torch.FloatTensor, scale: float = 1.0
) -> torch.FloatTensor:
for resnet in self.resnets:
hidden_states = resnet(hidden_states, temb=None, scale=scale)
if self.downsamplers is not None:
for downsampler in self.downsamplers:
hidden_states = downsampler(hidden_states, scale)
return hidden_states
class UpDecoderBlockCausal3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
resolution_idx: Optional[int] = None,
dropout: float = 0.0,
num_layers: int = 1,
resnet_eps: float = 1e-6,
resnet_time_scale_shift: str = "default", # default, spatial
resnet_act_fn: str = "swish",
resnet_groups: int = 32,
resnet_pre_norm: bool = True,
output_scale_factor: float = 1.0,
add_upsample: bool = True,
upsample_scale_factor=(2, 2, 2),
temb_channels: Optional[int] = None,
):
super().__init__()
resnets = []
for i in range(num_layers):
input_channels = in_channels if i == 0 else out_channels
resnets.append(
ResnetBlockCausal3D(
in_channels=input_channels,
out_channels=out_channels,
temb_channels=temb_channels,
eps=resnet_eps,
groups=resnet_groups,
dropout=dropout,
time_embedding_norm=resnet_time_scale_shift,
non_linearity=resnet_act_fn,
output_scale_factor=output_scale_factor,
pre_norm=resnet_pre_norm,
)
)
self.resnets = nn.ModuleList(resnets)
if add_upsample:
self.upsamplers = nn.ModuleList(
[
UpsampleCausal3D(
out_channels,
use_conv=True,
out_channels=out_channels,
upsample_factor=upsample_scale_factor,
)
]
)
else:
self.upsamplers = None
self.resolution_idx = resolution_idx
def forward(
self,
hidden_states: torch.FloatTensor,
temb: Optional[torch.FloatTensor] = None,
scale: float = 1.0,
) -> torch.FloatTensor:
for resnet in self.resnets:
hidden_states = resnet(hidden_states, temb=temb, scale=scale)
if self.upsamplers is not None:
for upsampler in self.upsamplers:
hidden_states = upsampler(hidden_states)
return hidden_states
-374
View File
@@ -1,374 +0,0 @@
from dataclasses import dataclass
from typing import Optional, Tuple
import numpy as np
import torch
import torch.nn as nn
from diffusers.utils import BaseOutput, is_torch_version
from diffusers.utils.torch_utils import randn_tensor
from diffusers.models.attention_processor import SpatialNorm
from .unet_causal_3d_blocks import (
CausalConv3d,
UNetMidBlockCausal3D,
get_down_block3d,
get_up_block3d,
)
@dataclass
class DecoderOutput(BaseOutput):
r"""
Output of decoding method.
Args:
sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`):
The decoded output sample from the last layer of the model.
"""
sample: torch.FloatTensor
class EncoderCausal3D(nn.Module):
r"""
The `EncoderCausal3D` layer of a variational autoencoder that encodes its input into a latent representation.
"""
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
down_block_types: Tuple[str, ...] = ("DownEncoderBlockCausal3D",),
block_out_channels: Tuple[int, ...] = (64,),
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
double_z: bool = True,
mid_block_add_attention=True,
time_compression_ratio: int = 4,
spatial_compression_ratio: int = 8,
):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(
in_channels, block_out_channels[0], kernel_size=3, stride=1
)
self.mid_block = None
self.down_blocks = nn.ModuleList([])
# down
output_channel = block_out_channels[0]
for i, down_block_type in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_downsample_layers = int(np.log2(spatial_compression_ratio))
num_time_downsample_layers = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial_downsample = bool(i < num_spatial_downsample_layers)
add_time_downsample = bool(
i >= (len(block_out_channels) - 1 - num_time_downsample_layers)
and not is_final_block
)
else:
raise ValueError(
f"Unsupported time_compression_ratio: {time_compression_ratio}."
)
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
downsample_stride_T = (2,) if add_time_downsample else (1,)
downsample_stride = tuple(downsample_stride_T + downsample_stride_HW)
down_block = get_down_block3d(
down_block_type,
num_layers=self.layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=bool(add_spatial_downsample or add_time_downsample),
downsample_stride=downsample_stride,
resnet_eps=1e-6,
downsample_padding=0,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
attention_head_dim=output_channel,
temb_channels=None,
)
self.down_blocks.append(down_block)
# mid
self.mid_block = UNetMidBlockCausal3D(
in_channels=block_out_channels[-1],
resnet_eps=1e-6,
resnet_act_fn=act_fn,
output_scale_factor=1,
resnet_time_scale_shift="default",
attention_head_dim=block_out_channels[-1],
resnet_groups=norm_num_groups,
temb_channels=None,
add_attention=mid_block_add_attention,
)
# out
self.conv_norm_out = nn.GroupNorm(
num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6
)
self.conv_act = nn.SiLU()
conv_out_channels = 2 * out_channels if double_z else out_channels
self.conv_out = CausalConv3d(
block_out_channels[-1], conv_out_channels, kernel_size=3
)
def forward(self, sample: torch.FloatTensor) -> torch.FloatTensor:
r"""The forward method of the `EncoderCausal3D` class."""
assert len(sample.shape) == 5, "The input tensor should have 5 dimensions"
sample = self.conv_in(sample)
# down
for down_block in self.down_blocks:
sample = down_block(sample)
# middle
sample = self.mid_block(sample)
# post-process
sample = self.conv_norm_out(sample)
sample = self.conv_act(sample)
sample = self.conv_out(sample)
return sample
class DecoderCausal3D(nn.Module):
r"""
The `DecoderCausal3D` layer of a variational autoencoder that decodes its latent representation into an output sample.
"""
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
up_block_types: Tuple[str, ...] = ("UpDecoderBlockCausal3D",),
block_out_channels: Tuple[int, ...] = (64,),
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
norm_type: str = "group", # group, spatial
mid_block_add_attention=True,
time_compression_ratio: int = 4,
spatial_compression_ratio: int = 8,
):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(
in_channels, block_out_channels[-1], kernel_size=3, stride=1
)
self.mid_block = None
self.up_blocks = nn.ModuleList([])
temb_channels = in_channels if norm_type == "spatial" else None
# mid
self.mid_block = UNetMidBlockCausal3D(
in_channels=block_out_channels[-1],
resnet_eps=1e-6,
resnet_act_fn=act_fn,
output_scale_factor=1,
resnet_time_scale_shift="default" if norm_type == "group" else norm_type,
attention_head_dim=block_out_channels[-1],
resnet_groups=norm_num_groups,
temb_channels=temb_channels,
add_attention=mid_block_add_attention,
)
# up
reversed_block_out_channels = list(reversed(block_out_channels))
output_channel = reversed_block_out_channels[0]
for i, up_block_type in enumerate(up_block_types):
prev_output_channel = output_channel
output_channel = reversed_block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_upsample_layers = int(np.log2(spatial_compression_ratio))
num_time_upsample_layers = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial_upsample = bool(i < num_spatial_upsample_layers)
add_time_upsample = bool(
i >= len(block_out_channels) - 1 - num_time_upsample_layers
and not is_final_block
)
else:
raise ValueError(
f"Unsupported time_compression_ratio: {time_compression_ratio}."
)
upsample_scale_factor_HW = (2, 2) if add_spatial_upsample else (1, 1)
upsample_scale_factor_T = (2,) if add_time_upsample else (1,)
upsample_scale_factor = tuple(
upsample_scale_factor_T + upsample_scale_factor_HW
)
up_block = get_up_block3d(
up_block_type,
num_layers=self.layers_per_block + 1,
in_channels=prev_output_channel,
out_channels=output_channel,
prev_output_channel=None,
add_upsample=bool(add_spatial_upsample or add_time_upsample),
upsample_scale_factor=upsample_scale_factor,
resnet_eps=1e-6,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
attention_head_dim=output_channel,
temb_channels=temb_channels,
resnet_time_scale_shift=norm_type,
)
self.up_blocks.append(up_block)
prev_output_channel = output_channel
# out
if norm_type == "spatial":
self.conv_norm_out = SpatialNorm(block_out_channels[0], temb_channels)
else:
self.conv_norm_out = nn.GroupNorm(
num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6
)
self.conv_act = nn.SiLU()
self.conv_out = CausalConv3d(block_out_channels[0], out_channels, kernel_size=3)
self.gradient_checkpointing = False
def forward(
self,
sample: torch.FloatTensor,
latent_embeds: Optional[torch.FloatTensor] = None,
) -> torch.FloatTensor:
r"""The forward method of the `DecoderCausal3D` class."""
assert len(sample.shape) == 5, "The input tensor should have 5 dimensions."
sample = self.conv_in(sample)
upscale_dtype = next(iter(self.up_blocks.parameters())).dtype
if self.training and self.gradient_checkpointing:
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
if is_torch_version(">=", "1.11.0"):
# middle
sample = torch.utils.checkpoint.checkpoint(
create_custom_forward(self.mid_block),
sample,
latent_embeds,
use_reentrant=False,
)
sample = sample.to(upscale_dtype)
# up
for up_block in self.up_blocks:
sample = torch.utils.checkpoint.checkpoint(
create_custom_forward(up_block),
sample,
latent_embeds,
use_reentrant=False,
)
else:
# middle
sample = torch.utils.checkpoint.checkpoint(
create_custom_forward(self.mid_block), sample, latent_embeds
)
sample = sample.to(upscale_dtype)
# up
for up_block in self.up_blocks:
sample = torch.utils.checkpoint.checkpoint(
create_custom_forward(up_block), sample, latent_embeds
)
else:
# middle
sample = self.mid_block(sample, latent_embeds)
sample = sample.to(upscale_dtype)
# up
for up_block in self.up_blocks:
sample = up_block(sample, latent_embeds)
# post-process
if latent_embeds is None:
sample = self.conv_norm_out(sample)
else:
sample = self.conv_norm_out(sample, latent_embeds)
sample = self.conv_act(sample)
sample = self.conv_out(sample)
return sample
class DiagonalGaussianDistribution(object):
def __init__(self, parameters: torch.Tensor, deterministic: bool = False):
if parameters.ndim == 3:
dim = 2 # (B, L, C)
elif parameters.ndim == 5 or parameters.ndim == 4:
dim = 1 # (B, C, T, H ,W) / (B, C, H, W)
else:
raise NotImplementedError
self.parameters = parameters
self.mean, self.logvar = torch.chunk(parameters, 2, dim=dim)
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
self.deterministic = deterministic
self.std = torch.exp(0.5 * self.logvar)
self.var = torch.exp(self.logvar)
if self.deterministic:
self.var = self.std = torch.zeros_like(
self.mean, device=self.parameters.device, dtype=self.parameters.dtype
)
def sample(self, generator: Optional[torch.Generator] = None) -> torch.FloatTensor:
# make sure sample is on the same device as the parameters and has same dtype
sample = randn_tensor(
self.mean.shape,
generator=generator,
device=self.parameters.device,
dtype=self.parameters.dtype,
)
x = self.mean + self.std * sample
return x
def kl(self, other: "DiagonalGaussianDistribution" = None) -> torch.Tensor:
if self.deterministic:
return torch.Tensor([0.0])
else:
reduce_dim = list(range(1, self.mean.ndim))
if other is None:
return 0.5 * torch.sum(
torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar,
dim=reduce_dim,
)
else:
return 0.5 * torch.sum(
torch.pow(self.mean - other.mean, 2) / other.var
+ self.var / other.var
- 1.0
- self.logvar
+ other.logvar,
dim=reduce_dim,
)
def nll(
self, sample: torch.Tensor, dims: Tuple[int, ...] = [1, 2, 3]
) -> torch.Tensor:
if self.deterministic:
return torch.Tensor([0.0])
logtwopi = np.log(2.0 * np.pi)
return 0.5 * torch.sum(
logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
dim=dims,
)
def mode(self) -> torch.Tensor:
return self.mean
@@ -1,553 +0,0 @@
import torch
import argparse
from safetensors.torch import save_file
import os
parser = argparse.ArgumentParser()
parser.add_argument("--diffusers_path", required=True, type=str)
parser.add_argument(
"--transformer_path", type=str, default=None, help="Path to save transformer model"
)
parser.add_argument(
"--vae_encoder_path", type=str, default=None, help="Path to save VAE encoder model"
)
parser.add_argument(
"--vae_decoder_path", type=str, default=None, help="Path to save VAE decoder model"
)
args = parser.parse_args()
def reverse_scale_shift(weight, dim):
scale, shift = weight.chunk(2, dim=0)
new_weight = torch.cat([shift, scale], dim=0)
return new_weight
def reverse_proj_gate(weight):
gate, proj = weight.chunk(2, dim=0)
new_weight = torch.cat([proj, gate], dim=0)
return new_weight
def convert_diffusers_transformer_to_mochi(state_dict):
original_state_dict = state_dict.copy()
new_state_dict = {}
# Convert patch_embed
new_state_dict["x_embedder.proj.weight"] = original_state_dict.pop(
"patch_embed.proj.weight"
)
new_state_dict["x_embedder.proj.bias"] = original_state_dict.pop(
"patch_embed.proj.bias"
)
# Convert time_embed
new_state_dict["t_embedder.mlp.0.weight"] = original_state_dict.pop(
"time_embed.timestep_embedder.linear_1.weight"
)
new_state_dict["t_embedder.mlp.0.bias"] = original_state_dict.pop(
"time_embed.timestep_embedder.linear_1.bias"
)
new_state_dict["t_embedder.mlp.2.weight"] = original_state_dict.pop(
"time_embed.timestep_embedder.linear_2.weight"
)
new_state_dict["t_embedder.mlp.2.bias"] = original_state_dict.pop(
"time_embed.timestep_embedder.linear_2.bias"
)
new_state_dict["t5_y_embedder.to_kv.weight"] = original_state_dict.pop(
"time_embed.pooler.to_kv.weight"
)
new_state_dict["t5_y_embedder.to_kv.bias"] = original_state_dict.pop(
"time_embed.pooler.to_kv.bias"
)
new_state_dict["t5_y_embedder.to_q.weight"] = original_state_dict.pop(
"time_embed.pooler.to_q.weight"
)
new_state_dict["t5_y_embedder.to_q.bias"] = original_state_dict.pop(
"time_embed.pooler.to_q.bias"
)
new_state_dict["t5_y_embedder.to_out.weight"] = original_state_dict.pop(
"time_embed.pooler.to_out.weight"
)
new_state_dict["t5_y_embedder.to_out.bias"] = original_state_dict.pop(
"time_embed.pooler.to_out.bias"
)
new_state_dict["t5_yproj.weight"] = original_state_dict.pop(
"time_embed.caption_proj.weight"
)
new_state_dict["t5_yproj.bias"] = original_state_dict.pop(
"time_embed.caption_proj.bias"
)
# Convert transformer blocks
num_layers = 48
for i in range(num_layers):
block_prefix = f"transformer_blocks.{i}."
new_prefix = f"blocks.{i}."
# norm1
new_state_dict[new_prefix + "mod_x.weight"] = original_state_dict.pop(
block_prefix + "norm1.linear.weight"
)
new_state_dict[new_prefix + "mod_x.bias"] = original_state_dict.pop(
block_prefix + "norm1.linear.bias"
)
if i < num_layers - 1:
new_state_dict[new_prefix + "mod_y.weight"] = original_state_dict.pop(
block_prefix + "norm1_context.linear.weight"
)
new_state_dict[new_prefix + "mod_y.bias"] = original_state_dict.pop(
block_prefix + "norm1_context.linear.bias"
)
else:
new_state_dict[new_prefix + "mod_y.weight"] = original_state_dict.pop(
block_prefix + "norm1_context.linear_1.weight"
)
new_state_dict[new_prefix + "mod_y.bias"] = original_state_dict.pop(
block_prefix + "norm1_context.linear_1.bias"
)
# Visual attention
q = original_state_dict.pop(block_prefix + "attn1.to_q.weight")
k = original_state_dict.pop(block_prefix + "attn1.to_k.weight")
v = original_state_dict.pop(block_prefix + "attn1.to_v.weight")
qkv_weight = torch.cat([q, k, v], dim=0)
new_state_dict[new_prefix + "attn.qkv_x.weight"] = qkv_weight
new_state_dict[new_prefix + "attn.q_norm_x.weight"] = original_state_dict.pop(
block_prefix + "attn1.norm_q.weight"
)
new_state_dict[new_prefix + "attn.k_norm_x.weight"] = original_state_dict.pop(
block_prefix + "attn1.norm_k.weight"
)
new_state_dict[new_prefix + "attn.proj_x.weight"] = original_state_dict.pop(
block_prefix + "attn1.to_out.0.weight"
)
new_state_dict[new_prefix + "attn.proj_x.bias"] = original_state_dict.pop(
block_prefix + "attn1.to_out.0.bias"
)
# Context attention
q = original_state_dict.pop(block_prefix + "attn1.add_q_proj.weight")
k = original_state_dict.pop(block_prefix + "attn1.add_k_proj.weight")
v = original_state_dict.pop(block_prefix + "attn1.add_v_proj.weight")
qkv_weight = torch.cat([q, k, v], dim=0)
new_state_dict[new_prefix + "attn.qkv_y.weight"] = qkv_weight
new_state_dict[new_prefix + "attn.q_norm_y.weight"] = original_state_dict.pop(
block_prefix + "attn1.norm_added_q.weight"
)
new_state_dict[new_prefix + "attn.k_norm_y.weight"] = original_state_dict.pop(
block_prefix + "attn1.norm_added_k.weight"
)
if i < num_layers - 1:
new_state_dict[new_prefix + "attn.proj_y.weight"] = original_state_dict.pop(
block_prefix + "attn1.to_add_out.weight"
)
new_state_dict[new_prefix + "attn.proj_y.bias"] = original_state_dict.pop(
block_prefix + "attn1.to_add_out.bias"
)
# MLP
new_state_dict[new_prefix + "mlp_x.w1.weight"] = reverse_proj_gate(
original_state_dict.pop(block_prefix + "ff.net.0.proj.weight")
)
new_state_dict[new_prefix + "mlp_x.w2.weight"] = original_state_dict.pop(
block_prefix + "ff.net.2.weight"
)
if i < num_layers - 1:
new_state_dict[new_prefix + "mlp_y.w1.weight"] = reverse_proj_gate(
original_state_dict.pop(block_prefix + "ff_context.net.0.proj.weight")
)
new_state_dict[new_prefix + "mlp_y.w2.weight"] = original_state_dict.pop(
block_prefix + "ff_context.net.2.weight"
)
# Output layers
new_state_dict["final_layer.mod.weight"] = reverse_scale_shift(
original_state_dict.pop("norm_out.linear.weight"), dim=0
)
new_state_dict["final_layer.mod.bias"] = reverse_scale_shift(
original_state_dict.pop("norm_out.linear.bias"), dim=0
)
new_state_dict["final_layer.linear.weight"] = original_state_dict.pop(
"proj_out.weight"
)
new_state_dict["final_layer.linear.bias"] = original_state_dict.pop("proj_out.bias")
new_state_dict["pos_frequencies"] = original_state_dict.pop("pos_frequencies")
print("Remaining Keys:", original_state_dict.keys())
return new_state_dict
def convert_diffusers_vae_to_mochi(state_dict):
original_state_dict = state_dict.copy()
encoder_state_dict = {}
decoder_state_dict = {}
# Convert encoder
prefix = "encoder."
encoder_state_dict["layers.0.weight"] = original_state_dict.pop(
f"{prefix}proj_in.weight"
)
encoder_state_dict["layers.0.bias"] = original_state_dict.pop(
f"{prefix}proj_in.bias"
)
# Convert block_in
for i in range(3):
encoder_state_dict[f"layers.{i+1}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.weight"
)
encoder_state_dict[f"layers.{i+1}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.bias"
)
encoder_state_dict[f"layers.{i+1}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.weight"
)
encoder_state_dict[f"layers.{i+1}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.bias"
)
encoder_state_dict[f"layers.{i+1}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.weight"
)
encoder_state_dict[f"layers.{i+1}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.bias"
)
encoder_state_dict[f"layers.{i+1}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.weight"
)
encoder_state_dict[f"layers.{i+1}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.bias"
)
# Convert down_blocks
down_block_layers = [3, 4, 6]
for block in range(3):
encoder_state_dict[
f"layers.{block+4}.layers.0.weight"
] = original_state_dict.pop(f"{prefix}down_blocks.{block}.conv_in.conv.weight")
encoder_state_dict[f"layers.{block+4}.layers.0.bias"] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.conv_in.conv.bias"
)
for i in range(down_block_layers[block]):
# Convert resnets
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.0.weight"
] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.weight"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.0.bias"
] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.bias"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.2.weight"
] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.weight"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.2.bias"
] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.bias"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.3.weight"
] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.weight"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.3.bias"
] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.bias"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.5.weight"
] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.weight"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.stack.5.bias"
] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.bias"
)
# Convert attentions
q = original_state_dict.pop(
f"{prefix}down_blocks.{block}.attentions.{i}.to_q.weight"
)
k = original_state_dict.pop(
f"{prefix}down_blocks.{block}.attentions.{i}.to_k.weight"
)
v = original_state_dict.pop(
f"{prefix}down_blocks.{block}.attentions.{i}.to_v.weight"
)
qkv_weight = torch.cat([q, k, v], dim=0)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.attn_block.attn.qkv.weight"
] = qkv_weight
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.weight"
] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.attentions.{i}.to_out.0.weight"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.bias"
] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.attentions.{i}.to_out.0.bias"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.attn_block.norm.weight"
] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.norms.{i}.norm_layer.weight"
)
encoder_state_dict[
f"layers.{block+4}.layers.{i+1}.attn_block.norm.bias"
] = original_state_dict.pop(
f"{prefix}down_blocks.{block}.norms.{i}.norm_layer.bias"
)
# Convert block_out
for i in range(3):
encoder_state_dict[f"layers.{i+7}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.weight"
)
encoder_state_dict[f"layers.{i+7}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.bias"
)
encoder_state_dict[f"layers.{i+7}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.weight"
)
encoder_state_dict[f"layers.{i+7}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.bias"
)
encoder_state_dict[f"layers.{i+7}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.weight"
)
encoder_state_dict[f"layers.{i+7}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.bias"
)
encoder_state_dict[f"layers.{i+7}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.weight"
)
encoder_state_dict[f"layers.{i+7}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.bias"
)
q = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_q.weight")
k = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_k.weight")
v = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_v.weight")
qkv_weight = torch.cat([q, k, v], dim=0)
encoder_state_dict[f"layers.{i+7}.attn_block.attn.qkv.weight"] = qkv_weight
encoder_state_dict[
f"layers.{i+7}.attn_block.attn.out.weight"
] = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_out.0.weight")
encoder_state_dict[
f"layers.{i+7}.attn_block.attn.out.bias"
] = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_out.0.bias")
encoder_state_dict[
f"layers.{i+7}.attn_block.norm.weight"
] = original_state_dict.pop(f"{prefix}block_out.norms.{i}.norm_layer.weight")
encoder_state_dict[
f"layers.{i+7}.attn_block.norm.bias"
] = original_state_dict.pop(f"{prefix}block_out.norms.{i}.norm_layer.bias")
# Convert output layers
encoder_state_dict["output_norm.weight"] = original_state_dict.pop(
f"{prefix}norm_out.norm_layer.weight"
)
encoder_state_dict["output_norm.bias"] = original_state_dict.pop(
f"{prefix}norm_out.norm_layer.bias"
)
encoder_state_dict["output_proj.weight"] = original_state_dict.pop(
f"{prefix}proj_out.weight"
)
# Convert decoder
prefix = "decoder."
decoder_state_dict["blocks.0.0.weight"] = original_state_dict.pop(
f"{prefix}conv_in.weight"
)
decoder_state_dict["blocks.0.0.bias"] = original_state_dict.pop(
f"{prefix}conv_in.bias"
)
# Convert block_in
for i in range(3):
decoder_state_dict[f"blocks.0.{i+1}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.weight"
)
decoder_state_dict[f"blocks.0.{i+1}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.bias"
)
decoder_state_dict[f"blocks.0.{i+1}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.weight"
)
decoder_state_dict[f"blocks.0.{i+1}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv1.conv.bias"
)
decoder_state_dict[f"blocks.0.{i+1}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.weight"
)
decoder_state_dict[f"blocks.0.{i+1}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.bias"
)
decoder_state_dict[f"blocks.0.{i+1}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.weight"
)
decoder_state_dict[f"blocks.0.{i+1}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_in.resnets.{i}.conv2.conv.bias"
)
# Convert up_blocks
up_block_layers = [6, 4, 3]
for block in range(3):
for i in range(up_block_layers[block]):
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.0.weight"
] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.weight"
)
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.0.bias"
] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.bias"
)
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.2.weight"
] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.weight"
)
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.2.bias"
] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.bias"
)
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.3.weight"
] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.weight"
)
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.3.bias"
] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.bias"
)
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.5.weight"
] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.weight"
)
decoder_state_dict[
f"blocks.{block+1}.blocks.{i}.stack.5.bias"
] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.bias"
)
decoder_state_dict[f"blocks.{block+1}.proj.weight"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.proj.weight"
)
decoder_state_dict[f"blocks.{block+1}.proj.bias"] = original_state_dict.pop(
f"{prefix}up_blocks.{block}.proj.bias"
)
# Convert block_out
for i in range(3):
decoder_state_dict[f"blocks.4.{i}.stack.0.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.weight"
)
decoder_state_dict[f"blocks.4.{i}.stack.0.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.bias"
)
decoder_state_dict[f"blocks.4.{i}.stack.2.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.weight"
)
decoder_state_dict[f"blocks.4.{i}.stack.2.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv1.conv.bias"
)
decoder_state_dict[f"blocks.4.{i}.stack.3.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.weight"
)
decoder_state_dict[f"blocks.4.{i}.stack.3.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.bias"
)
decoder_state_dict[f"blocks.4.{i}.stack.5.weight"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.weight"
)
decoder_state_dict[f"blocks.4.{i}.stack.5.bias"] = original_state_dict.pop(
f"{prefix}block_out.resnets.{i}.conv2.conv.bias"
)
# Convert output layers
decoder_state_dict["output_proj.weight"] = original_state_dict.pop(
f"{prefix}proj_out.weight"
)
decoder_state_dict["output_proj.bias"] = original_state_dict.pop(
f"{prefix}proj_out.bias"
)
return encoder_state_dict, decoder_state_dict
def ensure_safetensors_extension(path):
if not path.endswith(".safetensors"):
path = path + ".safetensors"
return path
def ensure_directory_exists(path):
directory = os.path.dirname(path)
if directory:
os.makedirs(directory, exist_ok=True)
def main(args):
from diffusers import MochiPipeline
pipe = MochiPipeline.from_pretrained(args.diffusers_path)
if args.transformer_path:
transformer_path = ensure_safetensors_extension(args.transformer_path)
ensure_directory_exists(transformer_path)
print(f"Converting transformer model...")
transformer_state_dict = convert_diffusers_transformer_to_mochi(
pipe.transformer.state_dict()
)
save_file(transformer_state_dict, transformer_path)
print(f"Saved transformer to {transformer_path}")
if args.vae_encoder_path and args.vae_decoder_path:
encoder_path = ensure_safetensors_extension(args.vae_encoder_path)
decoder_path = ensure_safetensors_extension(args.vae_decoder_path)
ensure_directory_exists(encoder_path)
ensure_directory_exists(decoder_path)
print(f"Converting VAE models...")
encoder_state_dict, decoder_state_dict = convert_diffusers_vae_to_mochi(
pipe.vae.state_dict()
)
save_file(encoder_state_dict, encoder_path)
print(f"Saved VAE encoder to {encoder_path}")
save_file(decoder_state_dict, decoder_path)
print(f"Saved VAE decoder to {decoder_path}")
elif args.vae_encoder_path or args.vae_decoder_path:
print(
"Warning: Both VAE encoder and decoder paths must be specified to convert VAE models."
)
if __name__ == "__main__":
main(args)
@@ -1,47 +0,0 @@
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_dit_input(model_type, latents):
if model_type == "mochi":
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
elif model_type == "hunyuan":
return latents * 0.476986
else:
raise NotImplementedError(f"model_type {model_type} not supported")
+28 -54
View File
@@ -2,15 +2,12 @@ import json
import torch.distributed as dist
import torch
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
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
):
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
@@ -22,16 +19,17 @@ def generate_video_and_latent(
generator=generator,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
output_type="latent_and_video",
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)
@@ -39,71 +37,47 @@ if __name__ == "__main__":
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("--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)
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
)
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()
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
# 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
)
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,
)
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"
)
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"
)
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)
@@ -111,7 +85,7 @@ if __name__ == "__main__":
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"
@@ -123,11 +97,11 @@ if __name__ == "__main__":
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:
with open(os.path.join(args.dataset_output_dir, "videos2caption.json"), 'w') as f:
json.dump(all_data, f, indent=4)
@@ -1,240 +0,0 @@
import torch
from diffusers import HunyuanVideoPipeline, HunyuanVideoTransformer3DModel, BitsAndBytesConfig
import imageio as iio
import math
import numpy as np
import io
import time
import argparse
import os
def export_to_video_bytes(fps, frames):
request = iio.core.Request("<bytes>", mode="w", extension=".mp4")
pyavobject = iio.plugins.pyav.PyAVPlugin(request)
if isinstance(frames, np.ndarray):
frames = (np.array(frames) * 255).astype('uint8')
else:
frames = np.array(frames)
new_bytes = pyavobject.write(frames, codec="libx264", fps=fps)
out_bytes = io.BytesIO(new_bytes)
return out_bytes
def export_to_video(frames, path, fps):
video_bytes = export_to_video_bytes(fps, frames)
video_bytes.seek(0)
with open(path, "wb") as f:
f.write(video_bytes.getbuffer())
def main(args):
torch.manual_seed(args.seed)
device = "cuda" if torch.cuda.is_available() else "cpu"
prompt_template = {
"template": (
"<|start_header_cid|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
"1. The main content and theme of the video."
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the contents, including objects, people, and anything else."
"3. Actions, events, behaviors temporal relationships, physical movement changes of the contents."
"4. Background environment, light, style, atmosphere, and qualities."
"5. Camera angles, movements, and transitions used in the video."
"6. Thematic and aesthetic concepts associated with the scene, i.e. realistic, futuristic, fairy tale, etc<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
),
"crop_start": 95,
}
model_id = args.model_path
if args.quantization == "nf4":
quantization_config = BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=torch.bfloat16, bnb_4bit_quant_type="nf4", llm_int8_skip_modules=["proj_out", "norm_out"])
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
model_id, subfolder="transformer/" ,torch_dtype=torch.bfloat16, quantization_config=quantization_config
)
if args.quantization == "int8":
quantization_config = BitsAndBytesConfig(load_in_8bit=True, llm_int8_skip_modules=["proj_out", "norm_out"])
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
model_id, subfolder="transformer/" ,torch_dtype=torch.bfloat16, quantization_config=quantization_config
)
elif not args.quantization:
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
model_id, subfolder="transformer/" ,torch_dtype=torch.bfloat16
).to(device)
print("Max vram for read transofrmer:", round(torch.cuda.max_memory_allocated(device="cuda") / 1024 ** 3, 3), "GiB")
torch.cuda.reset_max_memory_allocated(device)
if not args.cpu_offload:
pipe = HunyuanVideoPipeline.from_pretrained(model_id, torch_dtype=torch.bfloat16).to(device)
pipe.transformer = transformer
else:
pipe = HunyuanVideoPipeline.from_pretrained(model_id, transformer=transformer, torch_dtype=torch.bfloat16)
torch.cuda.reset_max_memory_allocated(device)
pipe.scheduler._shift = args.flow_shift
pipe.vae.enable_tiling()
if args.cpu_offload:
pipe.enable_model_cpu_offload()
print("Max vram for init pipeline:", round(torch.cuda.max_memory_allocated(device="cuda") / 1024 ** 3, 3), "GiB")
with open(args.prompt) as f:
prompts = f.readlines()
generator = torch.Generator("cpu").manual_seed(args.seed)
os.makedirs(os.path.dirname(args.output_path), exist_ok=True)
torch.cuda.reset_max_memory_allocated(device)
for prompt in prompts:
start_time = time.perf_counter()
output = pipe(
prompt=prompt,
height = args.height,
width = args.width,
num_frames = args.num_frames,
prompt_template=prompt_template,
num_inference_steps = args.num_inference_steps,
generator=generator,
).frames[0]
export_to_video(output, os.path.join(args.output_path, f"{prompt[:100]}.mp4"), fps=args.fps)
print("Time:", round(time.perf_counter() - start_time, 2), "seconds")
print("Max vram for denoise:", round(torch.cuda.max_memory_allocated(device="cuda") / 1024 ** 3, 3), "GiB")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# Basic parameters
parser.add_argument("--prompt", type=str, help="prompt file for inference")
parser.add_argument("--num_frames", type=int, default=16)
parser.add_argument("--height", type=int, default=256)
parser.add_argument("--width", type=int, default=256)
parser.add_argument("--num_inference_steps", type=int, default=50)
parser.add_argument("--model_path", type=str, default="data/hunyuan")
parser.add_argument("--output_path", type=str, default="./outputs/video")
parser.add_argument("--fps", type=int, default=24)
parser.add_argument("--quantization", type=str, default=None)
parser.add_argument("--cpu_offload", action="store_true")
# Additional parameters
parser.add_argument(
"--denoise-type",
type=str,
default="flow",
help="Denoise type for noised inputs.",
)
parser.add_argument("--seed", type=int, default=None, help="Seed for evaluation.")
parser.add_argument(
"--neg_prompt", type=str, default=None, help="Negative prompt for sampling."
)
parser.add_argument(
"--guidance_scale",
type=float,
default=1.0,
help="Classifier free guidance scale.",
)
parser.add_argument(
"--embedded_cfg_scale",
type=float,
default=6.0,
help="Embedded classifier free guidance scale.",
)
parser.add_argument(
"--flow_shift", type=int, default=7, help="Flow shift parameter."
)
parser.add_argument(
"--batch_size", type=int, default=1, help="Batch size for inference."
)
parser.add_argument(
"--num_videos",
type=int,
default=1,
help="Number of videos to generate per prompt.",
)
parser.add_argument(
"--load-key",
type=str,
default="module",
help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.",
)
parser.add_argument(
"--use-cpu-offload",
action="store_true",
help="Use CPU offload for the model load.",
)
parser.add_argument(
"--dit-weight",
type=str,
default="data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
)
parser.add_argument(
"--reproduce",
action="store_true",
help="Enable reproducibility by setting random seeds and deterministic algorithms.",
)
parser.add_argument(
"--disable-autocast",
action="store_true",
help="Disable autocast for denoising loop and vae decoding in pipeline sampling.",
)
# Flow Matching
parser.add_argument(
"--flow-reverse",
action="store_true",
help="If reverse, learning/sampling from t=1 -> t=0.",
)
parser.add_argument(
"--flow-solver", type=str, default="euler", help="Solver for flow matching."
)
parser.add_argument(
"--use-linear-quadratic-schedule",
action="store_true",
help="Use linear quadratic schedule for flow matching. Following MovieGen (https://ai.meta.com/static-resource/movie-gen-research-paper)",
)
parser.add_argument(
"--linear-schedule-end",
type=int,
default=25,
help="End step for linear quadratic schedule for flow matching.",
)
# Model parameters
parser.add_argument("--model", type=str, default="HYVideo-T/2-cfgdistill")
parser.add_argument("--latent-channels", type=int, default=16)
parser.add_argument(
"--precision", type=str, default="bf16", choices=["fp32", "fp16", "bf16", "fp8"]
)
parser.add_argument(
"--rope-theta", type=int, default=256, help="Theta used in RoPE."
)
parser.add_argument("--vae", type=str, default="884-16c-hy")
parser.add_argument(
"--vae-precision", type=str, default="fp16", choices=["fp32", "fp16", "bf16"]
)
parser.add_argument("--vae-tiling", action="store_true", default=True)
parser.add_argument("--text-encoder", type=str, default="llm")
parser.add_argument(
"--text-encoder-precision",
type=str,
default="fp16",
choices=["fp32", "fp16", "bf16"],
)
parser.add_argument("--text-states-dim", type=int, default=4096)
parser.add_argument("--text-len", type=int, default=256)
parser.add_argument("--tokenizer", type=str, default="llm")
parser.add_argument("--prompt-template", type=str, default="dit-llm-encode")
parser.add_argument(
"--prompt-template-video", type=str, default="dit-llm-encode-video"
)
parser.add_argument("--hidden-state-skip-layer", type=int, default=2)
parser.add_argument("--apply-final-norm", action="store_true")
parser.add_argument("--text-encoder-2", type=str, default="clipL")
parser.add_argument(
"--text-encoder-precision-2",
type=str,
default="fp16",
choices=["fp32", "fp16", "bf16"],
)
parser.add_argument("--text-states-dim-2", type=int, default=768)
parser.add_argument("--tokenizer-2", type=str, default="clipL")
parser.add_argument("--text-len-2", type=int, default=77)
args = parser.parse_args()
main(args)
-233
View File
@@ -1,233 +0,0 @@
import os
import imageio
import time
from einops import rearrange
import torch
import torchvision
import numpy as np
from pathlib import Path
from loguru import logger
from datetime import datetime
import argparse
from diffusers.utils import export_to_video
from fastvideo.models.hunyuan.utils.file_utils import save_videos_grid
from fastvideo.models.hunyuan.inference import HunyuanVideoSampler
import torch.distributed as dist
from fastvideo.utils.parallel_states import (
initialize_sequence_parallel_state,
nccl_info,
)
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 main(args):
initialize_distributed()
print(nccl_info.sp_size)
device = torch.cuda.current_device()
print(args)
models_root_path = Path(args.model_path)
if not models_root_path.exists():
raise ValueError(f"`models_root` not exists: {models_root_path}")
# Create save folder to save the samples
save_path = args.output_path
os.makedirs(os.path.dirname(save_path), exist_ok=True)
# Load models
hunyuan_video_sampler = HunyuanVideoSampler.from_pretrained(
models_root_path, args=args
)
# Get the updated args
args = hunyuan_video_sampler.args
# Start sampling
samples = []
with open(args.prompt) as f:
prompts = f.readlines()
for prompt in prompts:
outputs = hunyuan_video_sampler.predict(
prompt=prompt,
height=args.height,
width=args.width,
video_length=args.num_frames,
seed=args.seed,
negative_prompt=args.neg_prompt,
infer_steps=args.num_inference_steps,
guidance_scale=args.guidance_scale,
num_videos_per_prompt=args.num_videos,
flow_shift=args.flow_shift,
batch_size=args.batch_size,
embedded_guidance_scale=args.embedded_cfg_scale,
)
videos = rearrange(outputs["samples"], "b c t h w -> t b c h w")
outputs = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
outputs.append((x * 255).numpy().astype(np.uint8))
os.makedirs(os.path.dirname(args.output_path), exist_ok=True)
imageio.mimsave(
os.path.join(args.output_path, f"{prompt[:100]}.mp4"), outputs, fps=args.fps
)
if __name__ == "__main__":
parser = argparse.ArgumentParser()
# Basic parameters
parser.add_argument("--prompt", type=str, help="prompt file for inference")
parser.add_argument("--num_frames", type=int, default=16)
parser.add_argument("--height", type=int, default=256)
parser.add_argument("--width", type=int, default=256)
parser.add_argument("--num_inference_steps", type=int, default=50)
parser.add_argument("--model_path", type=str, default="data/hunyuan")
parser.add_argument("--output_path", type=str, default="./outputs/video")
parser.add_argument("--fps", type=int, default=24)
# Additional parameters
parser.add_argument(
"--denoise-type",
type=str,
default="flow",
help="Denoise type for noised inputs.",
)
parser.add_argument("--seed", type=int, default=None, help="Seed for evaluation.")
parser.add_argument(
"--neg_prompt", type=str, default=None, help="Negative prompt for sampling."
)
parser.add_argument(
"--guidance_scale",
type=float,
default=1.0,
help="Classifier free guidance scale.",
)
parser.add_argument(
"--embedded_cfg_scale",
type=float,
default=6.0,
help="Embedded classifier free guidance scale.",
)
parser.add_argument(
"--flow_shift", type=int, default=7, help="Flow shift parameter."
)
parser.add_argument(
"--batch_size", type=int, default=1, help="Batch size for inference."
)
parser.add_argument(
"--num_videos",
type=int,
default=1,
help="Number of videos to generate per prompt.",
)
parser.add_argument(
"--load-key",
type=str,
default="module",
help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.",
)
parser.add_argument(
"--use-cpu-offload",
action="store_true",
help="Use CPU offload for the model load.",
)
parser.add_argument(
"--dit-weight",
type=str,
default="data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
)
parser.add_argument(
"--reproduce",
action="store_true",
help="Enable reproducibility by setting random seeds and deterministic algorithms.",
)
parser.add_argument(
"--disable-autocast",
action="store_true",
help="Disable autocast for denoising loop and vae decoding in pipeline sampling.",
)
# Flow Matching
parser.add_argument(
"--flow-reverse",
action="store_true",
help="If reverse, learning/sampling from t=1 -> t=0.",
)
parser.add_argument(
"--flow-solver", type=str, default="euler", help="Solver for flow matching."
)
parser.add_argument(
"--use-linear-quadratic-schedule",
action="store_true",
help="Use linear quadratic schedule for flow matching. Following MovieGen (https://ai.meta.com/static-resource/movie-gen-research-paper)",
)
parser.add_argument(
"--linear-schedule-end",
type=int,
default=25,
help="End step for linear quadratic schedule for flow matching.",
)
# Model parameters
parser.add_argument("--model", type=str, default="HYVideo-T/2-cfgdistill")
parser.add_argument("--latent-channels", type=int, default=16)
parser.add_argument(
"--precision", type=str, default="bf16", choices=["fp32", "fp16", "bf16"]
)
parser.add_argument(
"--rope-theta", type=int, default=256, help="Theta used in RoPE."
)
parser.add_argument("--vae", type=str, default="884-16c-hy")
parser.add_argument(
"--vae-precision", type=str, default="fp16", choices=["fp32", "fp16", "bf16"]
)
parser.add_argument("--vae-tiling", action="store_true", default=True)
parser.add_argument("--text-encoder", type=str, default="llm")
parser.add_argument(
"--text-encoder-precision",
type=str,
default="fp16",
choices=["fp32", "fp16", "bf16"],
)
parser.add_argument("--text-states-dim", type=int, default=4096)
parser.add_argument("--text-len", type=int, default=256)
parser.add_argument("--tokenizer", type=str, default="llm")
parser.add_argument("--prompt-template", type=str, default="dit-llm-encode")
parser.add_argument(
"--prompt-template-video", type=str, default="dit-llm-encode-video"
)
parser.add_argument("--hidden-state-skip-layer", type=int, default=2)
parser.add_argument("--apply-final-norm", action="store_true")
parser.add_argument("--text-encoder-2", type=str, default="clipL")
parser.add_argument(
"--text-encoder-precision-2",
type=str,
default="fp16",
choices=["fp32", "fp16", "bf16"],
)
parser.add_argument("--text-states-dim-2", type=int, default=768)
parser.add_argument("--tokenizer-2", type=str, default="clipL")
parser.add_argument("--text-len-2", type=int, default=77)
args = parser.parse_args()
main(args)
+121 -90
View File
@@ -1,15 +1,12 @@
import torch
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
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,
)
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, nccl_info
import argparse
import os
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
import json
from typing import Optional
from safetensors.torch import save_file, load_file
@@ -20,116 +17,149 @@ import pdb
import copy
from typing import Dict
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import convert_unet_state_dict_to_peft
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)
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
)
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()
# Peiyuan: GPU seed will cause A100 and H100 to produce different results .....
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,
)
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/"
)
pipe = MochiPipeline.from_pretrained(
args.model_path, transformer=transformer, scheduler=scheduler
)
pipe.enable_vae_tiling()
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder = 'transformer/')
if args.lora_checkpoint_dir is not None:
print(f"Loading LoRA weights from {args.lora_checkpoint_dir}")
config_path = os.path.join(args.lora_checkpoint_dir, "lora_config.json")
with open(config_path, "r") as f:
lora_config_dict = json.load(f)
rank = lora_config_dict["lora_params"]["lora_rank"]
lora_alpha = lora_config_dict["lora_params"]["lora_alpha"]
lora_scaling = lora_alpha / rank
pipe.load_lora_weights(args.lora_checkpoint_dir, adapter_name="default")
pipe.set_adapters(["default"], [lora_scaling])
print(f"Successfully Loaded LoRA weights from {args.lora_checkpoint_dir}")
# 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)
)
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:
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:
generator = torch.Generator("cpu").manual_seed(args.seed)
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
if nccl_info.global_rank <= 0:
os.makedirs(args.output_path, exist_ok=True)
suffix = prompt.split(".")[0]
export_to_video(
video[0],
os.path.join(args.output_path, f"{suffix}.mp4"),
fps=30,
)
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):
generator = torch.Generator("cpu").manual_seed(args.seed)
videos = pipe(
prompt_embeds=prompt_embeds,
prompt_attention_mask=encoder_attention_mask,
@@ -141,14 +171,20 @@ def main(args):
generator=generator,
).frames
if nccl_info.global_rank <= 0:
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
# arg parse
parser = argparse.ArgumentParser()
parser.add_argument("--prompts", nargs="+", default=[])
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)
@@ -162,12 +198,7 @@ if __name__ == "__main__":
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('--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)
+8 -15
View File
@@ -1,11 +1,9 @@
import torch
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
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)
@@ -14,12 +12,8 @@ def main(args):
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
)
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()
@@ -35,15 +29,14 @@ def main(args):
num_inference_steps=args.num_inference_steps,
guidance_scale=args.guidance_scale,
).frames
for prompt, video in zip(args.prompts, videos):
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
# arg parse
parser = argparse.ArgumentParser()
parser.add_argument("--prompts", nargs="+", default=[])
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)
+204 -460
View File
@@ -5,14 +5,10 @@ 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.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.models.mochi_hf.mochi_latents_utils import normalize_dit_input
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
@@ -28,39 +24,33 @@ from fastvideo.utils.dataset_utils import LengthGroupedSampler
import wandb
from accelerate.utils import set_seed
from tqdm.auto import tqdm
from fastvideo.utils.fsdp_util import get_dit_fsdp_kwargs, apply_fsdp_checkpointing
from diffusers.utils import convert_unet_state_dict_to_peft
from diffusers import FlowMatchEulerDiscreteScheduler
from fastvideo.utils.load import load_transformer
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.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
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, get_peft_model_state_dict, set_peft_model_state_dict
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from fastvideo.utils.checkpoint import (
save_checkpoint,
save_lora_checkpoint,
resume_lora_optimizer,
from peft import LoraConfig, inject_adapter_in_model
from torch.distributed.fsdp import (
FullyShardedDataParallel as FSDP,
)
from fastvideo.utils.logging_ import main_print
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
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,
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.
@@ -71,13 +61,7 @@ def compute_density_for_timestep_sampling(
"""
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.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)
@@ -86,7 +70,6 @@ def compute_density_for_timestep_sampling(
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)
@@ -99,36 +82,16 @@ def get_sigmas(noise_scheduler, device, timesteps, n_dim=4, dtype=torch.float32)
return sigma
def train_one_step(
transformer,
model_type,
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,
):
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_dit_input(model_type, latents)
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(
u = compute_density_for_timestep_sampling(
weighting_scheme=weighting_scheme,
batch_size=batch_size,
generator=noise_random_generator,
@@ -141,58 +104,56 @@ def train_one_step(
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,
)
sigmas = get_sigmas(noise_scheduler, latents.device, timesteps, n_dim=latents.ndim, dtype=latents.dtype)
noisy_model_input = (1.0 - sigmas) * latents + sigmas * noise
# if rank<=0:
# print("2222222222222222222222222222222222222222222222")
# print(type(latents_attention_mask))
# print(latents_attention_mask)
with torch.autocast("cuda", torch.bfloat16):
model_pred = transformer(
noisy_model_input,
encoder_hidden_states,
timesteps,
encoder_attention_mask, # B, L
return_dict=False,
encoder_attention_mask, # B, L
return_dict= False
)[0]
# if rank<=0:
# print("333333333333333333333333333333333333333333333333")
if precondition_outputs:
model_pred = noisy_model_input - model_pred * sigmas
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
)
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()
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"])
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()
@@ -206,86 +167,53 @@ def main(args):
noise_random_generator = None
# Handle the repository creation
if rank <= 0 and args.output_dir is not None:
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 = load_transformer(
args.model_type,
args.dit_model_name_or_path,
transformer = MochiTransformer3DModel.from_pretrained(
args.pretrained_model_name_or_path,
torch.float32 if args.master_weight_type == "fp32" else torch.bfloat16,
subfolder="transformer",
torch_dtype = torch.float32,
#torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
)
if args.use_lora:
assert args.model_type == "mochi", "LoRA is only supported for Mochi model."
transformer.requires_grad_(False)
transformer_lora_config = LoraConfig(
lora_config = LoraConfig(
r=args.lora_rank,
lora_alpha=args.lora_alpha,
init_lora_weights=True,
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
init_lora_weights=True,
)
transformer.add_adapter(transformer_lora_config)
if args.resume_from_lora_checkpoint:
lora_state_dict = MochiPipeline.lora_state_dict(
args.resume_from_lora_checkpoint
)
transformer_state_dict = {
f'{k.replace("transformer.", "")}': v
for k, v in lora_state_dict.items()
if k.startswith("transformer.")
}
transformer_state_dict = convert_unet_state_dict_to_peft(transformer_state_dict)
incompatible_keys = set_peft_model_state_dict(
transformer, transformer_state_dict, adapter_name="default"
)
if incompatible_keys is not None:
# check only for unexpected keys
unexpected_keys = getattr(incompatible_keys, "unexpected_keys", None)
if unexpected_keys:
main_print(
f"Loading adapter weights from state_dict led to unexpected keys not found in the model: "
f" {unexpected_keys}. "
)
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, no_split_modules = get_dit_fsdp_kwargs(
transformer,
args.fsdp_sharding_startegy,
args.use_lora,
args.use_cpu_offload,
args.master_weight_type,
)
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 = [
no_split_module.__name__ for no_split_module in no_split_modules
]
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
transformer = FSDP(transformer, **fsdp_kwargs,)
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, no_split_modules, args.selective_checkpointing
)
apply_fsdp_checkpointing(transformer, args.selective_checkpointing)
# Set model as trainable.
transformer.train()
@@ -298,44 +226,39 @@ def main(args):
optimizer = torch.optim.AdamW(
params_to_optimize,
lr=args.learning_rate,
betas=(0.9, 0.999),
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_optimizer(
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,
num_training_steps=args.max_train_steps,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
last_epoch=init_steps - 1,
)
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
)
)
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,
@@ -343,339 +266,188 @@ def main(args):
pin_memory=True,
batch_size=args.train_batch_size,
num_workers=args.dataloader_num_workers,
drop_last=True,
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
)
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
)
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" 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"
)
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
# 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,
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
#todo future
for i in range(init_steps):
next(loader)
for step in range(init_steps + 1, args.max_train_steps + 1):
for step in range(init_steps + 1, args.max_train_steps+1):
start_time = time.time()
loss, grad_norm = train_one_step(
transformer,
args.model_type,
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,
)
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.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:
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
)
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.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
)
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(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()
parser.add_argument(
"--model_type", type=str, default="mochi", help="The type of model to train."
)
# 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
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, default=None)
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
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.",
)
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=str,
default="64",
help="use ',' to split multi sampling steps",
)
parser.add_argument(
"--validation_guidance_scale",
type=str,
default="4.5",
help="use ',' to split multi scale",
)
parser.add_argument("--validation_steps", type=int, default=50)
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***."
),
)
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("--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("--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("--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("--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(
@@ -685,16 +457,10 @@ if __name__ == "__main__":
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.",
"--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.",
"--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme."
)
parser.add_argument(
"--mode_scale",
@@ -703,36 +469,14 @@ if __name__ == "__main__":
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",
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."
)
parser.add_argument(
"--master_weight_type",
type=str,
default="fp32",
help="Weight type to use - fp32 or bf16.",
)
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)
main(args)
+104 -130
View File
@@ -1,39 +1,29 @@
# import
# 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 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.default_planner import DefaultSavePlanner, DefaultLoadPlanner
from torch.distributed.checkpoint.optimizer import load_sharded_optimizer_state_dict
from torch.distributed.fsdp import FullOptimStateDictConfig
from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
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),
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
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
# 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)
@@ -49,24 +39,21 @@ def save_checkpoint(model, optimizer, rank, output_dir, step, discriminator=Fals
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,
):
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),
model, StateDictType.FULL_STATE_DICT, FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
):
cpu_state = model.state_dict()
# todo move to get_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
# save using safetensors
if rank <= 0:
config_dict = dict(model.config)
config_path = os.path.join(hf_weight_dir, "config.json")
@@ -75,7 +62,8 @@ def save_checkpoint_generator_discriminator(
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)
@@ -86,53 +74,44 @@ def save_checkpoint_generator_discriminator(
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(),
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(),
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),
):
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"
)
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
)
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,
state_dict = weight_state_dict,
storage_reader=dist_cp.FileSystemReader(model_dir),
planner=DefaultLoadPlanner(),
)
@@ -141,62 +120,38 @@ def load_sharded_model(model, optimizer, model_dir, optimizer_dir):
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),
):
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:
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
)
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}"
)
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
):
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
)
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"
)
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),
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)
@@ -207,64 +162,83 @@ def resume_training(model, optimizer, checkpoint_dir, discriminator=False):
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
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):
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),
transformer,
StateDictType.FULL_STATE_DICT,
FullStateDictConfig(offload_to_cpu=True, rank0_only=True)
):
full_state_dict = transformer.state_dict()
lora_optim_state = FSDP.optim_state_dict(transformer, optimizer,)
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)
# save optimizer
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)
# save lora weight
main_print(f"--> saving LoRA checkpoint at step {step}")
transformer_lora_layers = get_peft_model_state_dict(
model=transformer, state_dict=full_state_dict
)
MochiPipeline.save_lora_weights(
save_directory=save_dir,
transformer_lora_layers=transformer_lora_layers,
is_main_process=True,
)
# save config
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,
},
'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_optimizer(transformer, checkpoint_dir, optimizer):
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
)
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 optimizer from step {step}")
return transformer, optimizer, step
step = config_dict['step']
main_print(f"--> Successfully resuming LoRA training from step {step}")
return transformer, optimizer, step
+51 -83
View File
@@ -10,12 +10,11 @@ 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:
@@ -113,6 +112,7 @@ class SeqAllToAll4D(torch.autograd.Function):
scatter_idx: int,
gather_idx: int,
) -> Tensor:
ctx.group = group
ctx.scatter_idx = scatter_idx
ctx.gather_idx = gather_idx
@@ -129,14 +129,18 @@ class SeqAllToAll4D(torch.autograd.Function):
None,
None,
)
def all_to_all_4D(
input_: torch.Tensor, scatter_dim: int = 2, gather_dim: int = 1,
input_: torch.Tensor,
scatter_dim: int = 2,
gather_dim: int = 1,
):
return SeqAllToAll4D.apply(nccl_info.group, input_, scatter_dim, gather_dim)
return SeqAllToAll4D.apply( nccl_info.group,input_, scatter_dim, gather_dim)
def _all_to_all(
input_: torch.Tensor,
world_size: int,
@@ -144,9 +148,7 @@ def _all_to_all(
scatter_dim: int,
gather_dim: int,
):
input_list = [
t.contiguous() for t in torch.tensor_split(input_, world_size, scatter_dim)
]
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()
@@ -168,9 +170,7 @@ class _AllToAll(torch.autograd.Function):
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
)
output = _all_to_all(input_, ctx.world_size, process_group, scatter_dim, gather_dim)
return output
@staticmethod
@@ -191,11 +191,14 @@ class _AllToAll(torch.autograd.Function):
def all_to_all(
input_: torch.Tensor, scatter_dim: int = 2, gather_dim: int = 1,
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.
@@ -234,7 +237,6 @@ class _AllGather(torch.autograd.Function):
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.
@@ -248,83 +250,49 @@ def all_gather(input_: torch.Tensor, dim: int = 1):
return _AllGather.apply(input_, dim)
def prepare_sequence_parallel_data(
hidden_states, encoder_hidden_states, attention_mask, encoder_attention_mask
):
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
):
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
)
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,
)
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),
)
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,
)
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)
+66 -151
View File
@@ -13,13 +13,11 @@ from collections import Counter
import random
IMG_EXTENSIONS = [".jpg", ".JPG", ".jpeg", ".JPEG", ".png", ".PNG"]
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."""
@@ -33,20 +31,17 @@ class DecordInit(object):
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
)
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})"
)
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:
@@ -55,8 +50,6 @@ def pad_to_multiple(number, ds_stride):
padding = ds_stride - remainder
return number + padding
# TODO
class Collate:
def __init__(self, args):
self.batch_size = args.train_batch_size
@@ -78,9 +71,9 @@ class Collate:
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]
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):
@@ -88,29 +81,13 @@ class Collate:
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"
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,
):
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
@@ -121,30 +98,13 @@ class Collate:
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,
)
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)]
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]
@@ -155,61 +115,50 @@ class Collate:
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_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_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)
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[0]-1) // ae_stride_thw[0] + 1),
max_tube_size[1] // ae_stride_thw[1],
max_tube_size[2] // ae_stride_thw[2],
]
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[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
]
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
]
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.
@@ -235,16 +184,13 @@ def split_to_even_chunks(indices, lengths, num_chunks, batch_size):
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))
]
chunk = chunk + [random.choice(chunk) for _ in range(batch_size - len(chunk))]
else:
chunk = random.choice(pad_chunks)
print(chunks[idx], "->", chunk)
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)
@@ -258,70 +204,48 @@ def megabatch_frame_alignment(megabatches, lengths):
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))
]
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,
):
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
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)
]
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
]
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
]
return [i for megabatch in shuffled_megabatches for batch in megabatch for i in batch]
class LengthGroupedSampler(Sampler):
@@ -335,9 +259,9 @@ class LengthGroupedSampler(Sampler):
batch_size: int,
rank: int,
world_size: int,
lengths: Optional[List[int]] = None,
group_frame=False,
group_resolution=False,
lengths: Optional[List[int]] = None,
group_frame=False,
group_resolution=False,
generator=None,
):
if lengths is None:
@@ -355,24 +279,15 @@ class LengthGroupedSampler(Sampler):
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,
)
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])
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
)
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")
-40
View File
@@ -1,40 +0,0 @@
import platform
import accelerate
import peft
import torch
import transformers
from transformers.utils import is_torch_cuda_available, is_torch_npu_available
VERSION = "1.2.0"
if __name__ == "__main__":
info = {
"FastVideo version": VERSION,
"Platform": platform.platform(),
"Python version": platform.python_version(),
"PyTorch version": torch.__version__,
"Transformers version": transformers.__version__,
"Accelerate version": accelerate.__version__,
"PEFT version": peft.__version__,
}
if is_torch_cuda_available():
info["PyTorch version"] += " (GPU)"
info["GPU type"] = torch.cuda.get_device_name()
if is_torch_npu_available():
info["PyTorch version"] += " (NPU)"
info["NPU type"] = torch.npu.get_device_name()
info["CANN version"] = torch.version.cann
try:
import bitsandbytes
info["Bitsandbytes version"] = bitsandbytes.__version__
except Exception:
pass
print("\n" + "\n".join([f"- {key}: {value}" for key, value in info.items()]) + "\n")
-347
View File
@@ -1,347 +0,0 @@
import torch
from fastvideo.models.mochi_hf.modeling_mochi import (
MochiTransformer3DModel,
MochiTransformerBlock,
)
from fastvideo.models.hunyuan.modules.models import (
HYVideoDiffusionTransformer,
MMDoubleStreamBlock,
MMSingleStreamBlock,
)
from fastvideo.models.hunyuan.vae.autoencoder_kl_causal_3d import AutoencoderKLCausal3D
from diffusers import AutoencoderKLMochi
from transformers import T5EncoderModel, AutoTokenizer
import os
from torch import nn
# Path
from pathlib import Path
import torch.nn.functional as F
from fastvideo.models.hunyuan.text_encoder import TextEncoder
from fastvideo.utils.logging_ import main_print
hunyuan_config = {
"mm_double_blocks_depth": 20,
"mm_single_blocks_depth": 40,
"rope_dim_list": [16, 56, 56],
"hidden_size": 3072,
"heads_num": 24,
"mlp_width_ratio": 4,
"guidance_embed": True,
}
PROMPT_TEMPLATE_ENCODE = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the image by detailing the color, shape, size, texture, "
"quantity, text, spatial relationships of the objects and background:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
)
PROMPT_TEMPLATE_ENCODE_VIDEO = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
"1. The main content and theme of the video."
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
"4. background environment, light, style and atmosphere."
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
)
NEGATIVE_PROMPT = "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion"
PROMPT_TEMPLATE = {
"dit-llm-encode": {"template": PROMPT_TEMPLATE_ENCODE, "crop_start": 36,},
"dit-llm-encode-video": {
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
"crop_start": 95,
},
}
class HunyuanTextEncoderWrapper(nn.Module):
def __init__(self, pretrained_model_name_or_path, device):
super().__init__()
text_len = 256
crop_start = PROMPT_TEMPLATE["dit-llm-encode-video"].get("crop_start", 0)
max_length = text_len + crop_start
# prompt_template
prompt_template = PROMPT_TEMPLATE["dit-llm-encode"]
# prompt_template_video
prompt_template_video = PROMPT_TEMPLATE["dit-llm-encode-video"]
text_encoder_path = os.path.join(pretrained_model_name_or_path, "text_encoder")
self.text_encoder = TextEncoder(
text_encoder_type="llm",
text_encoder_path=text_encoder_path,
max_length=max_length,
text_encoder_precision="fp16",
tokenizer_type="llm",
prompt_template=prompt_template,
prompt_template_video=prompt_template_video,
hidden_state_skip_layer=2,
apply_final_norm=False,
reproduce=False,
logger=None,
device=device,
)
text_encoder_path_2 = os.path.join(
pretrained_model_name_or_path, "text_encoder_2"
)
self.text_encoder_2 = TextEncoder(
text_encoder_type="clipL",
text_encoder_path=text_encoder_path_2,
max_length=77,
text_encoder_precision="fp16",
tokenizer_type="clipL",
reproduce=False,
logger=None,
device=device,
)
def encode_(self, prompt, text_encoder, clip_skip=None):
# TODO
device = self.text_encoder.device
data_type = "video"
num_videos_per_prompt = 1
text_inputs = text_encoder.text2tokens(prompt, data_type=data_type)
if clip_skip is None:
prompt_outputs = text_encoder.encode(
text_inputs, data_type="video", device=device
)
prompt_embeds = prompt_outputs.hidden_state
else:
prompt_outputs = text_encoder.encode(
text_inputs,
output_hidden_states=True,
data_type=data_type,
device=device,
)
prompt_embeds = prompt_outputs.hidden_states_list[-(clip_skip + 1)]
prompt_embeds = text_encoder.model.text_model.final_layer_norm(
prompt_embeds
)
attention_mask = prompt_outputs.attention_mask
if attention_mask is not None:
attention_mask = attention_mask.to(device)
bs_embed, seq_len = attention_mask.shape
attention_mask = attention_mask.repeat(1, num_videos_per_prompt)
attention_mask = attention_mask.view(
bs_embed * num_videos_per_prompt, seq_len
)
if text_encoder is not None:
prompt_embeds_dtype = text_encoder.dtype
elif self.transformer is not None:
prompt_embeds_dtype = self.transformer.dtype
else:
prompt_embeds_dtype = prompt_embeds.dtype
prompt_embeds = prompt_embeds.to(dtype=prompt_embeds_dtype, device=device)
if prompt_embeds.ndim == 2:
bs_embed, _ = prompt_embeds.shape
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt)
prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, -1)
else:
bs_embed, seq_len, _ = prompt_embeds.shape
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
prompt_embeds = prompt_embeds.view(
bs_embed * num_videos_per_prompt, seq_len, -1
)
return (prompt_embeds, attention_mask)
def encode_prompt(self, prompt):
prompt_embeds, attention_mask = self.encode_(prompt, self.text_encoder)
prompt_embeds_2, attention_mask_2 = self.encode_(prompt, self.text_encoder_2)
prompt_embeds_2 = F.pad(
prompt_embeds_2,
(0, prompt_embeds.shape[2] - prompt_embeds_2.shape[1]),
value=0,
).unsqueeze(1)
prompt_embeds = torch.cat([prompt_embeds_2, prompt_embeds], dim=1)
return prompt_embeds, attention_mask
class MochiTextEncoderWrapper(nn.Module):
def __init__(self, pretrained_model_name_or_path, device):
super().__init__()
self.text_encoder = T5EncoderModel.from_pretrained(
os.path.join(pretrained_model_name_or_path, "text_encoder")
).to(device)
self.tokenizer = AutoTokenizer.from_pretrained(
os.path.join(pretrained_model_name_or_path, "tokenizer")
)
self.max_sequence_length = 256
def encode_prompt(self, prompt):
device = self.text_encoder.device
dtype = 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=self.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[:, self.max_sequence_length - 1 : -1]
)
main_print(
f"Truncated text input: {prompt} to: {removed_text} for model input."
)
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.view(batch_size, seq_len, -1)
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
return prompt_embeds, prompt_attention_mask
def load_hunyuan_state_dict(model, dit_model_name_or_path):
load_key = "module"
model_path = dit_model_name_or_path
bare_model = "unknown"
state_dict = torch.load(
model_path, map_location=lambda storage, loc: storage, weights_only=True
)
if bare_model == "unknown" and ("ema" in state_dict or "module" in state_dict):
bare_model = False
if bare_model is False:
if load_key in state_dict:
state_dict = state_dict[load_key]
else:
raise KeyError(
f"Missing key: `{load_key}` in the checkpoint: {model_path}. The keys in the checkpoint "
f"are: {list(state_dict.keys())}."
)
model.load_state_dict(state_dict, strict=True)
return model
def load_transformer(
model_type,
dit_model_name_or_path,
pretrained_model_name_or_path,
master_weight_type,
):
if model_type == "mochi":
if dit_model_name_or_path:
transformer = MochiTransformer3DModel.from_pretrained(
dit_model_name_or_path,
torch_dtype=master_weight_type,
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
)
else:
transformer = MochiTransformer3DModel.from_pretrained(
pretrained_model_name_or_path,
subfolder="transformer",
torch_dtype=master_weight_type,
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
)
elif model_type == "hunyuan":
transformer = HYVideoDiffusionTransformer(
in_channels=16, out_channels=16, **hunyuan_config, dtype=master_weight_type,
)
transformer = load_hunyuan_state_dict(transformer, dit_model_name_or_path)
else:
raise ValueError(f"Unsupported model type: {model_type}")
return transformer
def load_vae(model_type, pretrained_model_name_or_path):
weight_dtype = torch.float32
if model_type == "mochi":
vae = AutoencoderKLMochi.from_pretrained(
pretrained_model_name_or_path, subfolder="vae", torch_dtype=weight_dtype
).to("cuda")
autocast_type = torch.bfloat16
fps = 30
elif model_type == "hunyuan":
vae_precision = torch.float32
vae_path = os.path.join(
pretrained_model_name_or_path, "hunyuan-video-t2v-720p/vae"
)
config = AutoencoderKLCausal3D.load_config(vae_path)
vae = AutoencoderKLCausal3D.from_config(config)
vae_ckpt = Path(vae_path) / "pytorch_model.pt"
assert vae_ckpt.exists(), f"VAE checkpoint not found: {vae_ckpt}"
ckpt = torch.load(vae_ckpt, map_location=vae.device, weights_only=True)
if "state_dict" in ckpt:
ckpt = ckpt["state_dict"]
if any(k.startswith("vae.") for k in ckpt.keys()):
ckpt = {
k.replace("vae.", ""): v
for k, v in ckpt.items()
if k.startswith("vae.")
}
vae.load_state_dict(ckpt)
vae = vae.to(dtype=vae_precision)
vae.requires_grad_(False)
vae = vae.to("cuda")
vae.eval()
autocast_type = torch.float32
fps = 24
return vae, autocast_type, fps
def load_text_encoder(model_type, pretrained_model_name_or_path, device):
if model_type == "mochi":
text_encoder = MochiTextEncoderWrapper(pretrained_model_name_or_path, device)
elif model_type == "hunyuan":
text_encoder = HunyuanTextEncoderWrapper(pretrained_model_name_or_path, device)
else:
raise ValueError(f"Unsupported model type: {model_type}")
return text_encoder
def get_no_split_modules(transformer):
# if of type MochiTransformer3DModel
if isinstance(transformer, MochiTransformer3DModel):
return (MochiTransformerBlock,)
elif isinstance(transformer, HYVideoDiffusionTransformer):
return (MMDoubleStreamBlock, MMSingleStreamBlock)
else:
raise ValueError(f"Unsupported transformer type: {type(transformer)}")
if __name__ == "__main__":
# test encode prompt
device = torch.cuda.current_device()
pretrained_model_name_or_path = "data/hunyuan"
text_encoder = load_text_encoder("hunyuan", pretrained_model_name_or_path, device)
prompt = "A man on stage claps his hands together while facing the audience. The audience, visible in the foreground, holds up mobile devices to record the event, capturing the moment from various angles. The background features a large banner with text identifying the man on stage. Throughout the sequence, the man's expression remains engaged and directed towards the audience. The camera angle remains constant, focusing on capturing the interaction between the man on stage and the audience."
prompt_embeds, attention_mask = text_encoder.encode_prompt(prompt)
@@ -1,24 +1,23 @@
import sys
import pdb
import os
def main_print(content):
if int(os.environ["LOCAL_RANK"]) <= 0:
if int(os.environ['LOCAL_RANK']) <= 0:
print(content)
# ForkedPdb().set_trace()
#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")
sys.stdin = open('/dev/stdin')
pdb.Pdb.interaction(self, *args, **kwargs)
finally:
sys.stdin = _stdin
-77
View File
@@ -1,77 +0,0 @@
from accelerate.logging import get_logger
import torch
logger = get_logger(__name__)
def get_optimizer(args, params_to_optimize, use_deepspeed: bool = False):
# Optimizer creation
supported_optimizers = ["adam", "adamw", "prodigy"]
if args.optimizer not in supported_optimizers:
logger.warning(
f"Unsupported choice of optimizer: {args.optimizer}. Supported optimizers include {supported_optimizers}. Defaulting to AdamW"
)
args.optimizer = "adamw"
if args.use_8bit_adam and not (args.optimizer.lower() not in ["adam", "adamw"]):
logger.warning(
f"use_8bit_adam is ignored when optimizer is not set to 'Adam' or 'AdamW'. Optimizer was "
f"set to {args.optimizer.lower()}"
)
if args.use_8bit_adam:
try:
import bitsandbytes as bnb
except ImportError:
raise ImportError(
"To use 8-bit Adam, please install the bitsandbytes library: `pip install bitsandbytes`."
)
if args.optimizer.lower() == "adamw":
optimizer_class = (
bnb.optim.AdamW8bit if args.use_8bit_adam else torch.optim.AdamW
)
optimizer = optimizer_class(
params_to_optimize,
betas=(args.adam_beta1, args.adam_beta2),
eps=args.adam_epsilon,
weight_decay=args.adam_weight_decay,
)
elif args.optimizer.lower() == "adam":
optimizer_class = bnb.optim.Adam8bit if args.use_8bit_adam else torch.optim.Adam
optimizer = optimizer_class(
params_to_optimize,
betas=(args.adam_beta1, args.adam_beta2),
eps=args.adam_epsilon,
weight_decay=args.adam_weight_decay,
)
elif args.optimizer.lower() == "prodigy":
try:
import prodigyopt
except ImportError:
raise ImportError(
"To use Prodigy, please install the prodigyopt library: `pip install prodigyopt`"
)
optimizer_class = prodigyopt.Prodigy
if args.learning_rate <= 0.1:
logger.warning(
"Learning rate is too low. When using prodigy, it's generally better to set learning rate around 1.0"
)
optimizer = optimizer_class(
params_to_optimize,
lr=args.learning_rate,
betas=(args.adam_beta1, args.adam_beta2),
beta3=args.prodigy_beta3,
weight_decay=args.adam_weight_decay,
eps=args.adam_epsilon,
decouple=args.prodigy_decouple,
use_bias_correction=args.prodigy_use_bias_correction,
safeguard_warmup=args.prodigy_safeguard_warmup,
)
return optimizer
+5 -16
View File
@@ -2,7 +2,6 @@ import torch
import torch.distributed as dist
import os
class COMM_INFO:
def __init__(self):
self.group = None
@@ -11,11 +10,8 @@ class COMM_INFO:
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:
@@ -23,29 +19,22 @@ def initialize_sequence_parallel_state(sequence_parallel_size):
initialize_sequence_parallel_group(sequence_parallel_size)
else:
nccl_info.sp_size = 1
nccl_info.global_rank = int(os.getenv("RANK", "0"))
nccl_info.global_rank = int(os.getenv('RANK', '0'))
nccl_info.rank_within_group = 0
nccl_info.group_id = int(os.getenv("RANK", "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
)
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
+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))
+91 -165
View File
@@ -1,29 +1,25 @@
from typing import Optional, Union, List
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 fastvideo.utils.communications import all_gather
from diffusers.utils.torch_utils import randn_tensor
from fastvideo.models.mochi_hf.pipeline_mochi import (
linear_quadratic_schedule,
retrieve_timesteps,
)
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.utils.logging import main_print
from fastvideo.distill.solver import PCMFMScheduler
from diffusers.utils import export_to_video
import os
import wandb
import gc
from fastvideo.utils.load import load_vae
def prepare_latents(
batch_size,
num_channels_latents,
@@ -42,10 +38,10 @@ def prepare_latents(
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,
@@ -64,9 +60,8 @@ def sample_validation_video(
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,
num_channels_latents=12,
vae_spatial_scale_factor = 8,
vae_temporal_scale_factor = 6,
):
device = vae.device
@@ -75,12 +70,11 @@ def sample_validation_video(
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
)
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,
@@ -91,14 +85,13 @@ def sample_validation_video(
device,
generator,
vae_spatial_scale_factor,
vae_temporal_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 = 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
@@ -107,11 +100,17 @@ def sample_validation_video(
sigmas = np.array(sigmas)
if scheduler_type == "euler":
timesteps, num_inference_steps = retrieve_timesteps(
scheduler, num_inference_steps, device, timesteps, sigmas,
scheduler,
num_inference_steps,
device,
timesteps,
sigmas,
)
else:
timesteps, num_inference_steps = retrieve_timesteps(
scheduler, num_inference_steps, device,
scheduler,
num_inference_steps,
device,
)
num_warmup_steps = max(len(timesteps) - num_inference_steps * scheduler.order, 0)
@@ -119,40 +118,31 @@ def sample_validation_video(
# 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:
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])
with torch.autocast("cuda", dtype=torch.bfloat16):
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]
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
)
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 = scheduler.step(noise_pred, t, latents.to(torch.float32), return_dict=False)[0]
latents = latents.to(latents_dtype)
if latents.dtype != latents_dtype:
@@ -160,183 +150,118 @@ def sample_validation_video(
# 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
):
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
)
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)
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)
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
with torch.autocast("cuda", dtype=vae.dtype):
video = vae.decode(latents, return_dict=False)[0]
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, # TODO
global_step,
scheduler_type="euler",
shift=1.0,
num_euler_timesteps=100,
linear_quadratic_threshold=0.025,
linear_range=0.5,
ema=False,
):
# TODO
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")
if args.model_type == "mochi":
vae_spatial_scale_factor = 8
vae_temporal_scale_factor = 6
num_channels_latents = 12
elif args.model_type == "hunyuan":
vae_spatial_scale_factor = 8
vae_temporal_scale_factor = 4
num_channels_latents = 16
else:
raise ValueError(f"Model type {args.model_type} not supported")
vae, autocast_type, fps = load_vae(
args.model_type, args.pretrained_model_name_or_path
)
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,
)
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
]
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
embe_dir = os.path.join(args.validation_prompt_dir, "prompt_embed")
mask_dir = os.path.join(args.validation_prompt_dir, "prompt_attention_mask")
embeds = sorted([f for f in os.listdir(embe_dir)])
masks = sorted([f for f in os.listdir(mask_dir)])
num_embeds = len(embeds)
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
num_sp_groups = int(os.getenv("WORLD_SIZE", '1')) // nccl_info.sp_size
# pad to multiple of groups
if num_embeds % num_sp_groups != 0:
validation_prompt_ids += [0] * (
num_sp_groups - num_embeds % num_sp_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
]
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(embe_dir, f"{embeds[i]}")
prompt_mask_path = os.path.join(mask_dir, f"{masks[i]}")
prompt_embeds = (
torch.load(prompt_embed_path, map_location="cpu", weights_only=True)
.to(device)
.unsqueeze(0)
)
prompt_attention_mask = (
torch.load(prompt_mask_path, map_location="cpu", weights_only=True)
.to(device)
.unsqueeze(0)
)
negative_prompt_embeds = torch.zeros(256, 4096).to(device).unsqueeze(0)
negative_prompt_attention_mask = (
torch.zeros(256).bool().to(device).unsqueeze(0)
)
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,
vae_spatial_scale_factor=vae_spatial_scale_factor,
vae_temporal_scale_factor=vae_temporal_scale_factor,
num_channels_latents=num_channels_latents,
)[0]
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
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
# 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=fps)
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 = {
@@ -346,3 +271,4 @@ def log_validation(
]
}
wandb.log(logs, step=global_step)
+39
View File
@@ -0,0 +1,39 @@
num_gpus=4
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/sample/sample_t2v_mochi.py \
--model_path data/mochi \
--prompt_path data/prompt.txt \
--transformer_path data/outputs/video_distill_synthetic/checkpoint-1500 \
--num_frames 163 \
--height 480 \
--width 848 \
--num_inference_steps 8 \
--guidance_scale 4.5 \
--output_path outputs_video/distill_lq_163_1500_precision_stochastic_0.7 \
--shift 8 \
--seed 12345 \
--scheduler_type "pcm_linear_quadratic"
num_gpus=4
torchrun --nnodes=1 --nproc_per_node=$num_gpus --master_port 29503 \
fastvideo/sample/sample_t2v_mochi.py \
--model_path data/mochi \
--prompt_embed_path "data/synthetic_debug2/prompt_embed/2.pt" \
--encoder_attention_mask_path "data/synthetic_debug2/prompt_attention_mask/1.pt" \
--num_frames 163 \
--height 480 \
--width 848 \
--num_inference_steps 32 \
--guidance_scale 4.5 \
--output_path outputs_video/debug \
--shift 8 \
--seed 12345 \
--scheduler_type "pcm_linear_quadratic"
+11
View File
@@ -0,0 +1,11 @@
python fastvideo/sample/sample_t2v_mochi_no_sp.py \
--model_path data/mochi \
--prompts "A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand's movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough." \
--num_frames 163 \
--height 480 \
--width 848 \
--num_inference_steps 64 \
--guidance_scale 0.0 \
--seed 12346 \
--transformer_path data/outputs/debug/checkpoint-100/transformer \
--output_path outputs_video/single_no_guidance
-130
View File
@@ -1,130 +0,0 @@
# Prediction interface for Cog ⚙️
# https://cog.run/python
from cog import BasePredictor, Input, Path
import os
import time
import torch
import imageio
import argparse
import subprocess
import torchvision
import numpy as np
from einops import rearrange
MODEL_CACHE = 'FastHunyuan'
os.environ['MODEL_BASE'] = './'+MODEL_CACHE
from fastvideo.models.hunyuan.inference import HunyuanVideoSampler
MODEL_URL = "https://weights.replicate.delivery/default/FastVideo/FastHunyuan/model.tar"
def download_weights(url, dest):
start = time.time()
print("downloading url: ", url)
print("downloading to: ", dest)
subprocess.check_call(["pget", "-xf", url, dest], close_fds=False)
print("downloading took: ", time.time() - start)
class Predictor(BasePredictor):
def setup(self):
"""Load the model into memory"""
print("Model Base: " + os.environ['MODEL_BASE'])
# Download weights
if not os.path.exists(MODEL_CACHE):
download_weights(MODEL_URL, MODEL_CACHE)
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
args = argparse.Namespace(
num_frames=125,
height=720,
width=1280,
num_inference_steps=6,
fps=24,
denoise_type='flow',
seed=1024,
neg_prompt=None,
guidance_scale=1.0,
embedded_cfg_scale=6.0,
flow_shift=17,
batch_size=1,
num_videos=1,
load_key='module',
use_cpu_offload=False,
dit_weight='FastHunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt',
reproduce=True,
disable_autocast=False,
flow_reverse=True,
flow_solver='euler',
use_linear_quadratic_schedule=False,
linear_schedule_end=25,
model='HYVideo-T/2-cfgdistill',
latent_channels=16,
precision='bf16',
rope_theta=256,
vae='884-16c-hy',
vae_precision='fp16',
vae_tiling=True,
text_encoder='llm',
text_encoder_precision='fp16',
text_states_dim=4096,
text_len=256,
tokenizer='llm',
prompt_template='dit-llm-encode',
prompt_template_video='dit-llm-encode-video',
hidden_state_skip_layer=2,
apply_final_norm=False,
text_encoder_2='clipL',
text_encoder_precision_2='fp16',
text_states_dim_2=768,
tokenizer_2='clipL',
text_len_2=77,
model_path=MODEL_CACHE,
)
self.model = HunyuanVideoSampler.from_pretrained(MODEL_CACHE, args=args)
def predict(
self,
prompt: str = Input(description="Text prompt for video generation", default="A cat walks on the grass, realistic style."),
negative_prompt: str = Input(description="Text prompt to specify what you don't want in the video.", default=""),
width: int = Input(description="Width of output video", default=1280, ge=256),
height: int = Input(description="Height of output video", default=720, ge=256),
num_frames: int = Input(description="Number of frames to generate", default=125, ge=16),
num_inference_steps: int = Input(description="Number of denoising steps", default=6, ge=1, le=50),
guidance_scale: float = Input(description="Classifier free guidance scale", default=1.0, ge=0.1, le=10.0),
embedded_cfg_scale: float = Input(description="Embedded classifier free guidance scale", default=6.0, ge=0.1, le=10.0),
flow_shift: int = Input(description="Flow shift parameter", default=17, ge=1, le=20),
fps: int = Input(description="Frames per second of output video", default=24, ge=1, le=60),
seed: int = Input(description="0 for Random seed. Set for reproducible generation", default=0),
) -> Path:
"""Run video generation"""
if seed <=0:
seed = int.from_bytes(os.urandom(2), "big")
print(f"Using seed: {seed}")
outputs = self.model.predict(
prompt=prompt,
height=height,
width=width,
video_length=num_frames,
seed=seed,
negative_prompt=negative_prompt,
infer_steps=num_inference_steps,
guidance_scale=guidance_scale,
embedded_guidance_scale=embedded_cfg_scale,
flow_shift=flow_shift,
flow_reverse=True,
batch_size=1,
num_videos_per_prompt=1,
)
# Process output video
videos = rearrange(outputs["samples"], "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
# Save video
output_path = Path("/tmp/output.mp4")
imageio.mimsave(str(output_path), frames, fps=fps)
return Path(output_path)
+11
View File
@@ -0,0 +1,11 @@
import json
import os
path = "data/outputs/BW_Testrun/checkpoint-0/config.json"
with open(path, 'r') as f:
data = json.load(f)
# save with indent
with open(path, 'w') as f:
json.dump(data, f, indent=4)
+6 -1
View File
@@ -21,7 +21,12 @@ dependencies = [
"timm==1.0.11", "torchdiffeq==0.2.4", "torchmetrics==1.5.1", "tqdm==4.66.5", "urllib3==2.2.0", "uvicorn==0.32.0",
"scikit-video==1.1.11", "imageio-ffmpeg==0.5.1", "sentencepiece==0.2.0", "beautifulsoup4==4.12.3", "ftfy==6.3.0",
"moviepy==1.0.3", "wandb==0.18.5", "tensorboard==2.18.0", "pydantic==2.9.2", "gradio==5.3.0", "huggingface_hub==0.26.1", "protobuf==5.28.3",
"watch", "gpustat", "peft==0.13.2", "liger_kernel==0.4.1", "einops==0.8.0", "wheel==0.44.0", "loguru", "diffusers==0.32.0", "bitsandbytes", "pytest", "requests-mock"]
"watch", "gpustat", "peft==0.13.2", "liger_kernel==0.4.1", "einops==0.8.0", "wheel==0.44.0"
]
[project.optional-dependencies]
train = ["deepspeed==0.15.3"]
dev = ["mypy==1.8.0"]
[tool.setuptools.packages.find]
@@ -0,0 +1,25 @@
# export WANDB_MODE="offline"
GPU_NUM=8
MODEL_PATH="/ephemeral/hao.zhang/outputfolder/ckptfolder/mochi_diffuser"
DATA_MERGE_PATH="/ephemeral/hao.zhang/resourcefolder/Mochi-Synthetic-Data-BW-Finetune/merge.txt"
OUTPUT_DIR="./data/BW-Finetune-Synthetic-Data_test"
rchrun --nproc_per_node=$GPU_NUM \
./fastvideo/utils/data_preprocess/finetune_data_VAE.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--train_batch_size=1 \
--max_height=480 \
--max_width=848 \
--num_frames=163 \
--dataloader_num_workers 1 \
--output_dir=$OUTPUT_DIR
to
torchrun --nproc_per_node=$GPU_NUM \
./fastvideo/utils/data_preprocess/finetune_data_T5.py \
--model_path $MODEL_PATH \
--output_dir=$OUTPUT_DIR
-40
View File
@@ -1,40 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
DATA_DIR=./data
torchrun --nnodes 1 --nproc_per_node 8\
fastvideo/distill.py\
--seed 42\
--pretrained_model_name_or_path $DATA_DIR/hunyuan\
--dit_model_name_or_path $DATA_DIR/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt\
--model_type "hunyuan" \
--cache_dir "$DATA_DIR/.cache"\
--data_json_path "$DATA_DIR/Hunyuan-30K-Distill-Data/videos2caption.json"\
--validation_prompt_dir "$DATA_DIR/Hunyuan-Distill-Data/validation"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 24\
--sp_size 1\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=2000\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "2,4,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/hy_phase1_shift17_bs_32"\
--tracker_project_name Hunyuan_Distill \
--num_frames 93 \
--shift 17 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--not_apply_cfg_solver
-38
View File
@@ -1,38 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
torchrun --nnodes 1 --nproc_per_node 4 \
fastvideo/distill.py \
--seed 42 \
--pretrained_model_name_or_path data/mochi \
--model_type "mochi" \
--cache_dir data/.cache \
--data_json_path data/Merge-30k-Data/video2caption.json \
--validation_prompt_dir data/Image-Vid-Finetune-Mochi/validation \
--gradient_checkpointing \
--train_batch_size=1 \
--num_latent_t 28 \
--sp_size 4 \
--train_sp_batch_size 2 \
--dataloader_num_workers 4 \
--gradient_accumulation_steps=1 \
--max_train_steps=4000 \
--learning_rate=1e-6 \
--mixed_precision=bf16 \
--checkpointing_steps=64 \
--validation_steps 1 \
--validation_sampling_steps 8 \
--checkpoints_total_limit 3 \
--allow_tf32 \
--ema_start_step 0 \
--cfg 0.0 \
--log_validation \
--output_dir="data/outputs/lq_euler_50_thres0.1_lrg_0.75_phase1_lr1e-6_repro" \
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale 0.5,1.5,2.5 \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--multi_phased_distill_schedule 4000-1
+19
View File
@@ -0,0 +1,19 @@
from huggingface_hub import snapshot_download, hf_hub_download
import argparse
# set args for repo_id, local_dir, repo_type,
if __name__ == '__main__':
parser = argparse.ArgumentParser(description='Download a dataset or model from the Hugging Face Hub')
parser.add_argument('--repo_id', type=str, help='The ID of the repository to download')
parser.add_argument('--local_dir', type=str, help='The local directory to download the repository to')
parser.add_argument('--repo_type', type=str, help='The type of repository to download (dataset or model)')
parser.add_argument('--file_name', type=str, help='The file name to download')
args = parser.parse_args()
if args.file_name:
hf_hub_download(repo_id=args.repo_id, filename=args.file_name, repo_type=args.repo_type, local_dir=args.local_dir)
else:
snapshot_download(repo_id=args.repo_id,
local_dir=args.local_dir,
repo_type=args.repo_type,
local_dir_use_symlinks=False,
resume_download=True)
+56
View File
@@ -0,0 +1,56 @@
export WANDB_MODE=online
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
export NCCL_DEBUG=INFO
export FI_PROVIDER=efa
export FI_EFA_USE_DEVICE_RDMA=1
export NCCL_PROTO=simple
torchrun --nnodes 2 --nproc_per_node 8\
--node_rank=0 \
--rdzv_id=456 \
--rdzv_backend=c10d \
--rdzv_endpoint=[MASTER_NODE_IP_ADDRESS]:29500 \
fastvideo/distill.py\
--seed 42\
--pretrained_model_name_or_path data/mochi\
--cache_dir "data/.cache"\
--data_json_path "data/Merge-30k-Data/video2caption.json"\
--validation_prompt_dir "data/validation_embeddings/validation_prompt_embed_mask"\
--uncond_prompt_dir "data/validation_embeddings/uncond_prompt_embed_mask"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 28\
--sp_size 4\
--train_sp_batch_size 2\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=4000\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=500\
--validation_steps 250\
--validation_sampling_steps 8 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--ema_decay 0.999\
--log_validation\
--output_dir="data/outputs/lq_euler_50_thresh_0.025"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale 4.5 \
--num_euler_timesteps 50
gsutil cp data/outputs/lq_euler_50/checkpoint-4000 gs://vid_gen/runlong_temp_folder_for_pandas70m_debugging/fastvid/lq_euler_50/checkpoint-4000
+49
View File
@@ -0,0 +1,49 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=online
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
export FI_PROVIDER=efa
export FI_EFA_USE_DEVICE_RDMA=1
export NCCL_PROTO=simple
DATA_DIR=/data
IP=10.4.139.86
torchrun --nnodes 2 --nproc_per_node 8\
--node_rank=0 \
--rdzv_id=456 \
--rdzv_backend=c10d \
--rdzv_endpoint=$IP:29500 \
fastvideo/distill.py\
--seed 42\
--pretrained_model_name_or_path $DATA_DIR/mochi\
--cache_dir "data/.cache"\
--data_json_path "$DATA_DIR/Merge-30k-Data/video2caption.json"\
--validation_prompt_dir "$DATA_DIR/validation_embeddings/validation_prompt_embed_mask"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 28\
--sp_size 4\
--train_sp_batch_size 4\
--dataloader_num_workers 4\
--gradient_accumulation_steps=2\
--max_train_steps=4000\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=500\
--validation_steps 125\
--validation_sampling_steps 8 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--ema_decay 0.999\
--log_validation\
--output_dir="$DATA_DIR/outputs/lq_euler_50_thresh0.05_bs32"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "2.5,3.5,4.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.05
+45
View File
@@ -0,0 +1,45 @@
export WANDB_MODE=online
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
export LD_LIBRARY_PATH=/opt/amazon/efa/lib:/opt/aws-ofi-nccl/lib:$LD_LIBRARY_PATH
export NCCL_DEBUG=INFO
export FI_PROVIDER=efa
export FI_EFA_USE_DEVICE_RDMA=1
export NCCL_PROTO=simple
torchrun --nnodes 2 --nproc_per_node 8\
--node_rank=0 \
--rdzv_id=456 \
--rdzv_backend=c10d \
--rdzv_endpoint=[MASTER_NODE_IP_ADDRESS]:29500 \
fastvideo/distill.py\
--seed 42\
--pretrained_model_name_or_path data/mochi\
--cache_dir "data/.cache"\
--data_json_path "data/Merge-30k-Data/video2caption.json"\
--validation_prompt_dir "data/validation_embeddings/validation_prompt_embed_mask"\
--gradient_checkpointing\
--train_batch_size=1\
--num_latent_t 28\
--sp_size 4\
--train_sp_batch_size 2\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=4000\
--learning_rate=1e-7\
--mixed_precision="bf16"\
--checkpointing_steps=500\
--validation_steps 125\
--validation_sampling_steps 8 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--ema_decay 0.999\
--log_validation\
--output_dir="data/outputs/lq_euler_50_thresh0.05_lr_1e-7"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "2.5,3.5,4.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.05

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