Compare commits

..
68 Commits
Author SHA1 Message Date
“BrianChen1129” d7b07adc08 pipeline retrace; bug for black video 2024-12-09 09:47:06 +00:00
“BrianChen1129” 8a41e61cec add orginal mochi but performance bad 2024-12-09 08:44:24 +00:00
BrianChen1129â9 60ce6a62a6 delete sage in mochi-genmo 2024-12-08 03:56:26 +00:00
BrianChen1129â9 7b2ca9abec inference success 2024-12-08 03:54:07 +00:00
BrianChen1129â9 b537e01a88 inference success 2024-12-08 03:53:00 +00:00
Yongqi Chen a233b58a6c syn 2024-12-07 22:38:31 -05:00
Yongqi Chen d24b25a3e1 syn 2024-12-07 22:34:01 -05:00
BrianChen1129â9 2a7c147c4e syn with main 2024-12-08 02:38:43 +00:00
BrianChen1129â9 0c1c939d59 genmo mochi inference ready: 2024-12-08 02:35:40 +00:00
BrianChen1129â9 881e1f130a genmo mochi inference ready: 2024-12-08 02:35:08 +00:00
BrianChen1129â9 00e899cd90 syn 2024-12-07 09:21:02 +00:00
Zhang Peiyuan de1e8d868e Cleanup (#75) 2024-12-06 20:56:06 -08:00
BrianChen1129 9f4151526b add sageattn 2024-12-07 00:33:12 +00:00
Brian ChenandBrianChenn1129 98b92be25e add web demo (#73)
Co-authored-by: BrianChenn1129 <yonqgich@umich>
2024-12-06 09:47:22 -08:00
Yongqi Chen 1ce3983d68 add cpu offload 2024-12-06 12:45:08 -05:00
Yongqi Chen ee241cfa4d syn 2024-12-06 01:32:55 -05:00
Yongqi Chen 8106ee3f3d syn 2024-12-06 01:30:43 -05:00
Yongqi Chen 4dde52be9c remove conflict in train.py 2024-12-06 01:27:39 -05:00
BrianChenn1129 58abda5c09 syn with main: 2024-12-06 06:19:49 +00:00
BrianChenn1129 3e6019415a add demo 2024-12-06 06:12:43 +00:00
Zhang Peiyuanandforeverpiano 8d41d505fe [cleanup] (#72)
Co-authored-by: foreverpiano <pianoqwz@gmail.com>
2024-12-05 21:02:10 -08:00
Brian Chen 8cfdf58a17 [Feat] HF Lora
yongqich@umich.edu
2024-12-05 18:50:53 -08:00
Yongqi Chen ccaf43c195 syn 2024-12-05 21:31:25 -05:00
Yongqi Chen 351538db29 syn 2024-12-05 21:27:03 -05:00
Yongqi Chen 361b24612d syn 2024-12-04 15:28:18 -05:00
Yongqi Chen 6ad03bea79 syn 2024-12-04 15:27:38 -05:00
Yongqi Chen 86e1f88877 syn 2024-12-04 00:11:22 -05:00
Yongqi Chen e212a9c6b9 syn with yongqi-dev2 and main 2024-12-03 23:51:46 -05:00
Zhang Peiyuan cf15594055 [Feat & Debug] fix uncond; multi guidance validaiton; multiphase schedule; linear range (#64)
typo

update

update

typo

[Debug] Typo (#65)

debug gradient accumulation loss

update gitignore

typo

typo

update scripts

update experiment 10

new script

typo

update

update

update

update

wandb offline and dir

runlong

update

update

update

update

update

update

update

fix finetune code bug

update

update

update

add l2

update
2024-11-30 16:43:24 -08:00
Zhang Peiyuan ce95c2df29 [Feat] EMA Distill; Distributed validation (#63)
update
2024-11-29 21:46:40 -08:00
Zhang Peiyuan d417e4c7c4 [Feat] Refactor GAN; State saving & Resume; Experiments script (#59)
typo

add upload command

add

add ema_transform

add ema transformer

remove harcode

ok

ema

remove hardcode

update env

update script; no sp

revert to sp=4, sp bs=2, full shard

add aws efo env

distributed validation

readme

distributed validation

add gupload

linear range; fix uncond; multi guidance validation

add script
2024-11-28 21:09:31 -08:00
Zhang Peiyuan 2a70d05b4f [Feat] Training precision (#57) 2024-11-27 16:25:07 -08:00
Zhang Peiyuan 03187fd83a [Feat][Debug] linear quadratic distill; HF precision bug (#56) 2024-11-27 15:13:20 -08:00
Zhang Peiyuan 8a128ad815 [Fix] Squeeze bug (#55) 2024-11-26 13:29:03 -08:00
Zhang Peiyuan 13f665e455 [Feat] PCM Distill; Refactor FM logit to be compatible with all SD3/Flux scheduler. (#54) 2024-11-26 12:26:24 -08:00
rlsu9andRunlong 94ba0ab6ea [feat]: Add batchy data preprocess (#53)
Co-authored-by:Runlong <rlsu9@ucsd.edu>
2024-11-25 21:52:36 -08:00
rlsu9andrunlong 5a5d0ef1a0 [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 45e4adca4d [Fix]: Resolve config bug and seed (#51)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2024-11-15 20:29:48 -08:00
Zhang Peiyuanandforeverpiano 7106eadffc [feat]: Add LADD (#45)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
Co-authored-by: foreverpiano <pianoqwz@gmail.com>
2024-11-15 19:28:41 -08:00
Yongqi ChenandPeiyuan Zhang 2f3a8661bf [feat]: add lr scheduler; precision bug fix; add naive dataloader resume (#49)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2024-11-14 22:57:38 -05:00
Yongqi ChenandPeiyuan Zhang 6d0082c1a9 [Feat] Lora resume (#48)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2024-11-13 01:30:58 -05:00
Zhang PeiyuanandYongqi Chen 3d9189571a [feat]: Add lora (#47)
Co-authored-by: Yongqi Chen <144848849+BrianChen1129@users.noreply.github.com>
2024-11-11 13:12:14 -08:00
Zhang Peiyuan b042e321a1 [Refactor] Switch to FSDP (#42) 2024-11-09 15:46:25 -08:00
rlsu9 44bda9f8a3 [feat]: Add adaptive fps dataloader and remove redundant code (#41)
No checkout mochi
2024-11-06 19:52:48 -08:00
rlsu9 035ba5f5cf [feat]: Add vae encoder embedded generator to main (#30) 2024-11-06 11:58:15 -08:00
Zhang Peiyuan 52ba538e8e [feat]: Add validation logging with SP (#36) 2024-11-05 14:46:20 -08:00
Zhang Peiyuanandforeverpiano f0bc297260 [Feat] Sequence Parallel (#31)
Co-authored-by: foreverpiano <pianoqwz@gmail.com>
2024-11-05 08:13:01 -08:00
rlsu9 7413b1dd5f [feat]: update data preprocess 2024-10-31 00:13:05 +00:00
Peiyuan Zhang 8ec82cbc37 Generate synthetic dataset
Delete Open Sora Plan Modeling

amend log validation

Change name to fast video

Remove OSP modeling

clean up name changing

remove files

 Deleted unnecessary files

commit first

commit first

training

ok

Update generate_synthetic.sh and deepspeed_zero2_config.yaml

14

debug

overfitting ....

debugging ..

Debug successful!

Add zero3

OK

update optimizer

load

small bug

random seed args

sp enable & still has bug

rename

update inference sp code / can run / still has bug / don't output normal mp4

switch to deepspeed dummyoptim

fix some bugs; still output green

SP inference done!

fix typos in readme and latent dataset debug file
2024-10-28 00:44:02 +00:00
Peiyuan Zhang 54e74aec3f Mochi Inferenfce & Diffusers
typo

Original pipeline
2024-10-27 22:57:14 +00:00
Peiyuan Zhang c1c276b616 Refactor OpenSora sample_t2v.py and update download_hf 2024-10-26 21:59:22 +00:00
Peiyuan Zhang 8e5fa4d383 remove vae loss 2024-10-26 19:20:04 +00:00
Peiyuan Zhang 06ff1c912f Remove files in causalvae 2024-10-26 19:06:46 +00:00
Peiyuan Zhang 857f5df51b normalize 255; vae reconstruct 2024-10-26 18:57:07 +00:00
Peiyuan Zhang 18c5ca131d Add mochi download & Change output dir 2024-10-26 18:08:50 +00:00
Peiyuan Zhang c16625242e Delete merge_data.txt 2024-10-26 18:00:16 +00:00
Peiyuan Zhang fcc45701c9 Remove unused adaptor files 2024-10-26 17:59:04 +00:00
Peiyuan Zhang 6c87e003aa Remove unused arguments in train_t2v_diffusers.py 2024-10-26 17:54:43 +00:00
Peiyuan Zhang 23181aac96 Update T5Base 2024-10-26 17:53:51 +00:00
Peiyuan Zhang 26b65a8baf Update PyTorch installation command 2024-10-26 17:36:23 +00:00
Peiyuan Zhang 25498a9d85 Update model path and cache directory 2024-10-25 10:21:31 +00:00
Peiyuan Zhang afb3ac9fc1 Update t2v_debug_multi.sh with video_length_tolerance_range and dataloader_num_workers 2024-10-25 09:58:41 +00:00
Peiyuan Zhang 3d385600e9 update version 2024-10-25 09:22:16 +00:00
Peiyuan Zhang 5680039dbe Add environment setup and training instructions to README.md
Update dependencies; Setup code for debugging

Delete unused files and code

Remove all npu code

Update torchvision imports

Update EMA model and t2v_debug.sh script

Delete npu related stuff and remove inpaint module

Remove compress kv

Fix warning with dataset handling and model loading

Update PyTorch index URLs and video length tolerance range

Remove UDIT and inpaint

fix typo for dataset download

 include pretrained open-sora

Add pretrained model for OpenSoraT2V-ROPE-L

Update max height and width for video processing
2024-10-24 17:47:34 +00:00
jzhang38 2c1eb3b0e2 Remove NPU related code and update training process 2024-10-24 02:36:18 +00:00
jzhang38 5e2e3ab06a Remove unused scripts and update TODO list 2024-10-24 02:27:03 +00:00
jzhang38 f5ca624aff Delete unnecessary files 2024-10-24 02:23:33 +00:00
d1ea86d351 Initial Commit based on Open-Sora-Plan-V1.2
Former-commit-id: 7a3dccf2c738eaf535e62996211c8bc5f3de5683

Update README.md

Former-commit-id: aee63a3a8422e029f1550e47825ca4beb916874b

Update README.md

Former-commit-id: fb46c80c1a2663c09849472e93b2da6cd068a98c

Update README.md

Former-commit-id: 5fd1e6afe1277ccffdfeaf507477d944a6a35f03

Update README.md

Former-commit-id: 1cd7d8b5df34ab9f61d9346b2647b643f78a8f5d

Update README.md

Former-commit-id: e27cc5d49a19e9580ba1a29d948051b918aede74

Update README.md

Former-commit-id: 0ff9314bddef1dfbbf429726ed4240b823b221b3

Update README.md

Former-commit-id: f295f4623448dd4ffa9be747a5cc6733cf48c663

Create Contribution_Guidelines.md

Former-commit-id: 897dbd63dc6f53095e1dda3f10a86dd40e4bd263

Update README.md

Former-commit-id: c652acd1fc5a003b4d6eefcfb83dfb310f5b5e81

Rename data.md to Data.md

Former-commit-id: 27484f7a2d8634bd1cd9781a2f10afa7f304f5cb

Update README.md

Former-commit-id: 351be764d8648ac60cf170c3fe4e6ed9c544c90e

reformat code

Former-commit-id: ce23b280fbb419c08347ebcbf1ed4075bdf9af60

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

Update README.md

Former-commit-id: 1dd9c945be2a170d87d568bb00e1ec30e7be592b

Update README.md

Former-commit-id: 6ecaba8e03d7c7f99c068a57c923f11c5456a8c9

Update Contribution_Guidelines.md

Former-commit-id: dc87e459ca5a8a090800331111118acb65214f65

Update README.md

Former-commit-id: be7ecc3d56f7289052fc66d535038e7111f7ac84

Update README.md

Former-commit-id: 164e76cf496b1157eebea299ed58ca6c1799d70d

Update README.md

Former-commit-id: 3e3d674134498512b1ef1b0dd72d6e19afd6ef65

[feat]: frame_interpolation

Former-commit-id: 711fc4bba80d5f051f712e3460e2c0685117cf68

Update README.md

Former-commit-id: 1436d8feea7f82d3d747c632a68c47b3ea74258f

fix_readme_requirements

Former-commit-id: eb5a5c69d12cc412c12a4fea7f00baafd02701ca

Update README.md

Former-commit-id: c5a79c7e993d9e6a77c6c70efa1c0ac2bb3a063b

add latte

Former-commit-id: 01ba10e5d65c9d2d957472fbe7958e1b07a971e9

Update README.md

Former-commit-id: 437bc5c68b5c7a85865ad97d544c4818be45cb52

first commit

Former-commit-id: ec3481bc35a8120d519082dc76585428f0896470

[refactor] reformat videogpt, support training videogpt on accelerate

Former-commit-id: 75385ed47fe60008b61b1bcf89f61ea696d095bc

[refactor] adapt to old training code

Former-commit-id: 358d362ecbc100b21ea02895cb0426cf5b386a72

Delete LICENSE.txt

Former-commit-id: b10d8b2c6620e3a8412527bd03ea0adc26c981a6

Create LICENSE

Former-commit-id: 369fe085e91e55b9be6165da4a1eb7f9c0cf5a7f

Update README.md

Former-commit-id: f1aaeca19bc9e4304961c8c903f9bf03cb42fd9f

Update README.md

Former-commit-id: 0f2bdd2bcd819d528f2a41a96ba93f5c3c04dc8c

Update README.md

Former-commit-id: 10ee64c122f28efea39fc1a4542d278683a58ce1

option to use rebased linear attention

Former-commit-id: f1b39b34df5fbdef14ef3f6a0e6d0972a8570445

support for ring attention

Former-commit-id: 0533c73b99217e88da121bd0f68e0949bc88016e

add vae and fix bug in dit

Former-commit-id: 2dd4254c22ce671d5cd5e7d22a2386b7660fbdc4

Update README.md

Former-commit-id: 898d35f15ce0fd124cd48078f0a8af2764bf55d1

support latte training

Former-commit-id: d471a9414842f422d1200c2fd3f3ddfcf9e97fd6

Update README.md

Former-commit-id: 9fd88eb8ac5122b0999ccb0071f195aaeb0979a0

Update README.md

Former-commit-id: 398cfcd7b19cf1ba827eb49fc11223f72fa93a85

Update README.md

Former-commit-id: 673659495b7503ba402fbcb6c46204499761b4cf

[fix] fix a bug when using ucf101_stride4x4x4

Former-commit-id: 9911403c688321db731450d9383bc02923a4be2d

[refactor] remove test.ipynb

Former-commit-id: e73c3a9e0d41e51a0fd755e57739424213a76312

[fix] fix a bug when using ucf101_stride4x4x4

Former-commit-id: f541e723d65ea44ae4560c68908123c7e6781bbc

[refactor] rename

Former-commit-id: ec317605d24223527b8c2d4b4d2ee876f86abbaf

update

add RGT

Former-commit-id: a491111e48cec588095381eb70ffca94600ca4f7

add sample script

Former-commit-id: 9ce847651bb6dba28c51ce5807e222ae2902a78a

[refactor] configuration default values

Former-commit-id: d7ca7685d6a8c6d015082dda241bc98ab156770d

[refactor] added more methods in abstract classes

Former-commit-id: a535e3a06eb766bebe7baa3d9df64b04d04d40e6

Update README.md

Former-commit-id: 59a595a1e31b156500a944f244afec40988e872d

[fix]: disable gradient computation to save GPU memory when reconstructing a video

Former-commit-id: b7bf763e7f8800959ed82bb89f41684db216dd8e

support accelerate training

Former-commit-id: b32671c9991b9a9db66cb148822ccb3e8be7dc11

Update train.sh

Former-commit-id: 6c4957d715eaa489691be21dd0abdb9309660e0d

Update train.sh

Former-commit-id: 66aa3709f80d89a30d0a57c5e033ce1ae46f998f

Update README.md

Former-commit-id: 604a29d1de43e0fcb90ada4dd97f5f080e37caea

Update README.md

Former-commit-id: d097e40998a3396f74c2bfbcf26d6ace8080d0b6

Update README.md

Former-commit-id: 1a7178a61547c33a744d032302fb76ee0a3c6671

denorm fun

Former-commit-id: 05707a07b14295f5e1932d0c10bb84a20a0cd328

Update README.md

Former-commit-id: b3f6bb7f252505efb0ffd8195003c79206fac313

refine readme.

Former-commit-id: 391d9511b72aa9c09c7ad8992ae5df8c89c88c63

complement emoji.

Former-commit-id: 50679cc482d8ecf4e3083c0f5a791f195a58a855

Update README.md

Former-commit-id: 29f24a819fe2d4593c48afe6ca990659784bd776

[feat]: add sit model

Former-commit-id: 2eda854d6c25254b71d0105a3ed092953a9d6f68

fix: remove replicated diffusion module in sit

Former-commit-id: 997cb049185421d303b14d59ab6c0603f1d1a66f

fix: rename sit model filename

Former-commit-id: f3b2cd7dcf644f680616eee0571db90478fede23

feat: Incorporating SiT sample

Former-commit-id: 91e2931e412f8c96ef748cb2322ee956309fb6aa

feat: update sample

Former-commit-id: e7dc492b9b4764db28aa31cc1eede914052b4ca6

attn_mask with bf16

Former-commit-id: d9667a8c15a5344e3c777197fd6badec15dcddbd

Update README.md

Former-commit-id: a1cc0b26a3b156d7708e8df0cfe42aad33b1c946

Update README.md

Former-commit-id: 4179ed075e68d48094655c38dbde784836fd6e0f

Update README.md

Former-commit-id: 36e69275598b0cff38fdc5006926e422cda069b3

[fix]: disable gradient computation to save GPU memory when reconstructing a video

Former-commit-id: 06da17790b53cf593a0c264dda17c76f68d2a541

Update README.md

Former-commit-id: 8e9658ed993380abc43b518fcbac21a02d8ba9ba

Update README.md

Former-commit-id: 8b553a4477c64b5452f8424ca6987c0fb8c111ce

Update README.md

Former-commit-id: 0898e4b917f5e245c4394118511bcb7639973817

Update README.md

Former-commit-id: 4a893ff6177d59219b6755c4a89764ef54fc7f62

Update README.md

Former-commit-id: 89f9fc8eed3ca8cc2434ec227acbb845bfbcff82

[refactor] rename train_videogpt.sh

Former-commit-id: e2ea2caacaf5be046f48cf50a53ccafa7a3cc351

[docs] Modify README.md and add VQVAE documentation.

Former-commit-id: a5bfe70f8a8998db969d9c4938fda123aa153bc0

fixed dit forward

Former-commit-id: 4cd537fab12231f7471107762b02b36217ca42d8

videogpt inference

Former-commit-id: c5d617450a8271ad3b8a5e8f76ee2d77a822e31b

Update README.md

Former-commit-id: 4d83f1feab7d6872aca1c6b41be95950ec3cf465

Update README.md

Former-commit-id: da5507e570a27becab754ba24a0671d669b7e08e

Update README.md

Former-commit-id: 4a04e71636dabccbef446623674f3a514395b7f9

Update README.md

Former-commit-id: 2a286e5aab48aa68e7d3f04df23ec1cdfe71c6de

Update README.md

Former-commit-id: f4c4b3b4b1ce3c1d0b56ba5a109cd6cc57da900d

Update README.md

Former-commit-id: 7af6b82b29c265e53e9c48477083b4cff0e0a3d3

Update README.md

Former-commit-id: 7bc544af55c0a987d0df426848cb06cb330fd52b

Update README.md

Former-commit-id: 2a68a1aa399a53f58246df4f9f64dfdcdd130bcf

[docs] Modify README.md to fix translation error in name "CloseGPT"

Former-commit-id: dd36043df4b0408e592a82f9046f6d654a5375c3

[docs] Modify README.md to fix other grammatical errors

Former-commit-id: f0de27fcc38bb4dab0eeafac225a2924080a3218

Update README.md

Former-commit-id: a881ff5198e1125032418f4f02fd95d2179ffb3a

Update README.md

Former-commit-id: 8a191abc1903ad5998fe83149c9c9ef5c8226b9e

Update README.md

Former-commit-id: 3b4eb14dcd8862b8251b72f30cba523a5dc1def8

Update README.md

Former-commit-id: e458366e47d4b13737a71bbf866e20ec10961ade

init pos embed

Former-commit-id: 28f8076136ec214c66039ead7547c9b25cfe7bfc

Update README.md

Former-commit-id: a1d1f07c83eb0c03d93b023b5c9162a6a3c26cd9

Update README.md

Former-commit-id: deddba1f03bf697fb1252a0ce4fdb5df5bcae8de

Update README.md

Former-commit-id: 076e6e70e717888f24265d624c6a7a8ffa2f5d1b

Update README.md

Former-commit-id: 7cd398721e556db6199f1d06c7a0d1de2a9720dd

[updata] updata video super resolution

updata video super resolution

Former-commit-id: 82397545859791d112d6f0569e33fe23437563df

train script

Former-commit-id: 3fd1ccf099c2deed20af50bf9556e2fc3b7a4382

Delete opensora/models/super_resolution/placeholder

Former-commit-id: 8dd5a6a3c85dc28989ff95e6ad41be7fdb0f0811

Update Data.md

Former-commit-id: 9531e33666f842125b331cb266eb26a3a03a4009

Update Data.md

Former-commit-id: 86be6b2bf5a9be1030bb3af661a2ec85e710a5c3

commit

Former-commit-id: 81359cf67248b6c594ff79f121a8a38558c1cc70

Update README.md

Former-commit-id: c4d7aafa3ceb6e4d48bf85b402ae2b2ce9394e53

Update README.md

Former-commit-id: bca7070fac79d5ae5c37830cd55ca1e81a289197

Update README.md

Former-commit-id: b37da2202e730a4930a194a80a703b591a535adf

Update __init__.py

Former-commit-id: 03a45aaf60d294b7a68d8fc44132d3d671e4472b

Update README.md

Former-commit-id: 97fe914aaa2b9bd7c2ff5d4239381631c47e8356

Update README.md

Former-commit-id: c6b1ea8f1ef01096aa8c3e5272bc9fb255b5532b

Delete scripts/train_vqvae.sh

Former-commit-id: 952a4ce3a3b0da159d599ff3a0a052265f48c2dd

update script

Former-commit-id: 43f539b68af150c45040d79918a84de1dce47c33

feature dataset

Former-commit-id: 9683caae8648735552270b871de22fc540ba2b55

update train.py

Former-commit-id: 844173711342f2dcc021294e87971ca39854bfe4

del use_fp16

Former-commit-id: 5424c0112e04b56d1df30ac76cbc104e848406a3

update sample.py

Former-commit-id: 10d8621051d1510270c6a55abe686eac86297170

fix extract feature

Former-commit-id: d4dce5ec1aaa058e452a41f15cd35fe1713f0b78

Update README.md

Former-commit-id: 70f6ebf421f0830a51dcc2838c2c3783cfb67376

[Updata] updata the video super resolution

updata the run.py and add README.md

Former-commit-id: 1e6d6ace7a0963149e228813801210b022aceacf

add CI for docker autobuild

Former-commit-id: c4e2d5ff1dce8662a42f1ed6a7786855d40173f7

[refactor]: add docker development support

Former-commit-id: fded3baff7915c0cfc9cec3765c271259359831c

update docker scripts and configs

Former-commit-id: 4567e6368f1bf413d8f6319a97a852b8289ca770

Update README.md

Former-commit-id: 99b0066895af072a8887737f2b603ec7fc893ccd

extract script

Former-commit-id: 4139f2f68ae09695302ce49afdb7318df83cabfd

[docs]:update the eval_code for calculate the FVD, clip_score, ssim, lpips, psnr

Former-commit-id: 1ee226c338126fbdc5178020b12271b980a7ce0c

[docs]:update the eval_code for calculate the FVD, clip_score, ssim, lpips, psnr

Former-commit-id: b504623d7ec2ac0b508530d21ffc88978f8a4f71

[docs]: update the evaluation

Former-commit-id: 02e3789b74daa036fbbff6d13dfc7e98372260ed

Update latte.py

revise the xformers inputs

Former-commit-id: 60db66c912192c3fbafe55ebe16157514b884e5b

Fix the memory error bug when loading metadata pickle file of vqvae dataset.

Former-commit-id: 97ab243bd12915cea7aa2deecfac4b99417f78a3

Update README.md

Former-commit-id: 8f7b41118a3f491d1d5b5789c53790c46b023191

Update pyproject.toml

Former-commit-id: 36da5d0e950625e629aebe954f89f92827397d0d

Update README.md

Former-commit-id: 826cc5e9162888bed82413088c6be662c278176b

Update README.md

Former-commit-id: 88e9f21cfc8ec775dbc9a65f740043e91e26bd91

Update README.md

Former-commit-id: 38e0e290e7a5909706c2cce2e4dd047647134fb8

[docs]:update the eval code

Former-commit-id: 937f2fc87d34ab9d4724cb95a6457b54a6bf4ebc

[docs]:update the eval code

Former-commit-id: a8c9e6cb9e39a9e929dab34077280a8d80d85573

flash-attn train and sample

Former-commit-id: e336565bd5ff2aaf4d0b4ad5f42732e89cd216b3

[refactor] add static type checking

Former-commit-id: 12c4240dc28888c61f1f79243627e7cdb445c840

Update README.md

Former-commit-id: 513883df1ad853c9a5e7457c364227261df17a83

Fix einops bug: module 'keras.backend' has no attribute 'is_tensor'

Former-commit-id: 849f8ed815295bac78c3c10d319e6731aa26ca02

[feat] Add 2D RoPE

Former-commit-id: ccae80d339e34fb23706827565b5095ae2c2f320

Update README.md

Former-commit-id: 6630c22baa2340a8b55badd61472daab7177bc63

xformers and flashattn

Former-commit-id: 002cad95253f7da36344237ce8c8410ffb517711

safe dit

Former-commit-id: f407cab7dc0e53257972d430e873f4fd77bb039d

[fix] 2d RoPE init

Former-commit-id: ff5931bae6efa94675b9ef6bc25f1b76217d128e

fixed attn_mask in flash-attn

Former-commit-id: 5e2504f429a342b75bc7a9004cee7b3f75582ee3

typo

Former-commit-id: 9e87366e906ed19726c8905afb27b5df4d4e1512

Decouple mixed_precision and gradient_accumulation to command args.

Former-commit-id: 903d473d1432eb492a87efe184d03439f643b92d

Fix mixed_precision bugs of latte.

Former-commit-id: d70a551972e04d95444d46cdbe2bfb16ad099fcb

Refactor latte.

Former-commit-id: 304534269f105578097bc16ab497636ad59fd9b2

Refactor latte.

Former-commit-id: 7408fd0ee16247afc3d4ba43b8c8e7dfb1091a97

Support deepspeed zero2 and zero2_offload for latte training.

Former-commit-id: c87e95dd357e6fa6d6fb2dcfee781bc49c039e13

Add some comments.

Former-commit-id: d5b5bd7743edde6a63dc74457bded0f8c56be1c2

Fix typos.

Former-commit-id: c10330d373466f3c1ab455bfb391569b9175bc3a

deepspeed

Former-commit-id: 6179a4b0e7d06caed4201a27679abc64615febe6

Update README.md

Former-commit-id: 582070d4f97c4a4398dcb9a6699f4ab6e980c812

Add base config.

Former-commit-id: d37c551a0421d6b73dee26060b517c46d8217e4f

train 1080p video

Former-commit-id: 74625142ef892f9e7fa459b2f98f2f7ae5b36d6a

Update README.md

Former-commit-id: 87636c1dc2275c4da77c6368a0f4425bea15dc12

Update README.md

Former-commit-id: 059d9e863958cfd679dbb8d365de6cd933e284fc

Update pyproject.toml

Former-commit-id: 5c379b6e4515ad75b89cdebbc5cd490c9a821d78

Update README.md

Former-commit-id: 7a62dec060c91b02ca556451266e6f3ef28e03e7

clean script

Former-commit-id: 02c61f63a8539a31516b92f5bb7c2d2e428ca7d7

Update README.md

Former-commit-id: 6d41cfe906f97dbed1cea8c9ed06a1f6dc84308e

Update README.md

Former-commit-id: 64d147863368aeef479ad25f2e5463fdcb76de46

Update README.md

Former-commit-id: 2808edf11005e04aa730971a786d1a4e3bca88b5

Update README.md

Former-commit-id: 1f94f2721f35670aeae8dbfb7c0a5abd06cd2aee

[feat]: Caption Refiner

Caption Refiner for Video Caption

Former-commit-id: 97c16bccb9979acf19fb91ae8d7b84ac4a3b4d97

[BUG]: Shorten the refined caption

Shorten the refined caption using GPT-3.5 summary

Former-commit-id: bb159f0d5fa928258475b2058983397ceb91af9d

Fix resolution typos.

Former-commit-id: 3ad55d28eaf13f9c8fd23912d2691965737527e9

clean dit

Former-commit-id: 117ec4196a6f2cd83da223b993b14e28f7a693f8

Update README.md

Former-commit-id: abb8ed57a29b1f66adec66d854429f36f9199540

Update README.md

Former-commit-id: c2c64665d4f281e264d87c20cc1cfd38e7e5f168

[fix]: correct attention_mode to attention-mode

Former-commit-id: 9957e2eebb2166b02600c6052fb54908f5d12bad

Update README.md

Former-commit-id: 266f9f97ea8f327fdefb2ca959c23d08b7084137

Update README.md

Former-commit-id: dd00e04220ff5232c5c4503a095e409471476a51

Update README.md

Former-commit-id: 5717aab628fceba311e284d1bdddeacc4c43c7c5

Create train_256.sh

Former-commit-id: bf62151617feb49663f8dfe0c78aa01f41e62092

Update README.md

Former-commit-id: c5eea8d3492872255a6a34d422230509d1c3531c

[feat]: use accelerate on multi-node

Former-commit-id: 0b12fc31f1d2a63f9da38821174ab63d9b3a2d72

Update README.md

Former-commit-id: 0d7dd6233811ec3c8741affb436d8004d8c676b8

Update sample.py

Former-commit-id: 6e09fc266a626f465109abc199d9fb452657929e

Update interpolation.py

Former-commit-id: e10e3416433e5f396aa13e7d8beaf451ff47a383

Add files via upload

Former-commit-id: b707b1968fa59d85891040604e7d54bef61b4549

Update and rename Frame Interpolation.md to readme.md

Former-commit-id: 87bdf1ebecb002c086497b3c47acb3f3cb2aaf34

Update readme.md

Former-commit-id: e287f94838007f2eda434a46bcd250099b627097

Support deepspeed for videogpt training.

Former-commit-id: 1f3dfe038845d91b81b0ecc9b6f3cabb06eed5d5

Fix config typos.

Former-commit-id: 455973b57bda19379d41e59aa418421fd269281e

Add hidden_size argument for videogpt config.

Former-commit-id: 1be9f22767fe9c80ad90cb17184751b97aff9012

Add deepspeed script for videogpt training.

Former-commit-id: 3bda2a05cb789fe785b2355a613385b3eb7abbbc

use fp16 temporarily.

Former-commit-id: 93381968e10aae7bbec23ec1ef40e7cb311fddf0

update videogpt.

Former-commit-id: 2e211563c684a8b6ef62dfc0c4f1897224dfb58d

Fix deepspeed zero loss bug.

Former-commit-id: e24af415bc5d6fd38d72a24d6312f4e65d3627dd

Can't fix zero loss bug.

Former-commit-id: 1f32f9d889743b7961b34b823b276bfd03974191

use imageio for write video which has better compatibility.

Former-commit-id: 005119dc9f78ea2a611d22a53772fd59d97d7139

Update training arguments of videogpt training.

Former-commit-id: c9660af891c0dae1f51a1075ef5b2df012f65e06

Update README.md

Former-commit-id: 8ec7524f4f171e650d7a25cc88ed29f7b5017454

Update train_256.sh

Former-commit-id: a2dfc9a8c70845446e731c88acd222e9b2db9122

Update README.md

Former-commit-id: 7e74356a3279b02d062bab9c9f361f61eb756179

Update README.md

Former-commit-id: cf1c56c0058288427c1dfd651c8c0e9a981c5653

refactor dit for supporting deepspeed.

Former-commit-id: e1909745cb05ea2f28f3a391c194f7656593a236

Fix typos.

Former-commit-id: 6add5ee7ef39a973e91b7a7c1b3a4b0980130b88

Update README.md

Former-commit-id: 2d5b6815f95e94a98847b8a0376473f01bb15d14

[bug] fix quantization loss

Former-commit-id: 111fcf55f969ffe4c9be284a317ae5712570a527

del t2v

Former-commit-id: 910398ca297be8a188d59c59c121857666774af1

Update sample.py

Former-commit-id: eb906691abec030d7aa2078b63fc9c7d678662ed

Update sample.sh

Former-commit-id: b149a300f1aff53d9ab6e66a91cc2bcd37900347

Update README.md

Former-commit-id: 7b96fd89eaf3743a4a9fbc995e6a2a791c734f2e

del t2v

Former-commit-id: 5daa2952997a29e74e2b5d235001d6d9dabf9ff6

[fix]: fix sample.sh

Former-commit-id: 09916dfcfb0b80a4217016919e0394a2cb665cfd

update scripts

Former-commit-id: 8b2a794095f2deb9b1e8090486cb8242bf81031d

Update pyproject.toml

Former-commit-id: 7fb921440b608e47781d2f541b2a8d893a48193e

Update README.md

Former-commit-id: 34e529fd02712d4466ab2bc1a63f9bd69976b6c3

Update README.md

Former-commit-id: d97edd0e714a237a54f191dcebb6fd39f73185d6

resume and compress kv

Former-commit-id: 0af7d58f88c1b65b4089da6f3e5c70456614da7c

refactor and fix resume bug

Former-commit-id: 1b1136671b9a7e0b15b70168e2c44dbff04a22e2

Update train.py

Former-commit-id: 783faf5cd86c148a713ea986ceed8527b108f2b8

Update train_t2v.py

Former-commit-id: d25b39ffac4474abd40b551fe0e51c421cb7bf5e

Update README.md

Former-commit-id: 6729039f059a23bf71faa944e2498657ef2c7442

[docs]: fix VQVAE script path

Former-commit-id: d5f5c681b8c50622a84256881d193d5803ff3c28

Update pyproject.toml

Former-commit-id: 5637656a79faba1b397b56918c42022cc209e731

[feat] add casual vqvae

Former-commit-id: 89087970b267908c8667984aba42b1ade2505467

Update __init__.py

Former-commit-id: 7c8af2909c4549e96bc0893d09fcc53d71fb6cc7

Create clip.py

Former-commit-id: e3397c567efcfeef6f58dfef2eb4251bedbbf33a

Update clip.py

Former-commit-id: 1ee53d152630b6acc69369be33287abfb143224d

train with image

Former-commit-id: 209b9a8ad2f19a50e46f9b7ec68dd781898e491a

support attention_mode

Former-commit-id: c72f354fbfd948e8cd4b0cc63c157bc4688c7665

Update README.md

Former-commit-id: 9065d12ee867c76452ce370fb71cc83bdc01470b

Update README.md

Former-commit-id: 95b28706d257d7769049abe30a7121b60e39a53f

t2v attention_mode

Former-commit-id: 4fd8bd87e0d0a98d5d0a3201d5db9a5ba055b894

Update README.md

Former-commit-id: 859c0f60dac1562b8a6bbb72b793a4709dbc3597

Update Data.md

Former-commit-id: d1c72e3209769d0204efb8e7d705bdf4f22b82ea

train t2v feature

Former-commit-id: 7f16b162cae0e3533c8ba9aa03186a846d9e02bf

scripts

Former-commit-id: 047046f684c2831e68e0cf256cb46fc3790cda33

add causalvae

Former-commit-id: ac936276fd3213ee1815dfe096cd64cecb601428

Create causalvideovae.md

Former-commit-id: 98414c2ad7dd39f740440b2d9059140f93492e83

Rename causalvideovae.md to CausalVideoVAE.md

Former-commit-id: 0a78e4193fcaec4662b7c6465a9d80643b354cd6

Update README.md

Former-commit-id: 8d70a890fc4191bfc98c3aca6dd3c8a70534fb3d

Update README.md

Former-commit-id: a41674f5ef71ced59b374226dddfa09c15a99f25

Update CausalVideoVAE.md

Former-commit-id: 1cd665f699632e9dc9357379225c085d9c1278d5

Update README.md

Former-commit-id: 90d6d298eaee6255098b342b31444e7a347fc50e

Update README.md

Former-commit-id: f294a0a28f288ae592e7de004320d3eb5e48e945

Update README.md

Former-commit-id: 02b0422fd745b4b7ec091cc6ee81d1d237245fe6

[feat] eval

Former-commit-id: d2c3cf5ef8d18feb7b0544d68d9796656e90ba8d

fix training bug

Former-commit-id: 09a6283ddd9ab11479f30324de788bd7415702ae

update model

Former-commit-id: e4e24650376d9dce974b290752efb416d5061983

update vae

Former-commit-id: 24298df5d374d5aaa5696f6200d4e23ae887715f

fix videovae trainer

Former-commit-id: 79e1feb412b6eeead3306f7f0d55a01f9b2a543e

Update CausalVideoVAE.md

Former-commit-id: fcde4d0ff783448c2006d6f4751a0e4f9cc23d0f

Update README.md

Former-commit-id: 2d9ae56309717a91e0690bd5c354fe3cbe5d5ed6

Update __init__.py

Former-commit-id: 2bdc29d86ebff3d6f29c38ded57c7b0dd1e6bbdb

update sample

Former-commit-id: 116726b9d565c689496123d4b6e62b98400536a6

tile conv

Former-commit-id: 5133615bcba221c4b86ae4e8b09b82ee5d34c66d

fix reshape bugs in AttnBlock3D

Former-commit-id: 299622e168a6cbee14ef2056c47cae722c20e005

tile only2d

Former-commit-id: 7f31ad03fdbaac97aea51a98bd29dc7327b700c9

released v1.0.0

Former-commit-id: 8df5897a4ffc341ced45faf617687fb8ebf1286c

sample pipeline

Former-commit-id: 8964b16c0dd65445a20050c9a0e82d909b55eaff

add tile training

Former-commit-id: 216079dc925beefe1ec722a4e87e46a4e1ca56be

update train script

Former-commit-id: ac16f77d19488cd46130587369cd7a5ea378052b

refactor vae to hf

Former-commit-id: a938d3f7f962709f7ecff5e9d49c9bf1ce40d7fa

release

Former-commit-id: 52fe2ac6610661d6a13a08390c84a8237f37986a

v1.0.0

Former-commit-id: 93f44f3b425c090d9bc92687a9c50a26b4f0d7c8

clean

Former-commit-id: 56ed87088e13cdd93929e2ceedbfae693e2f013a

update gradio demo

Former-commit-id: 80c2cc0bcc2a7532bc2440df16c1f811b0b30516

Update and rename CausalVideoVAE.md to Report-v1.0.0.md

Former-commit-id: 3b848f8de33cf4f0ed802428e9ad31d2fa26de63

Update README.md

Former-commit-id: e6fb5d864ac8af8214cdc89c71b42f296dda14e7

Update README.md

Former-commit-id: a61200bea2eea4fd22bd8095f192fef4c22c49b4

Update README.md

Former-commit-id: a2ec90a19eacb099cb9f36c1b5dd6e24be307fae

Add files via upload

Former-commit-id: 1ae8512b7846e9faafdb38ca532898d7f1bdcce2

Add files via upload

Former-commit-id: 292cbcfa4e4f384fe6958badc27d3067173119f9

Update Report-v1.0.0.md

Former-commit-id: 92d906b340f5970a77377586d7295f46b5de99ec

Update README.md

Former-commit-id: 037d7d3e084d7ff426347db49c4b210a2ee3927d

Update README.md

Former-commit-id: 7d4118566b77e09de4a3615612a5b368c76c293e

Update README.md

Former-commit-id: 3210474a46845ae3a6c0bde7fb6c7f9845866628

Add files via upload

Former-commit-id: a232ec2d98b245a9b502a2c49fc4f468a86fbd64

Update README.md

Former-commit-id: d5ab411b2e2fd5d6abeda500895e3339167572ab

Update README.md

Former-commit-id: 30f61958e221c3d41c4aa5917f4511e8c0482cf9

Update README.md

Former-commit-id: f8cb7f38751d7a6ef586c293fb769f5d60ad5c59

Update Report-v1.0.0.md

Former-commit-id: bd9aa7991212d98a6224b855a3bfb2c12feb5950

Update Report-v1.0.0.md

Former-commit-id: aea1f4a82a525d3e6e1d1a1bb7631326b294cdcc

Update README.md

Former-commit-id: f709d75e7a7d8a08ec60132aa906f6b880af8e58

Add files via upload

Former-commit-id: 9a207435d9efac0b19771df5d2288160af1bdd6e

Update README.md

Former-commit-id: 9bae31c73ea902581e1ee94cfd9012efdf78d349

Update README.md

Former-commit-id: 9f926abb08f156807091dd044c133788649a6bac

Update Report-v1.0.0.md

Former-commit-id: 7cec40420601f7da57449586508e50119bd02a61

Update README.md

Former-commit-id: ce2bbf82045041495ed448b46d5f4bdba84aa740

Fix some compatibility issues

Former-commit-id: 942d8598b28516369780c3519607929869b427b7

Fix bugs in inference

Former-commit-id: c251a86ac7c91d6f2bb12ae645628e9cf4b83d3d

fix bugs in inference

Former-commit-id: e17c7a7b7704920cb86894419bf573d0fb0c06fa

Update README.md

Former-commit-id: c1492211ced3cdbec93ef28678a1467245e122e9

Update README.md

Former-commit-id: c4a976e9b93536299b2e350454c780fec54fbe95

Update README.md

Former-commit-id: bbf3cef43a129a25a80ce134a8895796e24353c6

Update README.md

Former-commit-id: 5731e810e2efb7ca7eda1aa8954eb7fd260f4250

Update README.md

Former-commit-id: 1b54444249fa91b29334b8fc364fb1978c25fb2d

Update README.md

Former-commit-id: 9c0f82ae9953b56f2540e9c7ada021e4e1bc6165

Update Report-v1.0.0.md

Former-commit-id: 38b779022418cc41cdf9fc508e461f2f6299fd60

Update Report-v1.0.0.md

Former-commit-id: 6ebe2c62c0904142026d3796027e95ce43afe16a

Update README.md

Former-commit-id: 2873d0eaedfad392c627a714448118949d7b7e13

Update Report-v1.0.0.md

Former-commit-id: 400f708a080f0799e02697bb8e66d8430ffb1f56

Update Report-v1.0.0.md

Former-commit-id: 37f181f34f64ca1141a1345b7c1f2af283be1cad

Update README.md

Former-commit-id: dc923fe2413b3be31d8d88e7dd83a931fcbf9458

Update README.md

Former-commit-id: 828be4267f0795969293e4a321e4e0bfa8e1f1e6

clean

Former-commit-id: a7f2c1c8b2243587cedd6e152e64298842287095

Update pyproject.toml

Former-commit-id: bb037420b75303b58486ba79da9180b2bf63cf88

Update Data.md

Former-commit-id: d9a5c8fbd36a34b7baf720bc27d69604ae5c9d57

add causalvae doc

Former-commit-id: 4e6d9e038c8edb81f8402834bd01eeeb823924f9

Update Train_And_Eval_CausalVideoVAE.md

Former-commit-id: 589657b089930379bb75f3d49063ebb72671efc6

Update Train_And_Eval_CausalVideoVAE.md

Former-commit-id: 156739c35b5dda594138b2e3b1fd6ec99556af23

Update README.md

Former-commit-id: 567ade18047a9cc81648997c5d9287fa8900614d

更新 README.md

Former-commit-id: 42607786f3f3d0eba9336f774b0a2ae076cce793

Update README.md

Former-commit-id: 88e7f2062c494c199caff468f5f14f3165f763de

Create Report-v1.0.0-cn.md

Former-commit-id: d484754e080bd65f6b35619b840040581fa6b69c

Update README.md

Former-commit-id: e386ec112337271bb9f55ea6ec872eaba94aff50

Update README.md

Former-commit-id: c437a18dbec7ac6e02287565a84d75f96664ee99

Update README.md

Former-commit-id: 0cc27a5b065b61f37d44dc2c90ba2782075eefb2

Update Report-v1.0.0-cn.md

Former-commit-id: d5cf6aed7c1673562a6d2022c3fd7f7b16006cfd

Update Data.md

Former-commit-id: 15defc1a3846335fd68e18fa9b8ffbe38e3e7100

Update Report-v1.0.0.md

Former-commit-id: 3c8eb75dc2045154e8654f860b1fabd7e597094b

Update Report-v1.0.0-cn.md

Former-commit-id: 0edac24fcaa447a534755bf3aac1f2027a491c06

Update Report-v1.0.0-cn.md

Former-commit-id: 59edce70edbaed957ba40f1d95b4a5050ece12bf

Update Report-v1.0.0-cn.md

Former-commit-id: 83124efeaa5a8cd528a15760bdf88acc403142bd

Add files via upload

Former-commit-id: 943dd9a063baaa05d7f0d0bd52cbde35cf2b3d98 [formerly 5542359738243a8e948a8f08ed4e8dedf6b4b818]
Former-commit-id: fbee07f6fef2ac86bc15ff66f2e2a1a68b64ed01

Update README.md

Former-commit-id: 6fd60c73b6ae69d6ad9dc14c7371bd9bff086b44 [formerly 2ed5d5b7bdcd88044da26ca94237c50af7ab47dd]
Former-commit-id: 443c333353963bc1ec64a4946afbe5e17f9b5a87

Update README.md

Former-commit-id: 54bad83bf054a2ff574199931121c686b4e28008 [formerly 06bb190a39a9b98c8de5e385fc304a86cc8c9fb2]
Former-commit-id: 23e8a235d6828d9ab7c02ef59d40fe13a46fa4b1

Update README.md

Former-commit-id: c48eba0bb15d594007809d4cc605dc46ba029cec [formerly 7deb8a0decf7a9757af4475e3029d415d78d185f]
Former-commit-id: 7f43bb2f1d9308bd08154171d78e9b36e00a837f

Update README.md

Former-commit-id: 68eaf2e6b00266cf0858b33f9b4d09681efbd2bc [formerly d8dcad250ed8357964ab8ee570b9973b45e29fcb]
Former-commit-id: 75f51b3d84203be3672f6b07bb3c70ba0ae5556c

Update README.md

Former-commit-id: be00a17729ac889493f63fd1d649054d7980ca42 [formerly 9b4d8dcfc8e6e2abe4ca727e01561d27d1ba3968]
Former-commit-id: f1b6f5013b99d0d288af0f05c46cb8300d79a8aa

Update train_videoae_65x256x256.sh

Former-commit-id: b2f5cfe72452dfb369aba166e4a101dfcf5758c9 [formerly 557b78ab38a61d7520e91044137f818e50ec24cb]
Former-commit-id: 0e46c9f261e7f5fd4d08e3d7735c17c7975b0e72

Update train_videoae_65x512x512.sh

Former-commit-id: 8e40db106efa2fd879cca42bd9ce2d9b79d777f8 [formerly 3e11781b045a838a5aa98ea75e274198f2b26a44]
Former-commit-id: f3b0804a7cf063e7ff04ee57dccee0320c1f32a5

fix using pretrained bug

Former-commit-id: 2dfbb0683b2f3bd1a1d619e5e4fabb5bd4c92d79 [formerly 27f352e822f732cb351d64ac5726c5df3b62c1bb]
Former-commit-id: 722d60ae5ef28499c47292e0a4807a87b9658b01

Update train_t2v_feature.py

Former-commit-id: 4a4f7ea2366f1dff5669d6d4834238ebd790e6e2 [formerly ca7d3158fb01f8e7228f833cea8e6741be127c2a]
Former-commit-id: 389f0e7dc66826026046b8e12d7a9f8692d32e42

[docs] update docs and train.sh

Former-commit-id: 1d9358e84d12863681b396c7621ae7244caa4dc2 [formerly b801491889b3abad1211225d7f07af3ae0389895]
Former-commit-id: 9e6a59ae2f246b50ace0a033cbd6548660241204

[fix] fix bug in #185

Former-commit-id: d74efbd0195380ca383598c06d3d8185dd9cc9bc [formerly 0f2687266bc2271e11b2f06edfe514f466e57b0f]
Former-commit-id: cb6bb234befc49b3870c5690653469a0714890c5

[docs] add evaluation detail

Former-commit-id: c15a717cc61159e5f6adafb718237ab578ac81e7 [formerly 9e92772776fa4dd18772c7fec939a42ad6dfbc0a]
Former-commit-id: 98c2a38af5064f27888d1b102c9fded87d2420d5

[docs] update train.sh

Former-commit-id: eeb00ac046692bc97a064ec1b47c994436798ab1 [formerly 6d0fd85e1b13a7824b8181ba9ccebec2a3c76d40]
Former-commit-id: 1e8e669ad9967aeeeb7622bc04a0759f348cb26f

Delete assets/we_want_you.jpg

Update README.md

fix typo

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`

[fix]: Fix variable naming errors

Update README.md

[docs] update inference example.

[feat] add time chunk inference

[docs] fix typo

[refactor] fix hardcode

Update LICENSE

Update Report-v1.0.0.md

Update pyproject.toml

add 2drope and dynamic training

refactor dynamic training

read image from folder

img training

img training

support newvae

vae temporal tiling

abs and rope

compress kv and rope pi

fix rope with compress

update dataset

mask loss

multi-data

fix dataset

train vis

update

fix mask

5.15

prepare release

prepare scripts

fix demo

5.27

update prompt

Update README.md

Update README.md

Create Report-v1.1.0.md

Update README.md

Update README.md

Update README.md

Update README.md

Update README.md

Update README.md

Update README.md

Update README.md

Update README.md

Update README.md

Update README.md

Update LICENSE

Update Report-v1.1.0.md

solve typos

Update pyproject.toml

Update README.md

Update dataset_utils.py

Update modeling_latte.py

Update train_t2v.py

Update train_t2v.py

Update t2v_datasets.py

fix wrong code

Update LICENSE

fix train bug

fix vis

[docs] update CausalVideoVAE docs

Update README.md

fix the bug

Update gradio_web_server.py

Update README.md

Update gradio_utils.py

Update sample_video_513.sh

Rename sample_video_513.sh to sample_video_221.sh

Update README.md

Update Report-v1.1.0.md

release v1.2.0

Update README.md

Update Report-v1.2.0.md

Update README.md

Update pyproject.toml

Update Report-v1.2.0.md

Update Report-v1.2.0.md

Update Report-v1.2.0.md

Update t2v_datasets.py

Update Report-v1.2.0.md

Update Report-v1.2.0.md

Update README.md

Update Report-v1.2.0.md

Update Report-v1.2.0.md

Update README.md

fix sample

Delete opensora/train/train_t2v_diffusers_lora.py

lora bug

Update README.md

Update Report-v1.2.0.md

Update train_t2v_diffusers.py

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>

Update Report-v1.2.0.md

add 29x480p link

Update README.md

Update train_inpaint.sh

do not set seed

[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-03-04 14:06:39 +00:00
127 changed files with 8492 additions and 6436 deletions
+197 -17
View File
@@ -1,21 +1,201 @@
MIT License
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
Copyright (c) 2024 PKU-YUAN's Group (袁粒课题组-北大信工) and Rabbitpre AI
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
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:
1. Definitions.
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
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.
"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.
+141 -87
View File
@@ -1,108 +1,162 @@
# Fast Video
This is currently based on Open-Sora-1.2.0: https://github.com/PKU-YuanGroup/Open-Sora-Plan/tree/294993ca78bf65dec1c3b6fb25541432c545eda9
# FastVideo
<div align="center">
<a href=""><img src="https://img.shields.io/static/v1?label=API:H100&message=Replicate&color=pink"></a> &ensp;
<a href=""><img src="https://img.shields.io/static/v1?label=Discuss&message=Discord&color=purple&logo=discord"></a> &ensp;
</div>
<br>
<div align="center">
<img src=assets/logo.png width="50%"/>
</div>
FastVideo is a scalable framework for post-training video diffusion models, addressing the growing challenges of fine-tuning, distillation, and inference as model sizes and sequence lengths increase. As a first step, it provides an efficient script for distilling and fine-tuning the 10B Mochi model, with plans to expand features and support for more models.
### Features
- FastMochi, a distilled Mochi model that can generate videos with merely 8 sampling steps.
- Finetuning with FSDP (both master weight and ema weight), sequence parallelism, and selective gradient checkpointing.
- LoRA coupled with pecomputed the latents and text embedding for minumum memory consumption.
- Finetuning with both image and videos.
## Change Log
- ```2024/12/06```: `FastMochi` v0.0.1 is released.
## Fast and High-Quality Text-to-video Generation
### 8-Step Results of FastMochi
<table class="center">
<td><img src=assets/8steps/1.gif width="320"></td></td>
<td><img src=assets/8steps/2.gif width="320"></td></td></td>
<tr>
<td style="text-align:center;" width="320">tmp</td>
<td style="text-align:center;" width="320">tmp</td>
<tr>
</table >
## Table of Contents
Jump to a specific section:
- [🔧 Installation](#-installation)
- [🚀 Inference](#-inference)
- [🎯 Distill](#-distill)
- [⚡ Finetune](#-lora-finetune)
## 🔧 Installation
## Envrironment
Change the index-url cuda version according to your system.
```
conda create -n fastvideo python=3.10.12
conda activate fastvideo
pip3 install torch==2.5.0 torchvision --index-url https://download.pytorch.org/whl/cu121
pip install git+https://github.com/huggingface/diffusers.git@76b7d86a9a5c0c2186efa09c4a67b5f5666ac9e3
conda create -n fastmochi python=3.10.0 -y && conda activate fastmochi
pip3 install torch==2.5.0 torchvision --index-url https://download.pytorch.org/whl/cu121
pip install packaging ninja && pip install flash-attn==2.7.0.post2 --no-build-isolation
pip install "git+https://github.com/huggingface/diffusers.git@bf64b32652a63a1865a0528a73a13652b201698b"
git clone https://github.com/hao-ai-lab/FastVideo.git
cd FastVideo && pip install -e .
```
```
pip install -e . && pip install -e ".[train]"
sudo apt-get update && apt install screen && pip install watch gpustat
## 🚀 Inference
Use [scripts/download_hf.py](scripts/download_hf.py) to download the hugging-face style model to a local directory. Use it like this:
```bash
python scripts/download_hf.py --repo_id=FastVideo/FastMochi --local_dir=data/FastMochi --repo_type=model
```
## 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)
Start the gradio UI with
```
python fastvideo/demo/gradio_web_demo.py --model_path data/FastMochi
```
We also provide CLI inference script featured with sequence parallelism.
```
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
export NUM_GPUS=4
torchrun --nnodes=1 --nproc_per_node=$NUM_GPUS \
fastvideo/sample/sample_t2v_mochi.py \
--model_path data/FastMochi \
--prompt_path assets/prompt.txt \
--num_frames 163 \
--height 480 \
--width 848 \
--num_inference_steps 8 \
--guidance_scale 1.5 \
--output_path outputs_video/demo_video \
--seed 12345 \
--scheduler_type "pcm_linear_quadratic" \
--linear_threshold 0.1 \
--linear_range 0.75
```
For the mochi style, simply following the scripts list in mochi repo.
```
git clone https://github.com/genmoai/mochi.git
cd mochi
# install env
...
python3 ./demos/cli.py --model_dir weights/ --cpu_offload
```
## 🎯 Distill
## 💰Hardware requirement
- VRAM is required for both distill 10B mochi model
To launch distillation, you will first need to prepare data in the following formats
```bash
asset/example_data
├── AAA.txt
├── AAA.png
├── BCC.txt
├── BCC.png
├── ......
├── CCC.txt
└── CCC.png
```
We provide a dataset example here. First download testing data. Use [scripts/download_hf.py](scripts/download_hf.py) to download the data to a local directory. Use it like this:
```bash
python scripts/download_hf.py --repo_id=Stealths-Video/Merge-425-Data --local_dir=data/Merge-425-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 ../..
```
## 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不行
Then the distillation can be launched by:
## Experiments
Scripts are located at scripts/experiment_N.sh
1. pcm_linear_quadratic, euler_steps 50, 0.025
2. pcm_linear_quadratic, euler_steps 50, 0.05
3. shift 8, euler_steps 100
4. shift 8, euler_steps 50
5. shift 8, euler_steps 100, adv
6. pcm_linear_quadratic, euler_steps 50, 0.025, adv
7. pcm_linear_quadratic, euler_steps 50, 0.05, multiphase 125
8. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75
9. pcm_linear_quadratic, euler_steps 50, 0.05, range 0.75
10. pcm_linear_quadratic, euler_steps 50, 0.05, batchsize 32
11. pcm_linear_quadratic, euler_steps 50, learning rate,1e-7
12. shift1, euler_steps 50
13. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75, phase 1
14. 4.5 cfg, validation no cfg, pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75
15. pcm_linear_quadratic, euler_steps 50, 0.15, linear_range 0.75
16. pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75 ema 0.95, decay 0.0
```
bash scripts/distill_t2v.sh
```
17. no cfg, validation no cfg, pcm_linear_quadratic, euler_steps 50, 0.1, linear_range 0.75
18. shift16, euler_steps 50
## ⚡ Lora Finetune
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
## 💰Hardware requirement
- VRAM is required for both distill 10B mochi model
To launch finetuning, you will first need to prepare data in the following formats.
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
Then the finetuning can be launched by:
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
```
bash scripts/lora_finetune.sh
```
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.
## Acknowledgement
We learned from and reused code from the following projects: [PCM](https://github.com/G-U-N/Phased-Consistency-Model), [diffusers](https://github.com/huggingface/diffusers), and [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan).
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.3 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.3 MiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 380 KiB

+9
View File
@@ -0,0 +1,9 @@
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.
En "The Matrix", Neo, interpretado por Keanu Reeves, personifica la lucha contra un sistema opresor a través de su icónica imagen, que incluye unos anteojos oscuros. Estos lentes no son solo un accesorio de moda; representan una barrera entre la realidad y la percepción. Al usar estos anteojos, Neo se sumerge en un mundo donde la verdad se oculta detrás de ilusiones y engaños. La oscuridad de los lentes simboliza la ignorancia y el control que las máquinas tienen sobre la humanidad, mientras que su propia búsqueda de la verdad lo lleva a descubrir sus auténticos poderes. La escena en que se los pone se convierte en un momento crucial, marcando su transformación de un simple programador a "El Elegido". Esta imagen se ha convertido en un ícono cultural, encapsulando el mensaje de que, al enfrentar la oscuridad, podemos encontrar la luz que nos guía hacia la libertad. Así, los anteojos de Neo se convierten en un símbolo de resistencia y autoconocimiento en un mundo manipulado.
Medium close up. Low-angle shot. A woman in a 1950s retro dress sits in a diner bathed in neon light, surrounded by classic decor and lively chatter. The camera starts with a medium shot of her sitting at the counter, then slowly zooms in as she blows a shiny pink bubblegum bubble. The bubble swells dramatically before popping with a soft, playful burst. The scene is vibrant and nostalgic, evoking the fun and carefree spirit of the 1950s.
Will Smith eats noodles.
A short clip of the blonde woman taking a sip from her whiskey glass, her eyes locking with the camera as she smirks playfully. The background shows a group of people laughing and enjoying the party, with vibrant neon signs illuminating the space. The shot is taken in a way that conveys the feeling of a tipsy, carefree night out. The camera then zooms in on her face as she winks, creating a cheeky, flirtatious vibe.
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 robot's immense intelligence. Photorealistic style with a cinematic, dark sci-fi aesthetic. Aspect ratio: 16:9 --v 6.1
A chimpanzee lead vocalist singing into a microphone on stage. The camera zooms in to show him singing. There is a spotlight on him.
+75 -49
View File
@@ -4,31 +4,51 @@ 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. * 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
])
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,
]
)
# 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)
@@ -37,32 +57,34 @@ 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)
@@ -70,7 +92,9 @@ 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:
@@ -81,5 +105,7 @@ 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")
+68 -25
View File
@@ -4,13 +4,14 @@ import json
import os
import random
class LatentDataset(Dataset):
def __init__(
self,
json_path,
num_latent_t,
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
@@ -18,8 +19,10 @@ 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,27 +31,44 @@ 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)
# TODO: Hack
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,
)
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
@@ -59,25 +79,48 @@ 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()
+123 -78
View File
@@ -13,15 +13,12 @@ from tqdm import tqdm
from PIL import Image
from accelerate.logging import get_logger
from fastvideo.utils.dataset_utils import DecordInit
from fastvideo.utils.utils import text_preprocessing
import torchvision
logger = get_logger(__name__)
class SingletonMeta(type):
"""
这是一个元类,用于创建单例类。
"""
_instances = {}
def __call__(cls, *args, **kwargs):
@@ -53,7 +50,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:
@@ -61,18 +58,20 @@ 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):
@@ -95,11 +94,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
@@ -121,35 +120,39 @@ class T2V_dataset(Dataset):
data = self.get_data(idx)
return data
except Exception as e:
logger.info(f'Error with {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)]
@@ -158,51 +161,70 @@ 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'
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,
)
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 = text_preprocessing(caps, support_Chinese=self.support_Chinese)
text = caps
input_ids, cond_mask = [], []
text = text if random.random() > self.cfg else ""
text_tokens_and_mask = self.tokenizer(
text,
max_length=self.text_max_length,
padding='max_length',
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors='pt'
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"],
)
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
@@ -213,78 +235,100 @@ 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()
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)}')
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()
@@ -294,19 +338,20 @@ 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}...')
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
+120 -66
View File
@@ -32,7 +32,9 @@ 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):
@@ -42,21 +44,37 @@ 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"):
@@ -107,11 +125,10 @@ 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
@@ -121,15 +138,16 @@ 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)
@@ -159,7 +177,9 @@ 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
@@ -219,7 +239,9 @@ 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
@@ -235,7 +257,7 @@ class RandomCropVideo:
class SpatialStrideCropVideo:
def __init__(self, stride):
self.stride = stride
self.stride = stride
def __call__(self, clip):
"""
@@ -258,17 +280,18 @@ 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
@@ -291,27 +314,31 @@ 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
@@ -325,10 +352,15 @@ 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:
@@ -336,19 +368,21 @@ 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)
@@ -363,7 +397,9 @@ 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
@@ -372,18 +408,20 @@ 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)
@@ -398,13 +436,15 @@ 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)
@@ -516,6 +556,7 @@ 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.
@@ -530,13 +571,16 @@ 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
@@ -544,18 +588,20 @@ 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
@@ -569,7 +615,9 @@ 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]
@@ -580,12 +628,18 @@ 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),
)
+155
View File
@@ -0,0 +1,155 @@
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 fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.asymm_models_joint import AsymmDiTJoint
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import export_to_video
from fastvideo.distill.solver import PCMFMScheduler
import tempfile
import os
import argparse
from safetensors.torch import load_file
def init_args():
parser = argparse.ArgumentParser()
parser.add_argument("--prompts", nargs='+', default=[])
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument("--height", type=int, default=480)
parser.add_argument("--width", type=int, default=848)
parser.add_argument("--num_inference_steps", type=int, default=64)
parser.add_argument("--guidance_scale", type=float, default=4.5)
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--transformer_path", type=str, default=None)
parser.add_argument("--scheduler_type", type=str, default="euler")
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=100)
parser.add_argument("--linear_threshold", type=float, default=0.025)
parser.add_argument("--linear_range", type=float, default=0.5)
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:
scheduler = PCMFMScheduler(1000, args.shift, args.num_euler_timesteps, False, args.linear_threshold, args.linear_range)
mochi_genmo = True
if mochi_genmo:
model_path = "/root/fastmochi_genmo/dit.safetensors"
state_dcit = load_file(model_path)
transformer = AsymmDiTJoint()
transformer.load_state_dict(state_dcit)
# from IPython import embed
# embed()
transformer.config.in_channels = 12
print("load gennmo mochi successfully")
else:
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_model_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()
pipe = load_model(args)
print("load model successfully")
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()
with gr.Blocks() as demo:
gr.Markdown("# 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=8, maximum=256, value=args.num_frames)
guidance_scale = gr.Slider(label="Guidance Scale", minimum=1, maximum=20, value=args.guidance_scale)
num_inference_steps = gr.Slider(label="Inference Steps", minimum=10, 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)
+505 -323
View File
File diff suppressed because it is too large Load Diff
+8 -11
View File
@@ -23,7 +23,6 @@ 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__()
@@ -48,9 +47,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)
@@ -58,10 +57,9 @@ class DiscriminatorHead(nn.Module):
class Discriminator(nn.Module):
def __init__(
self,
stride = 8,
stride=8,
num_h_per_head=1,
adapter_channel_dims=[3072],
):
@@ -82,24 +80,23 @@ 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) // self.stride == len(self.heads)
for i in range(0, len(features), self.stride):
for h in self.heads[i//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
+8 -7
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.model.pipeline_mochi import linear_quadratic_schedule
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -17,13 +17,14 @@ 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):
class PCMFMScheduler(SchedulerMixin, ConfigMixin):
_compatibles = []
order = 1
@@ -34,13 +35,14 @@ 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(
@@ -238,6 +240,7 @@ 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
@@ -279,7 +282,6 @@ class EulerSolver:
multiphase,
is_target=False,
):
inference_indices = np.linspace(
0, len(self.euler_timesteps), num=multiphase, endpoint=False
)
@@ -305,4 +307,3 @@ class EulerSolver:
x_prev = sample + (sigma_prev - sigma) * model_pred
return x_prev, timestep_index_end
+508 -259
View File
File diff suppressed because it is too large Load Diff
+41 -40
View File
@@ -17,8 +17,8 @@ from torch.distributed.fsdp import (
# ShardedStateDictConfig, # un-flattened param but shards, usable by other parallel schemes.
)
from fastvideo.model.modeling_mochi import MochiTransformerBlock
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformerBlock
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.asymm_models_joint import AsymmetricJointBlock
from functools import partial
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
@@ -58,38 +58,43 @@ def apply_fsdp_checkpointing(model, p=1):
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,
)
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_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
def get_dit_fsdp_kwargs(sharding_strategy, use_lora=False, cpu_offload=False):
def get_dit_fsdp_kwargs(
sharding_strategy, use_lora=False, cpu_offload=False, master_weight_type="fp32"
):
if use_lora:
auto_wrap_policy = fsdp_auto_wrap_policy
else:
auto_wrap_policy = functools.partial(
transformer_auto_wrap_policy,
transformer_layer_cls={
MochiTransformerBlock,
MochiTransformerBlock, # AsymmetricJointBlock
},
)
# we use float32 for fsdp but autocast during training
mixed_precision = float32
mixed_precision = get_mixed_precision(master_weight_type)
if sharding_strategy == "full":
sharding_strategy = ShardingStrategy.FULL_SHARD
elif sharding_strategy == "hybrid_full":
@@ -98,10 +103,12 @@ def get_dit_fsdp_kwargs(sharding_strategy, use_lora=False, cpu_offload=False):
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,
@@ -110,29 +117,26 @@ def get_dit_fsdp_kwargs(sharding_strategy, use_lora=False, cpu_offload=False):
"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
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 = float32
sharding_strategy = ShardingStrategy.NO_SHARD
mixed_precision = get_mixed_precision(master_weight_type)
sharding_strategy = ShardingStrategy.NO_SHARD
device_id = torch.cuda.current_device()
fsdp_kwargs = {
"auto_wrap_policy": auto_wrap_policy,
@@ -141,8 +145,5 @@ def get_discriminator_fsdp_kwargs():
"device_id": device_id,
"limit_all_gathers": True,
}
return fsdp_kwargs
-123
View File
@@ -1,123 +0,0 @@
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
@@ -1,39 +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_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
-155
View File
@@ -1,155 +0,0 @@
import torch
from fastvideo.model.pipeline_mochi import MochiPipeline
import torch.distributed as dist
from diffusers.utils import export_to_video
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state, nccl_info
import argparse
import os
from diffusers.models.transformers.transformer_mochi import MochiTransformerBlock
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
import sys
import pdb
class ForkedPdb(pdb.Pdb):
"""
PDB Subclass for debugging multi-processed code
Suggested in: https://stackoverflow.com/questions/4716533/how-to-attach-debugger-to-a-python-subproccess
"""
def interaction(self, *args, **kwargs):
_stdin = sys.stdin
try:
sys.stdin = open('/dev/stdin')
pdb.Pdb.interaction(self, *args, **kwargs)
finally:
sys.stdin = _stdin
def assert_all_close_list(input_list):
for i in range(len(input_list) - 1):
assert torch.allclose(input_list[i], input_list[i + 1]), f"input_list[{i}]: {input_list[i]}, input_list[{i+1}]: {input_list[i+1]}"
weight_dtype = torch.float32
def initialize_distributed():
local_rank = int(os.getenv('RANK', 0))
world_size = int(os.getenv('WORLD_SIZE', 1))
print('world_size', world_size)
torch.cuda.set_device(local_rank)
dist.init_process_group(backend='nccl', init_method='env://', world_size=world_size, rank=local_rank)
return world_size
def main_print(content):
if int(os.getenv('RANK', 0)) <= 0:
print(content)
@torch.inference_mode
def test_single_block(batch_size, device, seed):
# set manual seed
torch.manual_seed(seed)
device = torch.cuda.current_device()
block = MochiTransformerBlock(
dim=768,
num_attention_heads=12,
attention_head_dim=64,
pooled_projection_dim=256,
qk_norm="rms_norm",
activation_fn="swiglu",
context_pre_only=False,
).to(device)
hidden_states = torch.randn(1, 16, 768).to(device).repeat(batch_size, 1, 1)
encoder_hidden_states = torch.randn(1, 4, 256).to(device).repeat(batch_size, 1, 1)
temb = torch.randn(1, 768).to(device).repeat(batch_size, 1)
# shard hiddent_states according to world_size
local_seq_length = hidden_states.shape[1] // nccl_info.sp_size
hidden_states = hidden_states.narrow(1, nccl_info.global_rank * local_seq_length, local_seq_length)
main_print(hidden_states.shape)
hidden_states, encoder_hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
temb=temb,
)
mean = hidden_states[0].mean()
torch.distributed.all_reduce(mean, op=torch.distributed.ReduceOp.SUM)
mean = mean / nccl_info.sp_size
return mean
@torch.inference_mode
def test_DiT(batch_size, transformer, seed):
generator = torch.Generator(torch.cuda.current_device()).manual_seed(seed)
device = torch.cuda.current_device()
latent = torch.randn((1, 12, 8, 12, 8), device=device, dtype=weight_dtype, generator=generator).repeat(batch_size, 1, 1, 1, 1)
prompt_embeds = torch.randn((1, 20, 4096), device=device, dtype=weight_dtype, generator=generator).repeat(batch_size, 1, 1)
prompt_attention_mask = torch.ones((1, 20), device=device, dtype=weight_dtype).repeat(batch_size, 1)
timestep = 0
timestep = torch.tensor(timestep, device=device, dtype=weight_dtype).unsqueeze(0).repeat(batch_size)
local_seq_length = latent.shape[2] // nccl_info.sp_size
latent = latent.narrow(2, nccl_info.global_rank * local_seq_length, local_seq_length)
# main_print(latent.shape)
hidden_states = transformer(
hidden_states=latent,
encoder_hidden_states=prompt_embeds,
encoder_attention_mask=prompt_attention_mask,
timestep=timestep,
return_dict=False,
)[0]
def calculate_mean(states):
mean = states.mean()
torch.distributed.all_reduce(mean, op=torch.distributed.ReduceOp.SUM)
mean = mean / int(os.getenv('WORLD_SIZE', 1))
return mean
mean1 = calculate_mean(hidden_states[0])
main_print(hidden_states.shape)
if hidden_states.shape[0] > 1:
mean2 = calculate_mean(hidden_states[1])
return mean1, mean2
return mean1
if __name__ == "__main__":
world_size = initialize_distributed()
device = torch.cuda.current_device()
parser = argparse.ArgumentParser()
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--test_single_block", action="store_true")
args = parser.parse_args()
seed = args.seed
if args.test_single_block:
pass
single_no_patch_bs_1 = test_single_block(1)
single_no_patch_bs_2 = test_single_block(2)
# check all close
assert torch.allclose(single_no_patch_bs_1, single_no_patch_bs_2)
single_patch_bs_1 = test_single_block(1)
single_patch_bs_2 = test_single_block(2)
assert torch.allclose(single_patch_bs_1, single_patch_bs_2)
assert torch.allclose(single_no_patch_bs_1, single_patch_bs_2)
initialize_sequence_parallel_state(world_size)
sp_patch_bs_1 = test_single_block(1)
sp_patch_bs_2 = test_single_block(2)
assert torch.allclose(sp_patch_bs_1, sp_patch_bs_2)
assert torch.allclose(single_no_patch_bs_1, sp_patch_bs_2)
else:
transformer = MochiTransformer3DModel.from_pretrained("data/mochi/transformer", torch_dtype=weight_dtype).to(device)
single_no_patch_bs_1 = test_DiT(1, transformer, seed)
single_no_patch_bs_2_a, single_no_patch_bs_2_b = test_DiT(2, transformer, seed)
single_patch_bs_1 = test_DiT(1, transformer, seed)
single_patch_bs_2_a, single_patch_bs_2_b = test_DiT(2, transformer, seed)
initialize_sequence_parallel_state(world_size)
sp_patch_bs_1 = test_DiT(1, transformer, seed)
sp_patch_bs_2_a, sp_patch_bs_2_b = test_DiT(2, transformer, seed)
assert_all_close_list([single_no_patch_bs_1, single_no_patch_bs_2_a, single_no_patch_bs_2_b, single_patch_bs_1, single_patch_bs_2_a, single_patch_bs_2_b, sp_patch_bs_1, sp_patch_bs_2_a, sp_patch_bs_2_b])
main_print(sp_patch_bs_1)
@@ -0,0 +1,29 @@
from contextlib import contextmanager
import torch
try:
from flash_attn import flash_attn_varlen_func as flash_varlen_attn
except ImportError:
flash_varlen_attn = None
try:
from sageattention import sageattn as sage_attn
except ImportError:
sage_attn = None
from torch.nn.attention import SDPBackend, sdpa_kernel
training_backends = [SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION]
eval_backends = list(training_backends)
if torch.cuda.get_device_properties(0).major >= 9.0:
# Enable fast CuDNN attention on Hopper.
# This gives NaN on the backward pass for some reason,
# so only use it for evaluation.
eval_backends.append(SDPBackend.CUDNN_ATTENTION)
@contextmanager
def sdpa_attn_ctx(training: bool = False):
with sdpa_kernel(training_backends if training else eval_backends):
yield
@@ -0,0 +1,87 @@
import contextlib
from typing import Any, Iterable, Iterator, Optional
try:
from tqdm import tqdm
except ImportError:
tqdm = None
try:
from ray.experimental.tqdm_ray import tqdm as ray_tqdm
except:
ray_tqdm = None
# Global state
_current_progress_type = "none"
_is_progress_bar_active = False
class DummyProgressBar:
"""A no-op progress bar that mimics tqdm interface"""
def __init__(self, iterable=None, **kwargs):
self.iterable = iterable
def __iter__(self):
return iter(self.iterable)
def update(self, n=1):
pass
def close(self):
pass
def set_description(self, desc):
pass
def get_new_progress_bar(iterable: Optional[Iterable] = None, **kwargs) -> Any:
if not _is_progress_bar_active:
return DummyProgressBar(iterable=iterable, **kwargs)
if _current_progress_type == "tqdm":
if tqdm is None:
raise ImportError("tqdm is required but not installed. Please install tqdm to use the tqdm progress bar.")
return tqdm(iterable=iterable, **kwargs)
elif _current_progress_type == "ray_tqdm":
if ray_tqdm is None:
raise ImportError("ray is required but not installed. Please install ray to use the ray_tqdm progress bar.")
return ray_tqdm(iterable=iterable, **kwargs)
return DummyProgressBar(iterable=iterable, **kwargs)
@contextlib.contextmanager
def progress_bar(type: str = "none", enabled=True):
"""
Context manager for setting progress bar type and options.
Args:
type: Type of progress bar ("none" or "tqdm")
**options: Options to pass to the progress bar (e.g., total, desc)
Raises:
ValueError: If progress bar type is invalid
RuntimeError: If progress bars are nested
Example:
with progress_bar(type="tqdm", total=100):
for i in get_new_progress_bar(range(100)):
process(i)
"""
if type not in ("none", "tqdm", "ray_tqdm"):
raise ValueError("Progress bar type must be 'none' or 'tqdm' or 'ray_tqdm'")
if not enabled:
type = "none"
global _current_progress_type, _is_progress_bar_active
if _is_progress_bar_active:
raise RuntimeError("Nested progress bars are not supported")
_is_progress_bar_active = True
_current_progress_type = type
try:
yield
finally:
_is_progress_bar_active = False
_current_progress_type = "none"
+67
View File
@@ -0,0 +1,67 @@
import os
import subprocess
import tempfile
import time
import numpy as np
from moviepy.editor import ImageSequenceClip
from PIL import Image
from genmo.lib.progress import get_new_progress_bar
class Timer:
def __init__(self):
self.times = {} # Dictionary to store times per stage
def __call__(self, name):
print(f"Timing {name}")
return self.TimerContextManager(self, name)
def print_stats(self):
total_time = sum(self.times.values())
# Print table header
print("{:<20} {:>10} {:>10}".format("Stage", "Time(s)", "Percent"))
for name, t in self.times.items():
percent = (t / total_time) * 100 if total_time > 0 else 0
print("{:<20} {:>10.2f} {:>9.2f}%".format(name, t, percent))
class TimerContextManager:
def __init__(self, outer, name):
self.outer = outer # Reference to the Timer instance
self.name = name
self.start_time = None
def __enter__(self):
self.start_time = time.perf_counter()
return self
def __exit__(self, exc_type, exc_value, traceback):
end_time = time.perf_counter()
elapsed = end_time - self.start_time
self.outer.times[self.name] = self.outer.times.get(self.name, 0) + elapsed
def save_video(final_frames, output_path, fps=30):
assert final_frames.ndim == 4 and final_frames.shape[3] == 3, f"invalid shape: {final_frames} (need t h w c)"
if final_frames.dtype != np.uint8:
final_frames = (final_frames * 255).astype(np.uint8)
ImageSequenceClip(list(final_frames), fps=fps).write_videofile(output_path)
def create_memory_tracker():
import torch
previous = [None] # Use list for mutable closure state
def track(label="all2all"):
current = torch.cuda.memory_allocated() / 1e9
if previous[0] is not None:
diff = current - previous[0]
sign = "+" if diff >= 0 else ""
print(f"GPU memory ({label}): {current:.2f} GB ({sign}{diff:.2f} GB)")
else:
print(f"GPU memory ({label}): {current:.2f} GB")
previous[0] = current # type: ignore
return track
@@ -0,0 +1,721 @@
import os
from typing import Dict, List, Optional, Tuple, Any
import warnings
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from torch.nn.attention import sdpa_kernel
from fastvideo.models.mochi_genmo.lib.attn_imports import flash_varlen_attn, sage_attn, sdpa_attn_ctx
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.layers import (
FeedForward,
PatchEmbed,
RMSNorm,
TimestepEmbedder,
)
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.lora import LoraLinear
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.mod_rmsnorm import modulated_rmsnorm
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.residual_tanh_gated_rmsnorm import (
residual_tanh_gated_rmsnorm,
)
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.rope_mixed import (
compute_mixed_rotation,
create_position_matrix,
)
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.temporal_rope import apply_rotary_emb_qk_real
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.utils import (
AttentionPool,
modulate,
pad_and_split_xy,
)
from fastvideo.models.mochi_genmo.mochi_preview.pipelines import compute_packed_indices
from diffusers.models.modeling_utils import ModelMixin
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders import PeftAdapterMixin
COMPILE_FINAL_LAYER = os.environ.get("COMPILE_DIT") == "1"
COMPILE_MMDIT_BLOCK = os.environ.get("COMPILE_DIT") == "1"
def ck(fn, *args, enabled=True, **kwargs) -> torch.Tensor:
if enabled:
return torch.utils.checkpoint.checkpoint(fn, *args, **kwargs, use_reentrant=False)
return fn(*args, **kwargs)
class AsymmetricAttention(nn.Module):
def __init__(
self,
dim_x: int,
dim_y: int,
num_heads: int = 8,
qkv_bias: bool = False,
qk_norm: bool = True,
update_y: bool = True,
out_bias: bool = True,
attention_mode: str = "flash",
softmax_scale: Optional[float] = None,
device: Optional[torch.device] = None,
# Disable LoRA by default ...
qkv_proj_lora_rank: int = 0,
qkv_proj_lora_alpha: int = 0,
qkv_proj_lora_dropout: float = 0.0,
out_proj_lora_rank: int = 0,
out_proj_lora_alpha: int = 0,
out_proj_lora_dropout: float = 0.0,
):
super().__init__()
self.attention_mode = attention_mode
self.dim_x = dim_x
self.dim_y = dim_y
self.num_heads = num_heads
self.head_dim = dim_x // num_heads
self.update_y = update_y
self.softmax_scale = softmax_scale
if dim_x % num_heads != 0:
raise ValueError(f"dim_x={dim_x} should be divisible by num_heads={num_heads}")
# Input layers.
self.qkv_bias = qkv_bias
qkv_lora_kwargs = dict(
bias=qkv_bias,
device=device,
r=qkv_proj_lora_rank,
lora_alpha=qkv_proj_lora_alpha,
lora_dropout=qkv_proj_lora_dropout,
)
self.qkv_x = LoraLinear(dim_x, 3 * dim_x, **qkv_lora_kwargs)
# Project text features to match visual features (dim_y -> dim_x)
self.qkv_y = LoraLinear(dim_y, 3 * dim_x, **qkv_lora_kwargs)
# Query and key normalization for stability.
assert qk_norm
self.q_norm_x = RMSNorm(self.head_dim, device=device)
self.k_norm_x = RMSNorm(self.head_dim, device=device)
self.q_norm_y = RMSNorm(self.head_dim, device=device)
self.k_norm_y = RMSNorm(self.head_dim, device=device)
# Output layers. y features go back down from dim_x -> dim_y.
proj_lora_kwargs = dict(
bias=out_bias,
device=device,
r=out_proj_lora_rank,
lora_alpha=out_proj_lora_alpha,
lora_dropout=out_proj_lora_dropout,
)
self.proj_x = LoraLinear(dim_x, dim_x, **proj_lora_kwargs)
self.proj_y = LoraLinear(dim_x, dim_y, **proj_lora_kwargs) if update_y else nn.Identity()
def run_qkv_y(self, y):
local_heads = self.num_heads
qkv_y = self.qkv_y(y) # (B, L, 3 * dim)
qkv_y = qkv_y.view(qkv_y.size(0), qkv_y.size(1), 3, local_heads, self.head_dim)
q_y, k_y, v_y = qkv_y.unbind(2)
q_y = self.q_norm_y(q_y)
k_y = self.k_norm_y(k_y)
return q_y, k_y, v_y
def prepare_qkv(
self,
x: torch.Tensor, # (B, M, dim_x)
y: torch.Tensor, # (B, L, dim_y)
*,
scale_x: torch.Tensor,
scale_y: torch.Tensor,
rope_cos: torch.Tensor,
rope_sin: torch.Tensor,
valid_token_indices: torch.Tensor,
max_seqlen_in_batch: int,
):
# Process visual features
x = modulated_rmsnorm(x, scale_x) # (B, M, dim_x) where M = N
qkv_x = self.qkv_x(x) # (B, M, 3 * dim_x)
assert qkv_x.dtype == torch.bfloat16
B, M, _ = qkv_x.size()
qkv_x = qkv_x.view(B, M, 3, self.num_heads, -1)
qkv_x = qkv_x.permute(2, 0, 1, 3, 4)
# Split qkv_x into q, k, v
q_x, k_x, v_x = qkv_x.unbind(0) # (B, N, local_h, head_dim)
q_x = self.q_norm_x(q_x)
q_x = apply_rotary_emb_qk_real(q_x, rope_cos, rope_sin)
k_x = self.k_norm_x(k_x)
k_x = apply_rotary_emb_qk_real(k_x, rope_cos, rope_sin)
# Concatenate streams
B, N, num_heads, head_dim = q_x.size()
D = num_heads * head_dim
# Process text features
if B == 1:
text_seqlen = max_seqlen_in_batch - N
if text_seqlen > 0:
y = y[:, :text_seqlen] # Remove padding tokens.
y = modulated_rmsnorm(y, scale_y) # (B, L, dim_y)
q_y, k_y, v_y = self.run_qkv_y(y) # (B, L, local_heads, head_dim)
q = torch.cat([q_x, q_y], dim=1)
k = torch.cat([k_x, k_y], dim=1)
v = torch.cat([v_x, v_y], dim=1)
else:
q, k, v = q_x, k_x, v_x
else:
y = modulated_rmsnorm(y, scale_y) # (B, L, dim_y)
q_y, k_y, v_y = self.run_qkv_y(y) # (B, L, local_heads, head_dim)
indices = valid_token_indices[:, None].expand(-1, D)
q = torch.cat([q_x, q_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
k = torch.cat([k_x, k_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
v = torch.cat([v_x, v_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
q = q.view(-1, num_heads, head_dim)
k = k.view(-1, num_heads, head_dim)
v = v.view(-1, num_heads, head_dim)
return q, k, v
@torch.autocast("cuda", enabled=False)
def flash_attention(self, q, k, v, cu_seqlens, max_seqlen_in_batch, total, local_dim):
out: torch.Tensor = flash_varlen_attn(
q, k, v,
cu_seqlens_q=cu_seqlens,
cu_seqlens_k=cu_seqlens,
max_seqlen_q=max_seqlen_in_batch,
max_seqlen_k=max_seqlen_in_batch,
dropout_p=0.0,
softmax_scale=self.softmax_scale,
) # (total, local_heads, head_dim)
return out.view(total, local_dim)
def sdpa_attention(self, q, k, v):
with sdpa_attn_ctx(training=self.training):
out = F.scaled_dot_product_attention(
q, k, v,
attn_mask=None,
dropout_p=0.0,
is_causal=False,
)
return out
@torch.autocast("cuda", enabled=False)
def sage_attention(self, q, k, v):
return sage_attn(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=False)
def run_attention(
self,
q: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
k: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
v: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
*,
B: int,
cu_seqlens: Optional[torch.Tensor] = None,
max_seqlen_in_batch: Optional[int] = None,
):
local_heads = self.num_heads
local_dim = local_heads * self.head_dim
# Check shapes
assert q.ndim == 3 and k.ndim == 3 and v.ndim == 3
total = q.size(0)
assert k.size(0) == total and v.size(0) == total
if self.attention_mode == "flash":
out = self.flash_attention(
q, k, v, cu_seqlens, max_seqlen_in_batch, total, local_dim) # (total, local_dim)
else:
assert B == 1, \
f"Non-flash attention mode {self.attention_mode} only supports batch size 1, got {B}"
q = rearrange(q, "(b s) h d -> b h s d", b=B)
k = rearrange(k, "(b s) h d -> b h s d", b=B)
v = rearrange(v, "(b s) h d -> b h s d", b=B)
if self.attention_mode == "sdpa":
out = self.sdpa_attention(q, k, v) # (B, local_heads, seq_len, head_dim)
elif self.attention_mode == "sage":
out = self.sage_attention(q, k, v) # (B, local_heads, seq_len, head_dim)
else:
raise ValueError(f"Unknown attention mode: {self.attention_mode}")
out = rearrange(out, "b h s d -> (b s) (h d)")
return out
def post_attention(
self,
out: torch.Tensor,
B: int,
M: int,
L: int,
dtype: torch.dtype,
valid_token_indices: torch.Tensor,
):
"""
Args:
out: (total <= B * (N + L), local_dim)
valid_token_indices: (total <= B * (N + L),)
B: Batch size
M: Number of visual tokens per context parallel rank
L: Number of text tokens
dtype: Data type of the input and output tensors
Returns:
x: (B, N, dim_x) tensor of visual tokens where N = M
y: (B, L, dim_y) tensor of text token features
"""
local_heads = self.num_heads
local_dim = local_heads * self.head_dim
N = M
# Split sequence into visual and text tokens, adding back padding.
if B == 1:
out = out.view(B, -1, local_dim)
if out.size(1) > N:
x, y = torch.tensor_split(out, (N,), dim=1) # (B, N, local_dim), (B, <= L, local_dim)
y = F.pad(y, (0, 0, 0, L - y.size(1))) # (B, L, local_dim)
else:
# Empty prompt.
x, y = out, out.new_zeros(B, L, local_dim)
else:
x, y = pad_and_split_xy(out, valid_token_indices, B, N, L, dtype)
assert x.size() == (B, N, local_dim)
assert y.size() == (B, L, local_dim)
# Communicate across context parallel ranks.
x = x.view(B, N, local_heads, self.head_dim)
x = x.view(x.size(0), x.size(1), x.size(2) * x.size(3)) # (B, M, dim_x = num_heads * head_dim)
x = self.proj_x(x)
y = self.proj_y(y)
return x, y
def forward(
self,
x: torch.Tensor, # (B, M, dim_x)
y: torch.Tensor, # (B, L, dim_y)
*,
scale_x: torch.Tensor, # (B, dim_x), modulation for pre-RMSNorm.
scale_y: torch.Tensor, # (B, dim_y), modulation for pre-RMSNorm.
packed_indices: Dict[str, torch.Tensor] = None,
checkpoint_qkv: bool = False,
checkpoint_post_attn: bool = False,
**rope_rotation,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Forward pass of asymmetric multi-modal attention.
Args:
x: (B, M, dim_x) tensor of visual tokens
y: (B, L, dim_y) tensor of text token features
packed_indices: Dict with keys for Flash Attention
num_frames: Number of frames in the video. N = num_frames * num_spatial_tokens
Returns:
x: (B, M, dim_x) tensor of visual tokens after multi-modal attention
y: (B, L, dim_y) tensor of text token features after multi-modal attention
"""
B, L, _ = y.shape
_, M, _ = x.shape
# Predict a packed QKV tensor from visual and text features.
q, k, v = ck(self.prepare_qkv,
x=x,
y=y,
scale_x=scale_x,
scale_y=scale_y,
rope_cos=rope_rotation.get("rope_cos"),
rope_sin=rope_rotation.get("rope_sin"),
valid_token_indices=packed_indices["valid_token_indices_kv"],
max_seqlen_in_batch=packed_indices["max_seqlen_in_batch_kv"],
enabled=checkpoint_qkv,
) # (total <= B * (N + L), 3, local_heads, head_dim)
# Self-attention is expensive, so don't checkpoint it.
out = self.run_attention(
q, k, v, B=B,
cu_seqlens=packed_indices["cu_seqlens_kv"],
max_seqlen_in_batch=packed_indices["max_seqlen_in_batch_kv"],
)
x, y = ck(self.post_attention,
out,
B=B, M=M, L=L,
dtype=v.dtype,
valid_token_indices=packed_indices["valid_token_indices_kv"],
enabled=checkpoint_post_attn,
)
return x, y
@torch.compile(disable=not COMPILE_MMDIT_BLOCK)
class AsymmetricJointBlock(nn.Module):
def __init__(
self,
hidden_size_x: int,
hidden_size_y: int,
num_heads: int,
*,
mlp_ratio_x: float = 8.0, # Ratio of hidden size to d_model for MLP for visual tokens.
mlp_ratio_y: float = 4.0, # Ratio of hidden size to d_model for MLP for text tokens.
update_y: bool = True, # Whether to update text tokens in this block.
device: Optional[torch.device] = None,
**block_kwargs,
):
super().__init__()
self.update_y = update_y
self.hidden_size_x = hidden_size_x
self.hidden_size_y = hidden_size_y
self.mod_x = nn.Linear(hidden_size_x, 4 * hidden_size_x, device=device)
if self.update_y:
self.mod_y = nn.Linear(hidden_size_x, 4 * hidden_size_y, device=device)
else:
self.mod_y = nn.Linear(hidden_size_x, hidden_size_y, device=device)
# Self-attention:
self.attn = AsymmetricAttention(
hidden_size_x,
hidden_size_y,
num_heads=num_heads,
update_y=update_y,
device=device,
**block_kwargs,
)
# MLP.
mlp_hidden_dim_x = int(hidden_size_x * mlp_ratio_x)
assert mlp_hidden_dim_x == int(1536 * 8)
self.mlp_x = FeedForward(
in_features=hidden_size_x,
hidden_size=mlp_hidden_dim_x,
multiple_of=256,
ffn_dim_multiplier=None,
device=device,
)
# MLP for text not needed in last block.
if self.update_y:
mlp_hidden_dim_y = int(hidden_size_y * mlp_ratio_y)
self.mlp_y = FeedForward(
in_features=hidden_size_y,
hidden_size=mlp_hidden_dim_y,
multiple_of=256,
ffn_dim_multiplier=None,
device=device,
)
def forward(
self,
x: torch.Tensor,
c: torch.Tensor,
y: torch.Tensor,
# TODO: These could probably just go into attn_kwargs
checkpoint_ff: bool = False,
checkpoint_qkv: bool = False,
checkpoint_post_attn: bool = False,
**attn_kwargs,
):
"""Forward pass of a block.
Args:
x: (B, N, dim) tensor of visual tokens
c: (B, dim) tensor of conditioned features
y: (B, L, dim) tensor of text tokens
num_frames: Number of frames in the video. N = num_frames * num_spatial_tokens
Returns:
x: (B, N, dim) tensor of visual tokens after block
y: (B, L, dim) tensor of text tokens after block
"""
N = x.size(1)
c = F.silu(c)
mod_x = self.mod_x(c)
scale_msa_x, gate_msa_x, scale_mlp_x, gate_mlp_x = mod_x.chunk(4, dim=1)
mod_y = self.mod_y(c)
if self.update_y:
scale_msa_y, gate_msa_y, scale_mlp_y, gate_mlp_y = mod_y.chunk(4, dim=1)
else:
scale_msa_y = mod_y
# Self-attention block.
x_attn, y_attn = self.attn(
x,
y,
scale_x=scale_msa_x,
scale_y=scale_msa_y,
checkpoint_qkv=checkpoint_qkv,
checkpoint_post_attn=checkpoint_post_attn,
**attn_kwargs,
)
assert x_attn.size(1) == N
x = residual_tanh_gated_rmsnorm(x, x_attn, gate_msa_x)
if self.update_y:
y = residual_tanh_gated_rmsnorm(y, y_attn, gate_msa_y)
# MLP block.
x = ck(self.ff_block_x, x, scale_mlp_x, gate_mlp_x, enabled=checkpoint_ff)
if self.update_y:
y = ck(self.ff_block_y, y, scale_mlp_y, gate_mlp_y, enabled=checkpoint_ff) # type: ignore
return x, y
def ff_block_x(self, x, scale_x, gate_x):
x_mod = modulated_rmsnorm(x, scale_x)
x_res = self.mlp_x(x_mod)
x = residual_tanh_gated_rmsnorm(x, x_res, gate_x) # Sandwich norm
return x
def ff_block_y(self, y, scale_y, gate_y):
y_mod = modulated_rmsnorm(y, scale_y)
y_res = self.mlp_y(y_mod)
y = residual_tanh_gated_rmsnorm(y, y_res, gate_y) # Sandwich norm
return y
@torch.compile(disable=not COMPILE_FINAL_LAYER)
class FinalLayer(nn.Module):
"""
The final layer of DiT.
"""
def __init__(
self,
hidden_size,
patch_size,
out_channels,
device: Optional[torch.device] = None,
):
super().__init__()
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, device=device)
self.mod = nn.Linear(hidden_size, 2 * hidden_size, device=device)
self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, device=device)
def forward(self, x, c):
c = F.silu(c)
shift, scale = self.mod(c).chunk(2, dim=1)
x = modulate(self.norm_final(x), shift, scale)
x = self.linear(x)
return x
class AsymmDiTJoint(ModelMixin, ConfigMixin, PeftAdapterMixin):
"""
Diffusion model with a Transformer backbone.
Ingests text embeddings instead of a label.
"""
@register_to_config
def __init__(
self,
*,
patch_size=2,
in_channels=12,
hidden_size_x=3072,
hidden_size_y=1536,
depth=48,
num_heads=24,
mlp_ratio_x=4.0,
mlp_ratio_y=4.0,
t5_feat_dim: int = 4096,
t5_token_length: int = 256,
patch_embed_bias: bool = True,
timestep_mlp_bias: bool = True,
timestep_scale: float = 1000.0,
use_extended_posenc: bool = False,
rope_theta: float = 10000.0,
device: Optional[torch.device] = None,
**block_kwargs,
):
super().__init__()
self.in_channels = in_channels
self.out_channels = in_channels
self.patch_size = patch_size
self.num_heads = num_heads
self.hidden_size_x = hidden_size_x
self.hidden_size_y = hidden_size_y
self.head_dim = hidden_size_x // num_heads # Head dimension and count is determined by visual.
self.use_extended_posenc = use_extended_posenc
self.t5_token_length = t5_token_length
self.t5_feat_dim = t5_feat_dim
self.rope_theta = rope_theta # Scaling factor for frequency computation for temporal RoPE.
self.timestep_scale = timestep_scale
self.x_embedder = PatchEmbed(
patch_size=patch_size,
in_chans=in_channels,
embed_dim=hidden_size_x,
bias=patch_embed_bias,
device=device,
)
# Conditionings
# Timestep
self.t_embedder = TimestepEmbedder(hidden_size_x, bias=timestep_mlp_bias, timestep_scale=timestep_scale)
# Caption Pooling (T5)
self.t5_y_embedder = AttentionPool(t5_feat_dim, num_heads=8, output_dim=hidden_size_x, device=device)
# Dense Embedding Projection (T5)
self.t5_yproj = nn.Linear(t5_feat_dim, hidden_size_y, bias=True, device=device)
# Initialize pos_frequencies as an empty parameter.
self.pos_frequencies = nn.Parameter(torch.empty(3, self.num_heads, self.head_dim // 2, device=device))
# for depth 48:
# b = 0: AsymmetricJointBlock, update_y=True
# b = 1: AsymmetricJointBlock, update_y=True
# ...
# b = 46: AsymmetricJointBlock, update_y=True
# b = 47: AsymmetricJointBlock, update_y=False. No need to update text features.
blocks = []
for b in range(depth):
# Joint multi-modal block
update_y = b < depth - 1
block = AsymmetricJointBlock(
hidden_size_x,
hidden_size_y,
num_heads,
mlp_ratio_x=mlp_ratio_x,
mlp_ratio_y=mlp_ratio_y,
update_y=update_y,
device=device,
**block_kwargs,
)
blocks.append(block)
self.blocks = nn.ModuleList(blocks)
self.final_layer = FinalLayer(hidden_size_x, patch_size, self.out_channels, device=device)
def embed_x(self, x: torch.Tensor) -> torch.Tensor:
"""
Args:
x: (B, C=12, T, H, W) tensor of visual tokens
Returns:
x: (B, C=3072, N) tensor of visual tokens with positional embedding.
"""
return self.x_embedder(x) # Convert BcTHW to BCN
@torch.compile(disable=not COMPILE_MMDIT_BLOCK)
def prepare(
self,
x: torch.Tensor,
sigma: torch.Tensor,
t5_feat: torch.Tensor,
t5_mask: torch.Tensor,
):
"""Prepare input and conditioning embeddings."""
# Visual patch embeddings with positional encoding.
T, H, W = x.shape[-3:]
pH, pW = H // self.patch_size, W // self.patch_size
x = self.embed_x(x) # (B, N, D), where N = T * H * W / patch_size ** 2
assert x.ndim == 3
B = x.size(0)
# Construct position array of size [N, 3].
# pos[:, 0] is the frame index for each location,
# pos[:, 1] is the row index for each location, and
# pos[:, 2] is the column index for each location.
N = T * pH * pW
assert x.size(1) == N
pos = create_position_matrix(T, pH=pH, pW=pW, device=x.device, dtype=torch.float32) # (N, 3)
rope_cos, rope_sin = compute_mixed_rotation(
freqs=self.pos_frequencies, pos=pos
) # Each are (N, num_heads, dim // 2)
# Global vector embedding for conditionings.
c_t = self.t_embedder(1 - sigma) # (B, D)
# Pool T5 tokens using attention pooler
# Note encoder_hidden_states[1] contains T5 token features.
assert (
t5_feat.size(1) == self.t5_token_length
), f"Expected L={self.t5_token_length}, got {t5_feat.shape} for encoder_hidden_states."
t5_y_pool = self.t5_y_embedder(t5_feat, t5_mask) # (B, D)
assert t5_y_pool.size(0) == B, f"Expected B={B}, got {t5_y_pool.shape} for t5_y_pool."
c = c_t + t5_y_pool
encoder_hidden_states = self.t5_yproj(t5_feat) # (B, L, t5_feat_dim) --> (B, L, D)
return x, c, encoder_hidden_states, rope_cos, rope_sin
def forward(
self,
hidden_states: torch.Tensor, # [1, 12, 7, 60, 106]
timestep: torch.Tensor, # [1]
encoder_hidden_states: torch.Tensor, # [1, 256, 4096]
encoder_attention_mask: torch.Tensor, # [1, 256]
attention_kwargs: Optional[Dict[str, Any]] = None,
return_dict: bool = False,
rope_cos: torch.Tensor = None,
rope_sin: torch.Tensor = None,
num_ff_checkpoint: int = 48, # 48
num_qkv_checkpoint: int = 48, # 48
num_post_attn_checkpoint: int = 0, # 0
):
"""Forward pass of DiT.
Args:
hidden_states: (B, C, T, H, W) tensor of spatial inputs (images or latent representations of images)
sigma: (B,) tensor of noise standard deviations
encoder_hidden_states: List((B, L, encoder_hidden_states_dim) tensor of caption token features. For SDXL text encoders: L=77, encoder_hidden_states_dim=2048)
encoder_attention_mask: List((B, L) boolean tensor indicating which tokens are not padding)
packed_indices: Dict with keys for Flash Attention. Result of compute_packed_indices.
{'cu_seqlens_kv': tensor([ 0, 11230], device='cuda:0', dtype=torch.int32),
'max_seqlen_in_batch_kv': 11230,
'valid_token_indices_kv': tensor([ 0, 1, 2, ..., 11227, 11228, 11229], device='cuda:0')}
"""
sigma = timestep / self.timestep_scale
num_latent_toks = np.prod(hidden_states.shape[-3:])
packed_indices = compute_packed_indices(hidden_states.device, encoder_attention_mask, int(num_latent_toks))
_, _, T, H, W = hidden_states.shape
if self.pos_frequencies.dtype != torch.float32:
warnings.warn(f"pos_frequencies dtype {self.pos_frequencies.dtype} != torch.float32")
# Use EFFICIENT_ATTENTION backend for T5 pooling, since we have a mask.
# Have to call sdpa_kernel outside of a torch.compile region.
with sdpa_kernel(torch.nn.attention.SDPBackend.EFFICIENT_ATTENTION):
hidden_states, c, encoder_hidden_states, rope_cos, rope_sin = self.prepare(hidden_states, sigma, encoder_hidden_states, encoder_attention_mask) # [1, 11130, 3072], [1, 3072], [1, 256, 1536], [11130, 24, 64]
del encoder_attention_mask
for i, block in enumerate(self.blocks):
hidden_states, encoder_hidden_states = block( # [1, 11130, 3072], [1, 256, 1536]
hidden_states,
c,
encoder_hidden_states,
rope_cos=rope_cos,
rope_sin=rope_sin,
packed_indices=packed_indices,
checkpoint_ff=i < num_ff_checkpoint,
checkpoint_qkv=i < num_qkv_checkpoint,
checkpoint_post_attn=i < num_post_attn_checkpoint,
) # (B, M, D), (B, L, D)
del encoder_hidden_states # Final layers don't use dense text features.
hidden_states = self.final_layer(hidden_states, c) # (B, M, patch_size ** 2 * out_channels) [1, 11130, 48]
hidden_states = rearrange( # [1, 12, 7, 60, 106]
hidden_states,
"B (T hp wp) (p1 p2 c) -> B c T (hp p1) (wp p2)",
T=T,
hp=H // self.patch_size,
wp=W // self.patch_size,
p1=self.patch_size,
p2=self.patch_size,
c=self.out_channels,
)
attn_outputs_list = None
return (-hidden_states, attn_outputs_list)
@@ -0,0 +1,737 @@
import os
from typing import Dict, List, Optional, Tuple
import warnings
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from torch.nn.attention import sdpa_kernel
import genmo.mochi_preview.dit.joint_model.context_parallel as cp
from genmo.lib.attn_imports import flash_varlen_attn, sage_attn, sdpa_attn_ctx
from genmo.mochi_preview.dit.joint_model.layers import (
FeedForward,
PatchEmbed,
RMSNorm,
TimestepEmbedder,
)
from genmo.mochi_preview.dit.joint_model.lora import LoraLinear
from genmo.mochi_preview.dit.joint_model.mod_rmsnorm import modulated_rmsnorm
from genmo.mochi_preview.dit.joint_model.residual_tanh_gated_rmsnorm import (
residual_tanh_gated_rmsnorm,
)
from genmo.mochi_preview.dit.joint_model.rope_mixed import (
compute_mixed_rotation,
create_position_matrix,
)
from genmo.mochi_preview.dit.joint_model.temporal_rope import apply_rotary_emb_qk_real
from genmo.mochi_preview.dit.joint_model.utils import (
AttentionPool,
modulate,
pad_and_split_xy,
)
COMPILE_FINAL_LAYER = os.environ.get("COMPILE_DIT") == "1"
COMPILE_MMDIT_BLOCK = os.environ.get("COMPILE_DIT") == "1"
def ck(fn, *args, enabled=True, **kwargs) -> torch.Tensor:
if enabled:
return torch.utils.checkpoint.checkpoint(fn, *args, **kwargs, use_reentrant=False)
return fn(*args, **kwargs)
class AsymmetricAttention(nn.Module):
def __init__(
self,
dim_x: int,
dim_y: int,
num_heads: int = 8,
qkv_bias: bool = True,
qk_norm: bool = False,
update_y: bool = True,
out_bias: bool = True,
attention_mode: str = "flash",
softmax_scale: Optional[float] = None,
device: Optional[torch.device] = None,
# Disable LoRA by default ...
qkv_proj_lora_rank: int = 0,
qkv_proj_lora_alpha: int = 0,
qkv_proj_lora_dropout: float = 0.0,
out_proj_lora_rank: int = 0,
out_proj_lora_alpha: int = 0,
out_proj_lora_dropout: float = 0.0,
):
super().__init__()
self.attention_mode = attention_mode
self.dim_x = dim_x
self.dim_y = dim_y
self.num_heads = num_heads
self.head_dim = dim_x // num_heads
self.update_y = update_y
self.softmax_scale = softmax_scale
if dim_x % num_heads != 0:
raise ValueError(f"dim_x={dim_x} should be divisible by num_heads={num_heads}")
# Input layers.
self.qkv_bias = qkv_bias
qkv_lora_kwargs = dict(
bias=qkv_bias,
device=device,
r=qkv_proj_lora_rank,
lora_alpha=qkv_proj_lora_alpha,
lora_dropout=qkv_proj_lora_dropout,
)
self.qkv_x = LoraLinear(dim_x, 3 * dim_x, **qkv_lora_kwargs)
# Project text features to match visual features (dim_y -> dim_x)
self.qkv_y = LoraLinear(dim_y, 3 * dim_x, **qkv_lora_kwargs)
# Query and key normalization for stability.
assert qk_norm
self.q_norm_x = RMSNorm(self.head_dim, device=device)
self.k_norm_x = RMSNorm(self.head_dim, device=device)
self.q_norm_y = RMSNorm(self.head_dim, device=device)
self.k_norm_y = RMSNorm(self.head_dim, device=device)
# Output layers. y features go back down from dim_x -> dim_y.
proj_lora_kwargs = dict(
bias=out_bias,
device=device,
r=out_proj_lora_rank,
lora_alpha=out_proj_lora_alpha,
lora_dropout=out_proj_lora_dropout,
)
self.proj_x = LoraLinear(dim_x, dim_x, **proj_lora_kwargs)
self.proj_y = LoraLinear(dim_x, dim_y, **proj_lora_kwargs) if update_y else nn.Identity()
def run_qkv_y(self, y):
cp_rank, cp_size = cp.get_cp_rank_size()
local_heads = self.num_heads // cp_size
if cp.is_cp_active():
# Only predict local heads.
assert not self.qkv_bias
W_qkv_y = self.qkv_y.weight.view(3, self.num_heads, self.head_dim, self.dim_y)
W_qkv_y = W_qkv_y.narrow(1, cp_rank * local_heads, local_heads)
W_qkv_y = W_qkv_y.reshape(3 * local_heads * self.head_dim, self.dim_y)
qkv_y = F.linear(y, W_qkv_y, None) # (B, L, 3 * local_h * head_dim)
else:
qkv_y = self.qkv_y(y) # (B, L, 3 * dim)
qkv_y = qkv_y.view(qkv_y.size(0), qkv_y.size(1), 3, local_heads, self.head_dim)
q_y, k_y, v_y = qkv_y.unbind(2)
q_y = self.q_norm_y(q_y)
k_y = self.k_norm_y(k_y)
return q_y, k_y, v_y
def prepare_qkv(
self,
x: torch.Tensor, # (B, M, dim_x)
y: torch.Tensor, # (B, L, dim_y)
*,
scale_x: torch.Tensor,
scale_y: torch.Tensor,
rope_cos: torch.Tensor,
rope_sin: torch.Tensor,
valid_token_indices: torch.Tensor,
max_seqlen_in_batch: int,
):
# Process visual features
x = modulated_rmsnorm(x, scale_x) # (B, M, dim_x) where M = N / cp_group_size
qkv_x = self.qkv_x(x) # (B, M, 3 * dim_x)
assert qkv_x.dtype == torch.bfloat16
qkv_x = cp.all_to_all_collect_tokens(qkv_x, self.num_heads) # (3, B, N, local_h, head_dim)
# Split qkv_x into q, k, v
q_x, k_x, v_x = qkv_x.unbind(0) # (B, N, local_h, head_dim)
q_x = self.q_norm_x(q_x)
q_x = apply_rotary_emb_qk_real(q_x, rope_cos, rope_sin)
k_x = self.k_norm_x(k_x)
k_x = apply_rotary_emb_qk_real(k_x, rope_cos, rope_sin)
# Concatenate streams
B, N, num_heads, head_dim = q_x.size()
D = num_heads * head_dim
# Process text features
if B == 1:
text_seqlen = max_seqlen_in_batch - N
if text_seqlen > 0:
y = y[:, :text_seqlen] # Remove padding tokens.
y = modulated_rmsnorm(y, scale_y) # (B, L, dim_y)
q_y, k_y, v_y = self.run_qkv_y(y) # (B, L, local_heads, head_dim)
q = torch.cat([q_x, q_y], dim=1)
k = torch.cat([k_x, k_y], dim=1)
v = torch.cat([v_x, v_y], dim=1)
else:
q, k, v = q_x, k_x, v_x
else:
y = modulated_rmsnorm(y, scale_y) # (B, L, dim_y)
q_y, k_y, v_y = self.run_qkv_y(y) # (B, L, local_heads, head_dim)
indices = valid_token_indices[:, None].expand(-1, D)
q = torch.cat([q_x, q_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
k = torch.cat([k_x, k_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
v = torch.cat([v_x, v_y], dim=1).view(-1, D).gather(0, indices) # (total, D)
q = q.view(-1, num_heads, head_dim)
k = k.view(-1, num_heads, head_dim)
v = v.view(-1, num_heads, head_dim)
return q, k, v
@torch.autocast("cuda", enabled=False)
def flash_attention(self, q, k, v, cu_seqlens, max_seqlen_in_batch, total, local_dim):
out: torch.Tensor = flash_varlen_attn(
q, k, v,
cu_seqlens_q=cu_seqlens,
cu_seqlens_k=cu_seqlens,
max_seqlen_q=max_seqlen_in_batch,
max_seqlen_k=max_seqlen_in_batch,
dropout_p=0.0,
softmax_scale=self.softmax_scale,
) # (total, local_heads, head_dim)
return out.view(total, local_dim)
def sdpa_attention(self, q, k, v):
with sdpa_attn_ctx(training=self.training):
out = F.scaled_dot_product_attention(
q, k, v,
attn_mask=None,
dropout_p=0.0,
is_causal=False,
)
return out
@torch.autocast("cuda", enabled=False)
def sage_attention(self, q, k, v):
return sage_attn(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=False)
def run_attention(
self,
q: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
k: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
v: torch.Tensor, # (total <= B * (N + L), num_heads, head_dim)
*,
B: int,
cu_seqlens: Optional[torch.Tensor] = None,
max_seqlen_in_batch: Optional[int] = None,
):
_, cp_size = cp.get_cp_rank_size()
assert self.num_heads % cp_size == 0
local_heads = self.num_heads // cp_size
local_dim = local_heads * self.head_dim
# Check shapes
assert q.ndim == 3 and k.ndim == 3 and v.ndim == 3
total = q.size(0)
assert k.size(0) == total and v.size(0) == total
if self.attention_mode == "flash":
out = self.flash_attention(
q, k, v, cu_seqlens, max_seqlen_in_batch, total, local_dim) # (total, local_dim)
else:
assert B == 1, \
f"Non-flash attention mode {self.attention_mode} only supports batch size 1, got {B}"
q = rearrange(q, "(b s) h d -> b h s d", b=B)
k = rearrange(k, "(b s) h d -> b h s d", b=B)
v = rearrange(v, "(b s) h d -> b h s d", b=B)
if self.attention_mode == "sdpa":
out = self.sdpa_attention(q, k, v) # (B, local_heads, seq_len, head_dim)
elif self.attention_mode == "sage":
out = self.sage_attention(q, k, v) # (B, local_heads, seq_len, head_dim)
else:
raise ValueError(f"Unknown attention mode: {self.attention_mode}")
out = rearrange(out, "b h s d -> (b s) (h d)")
return out
def post_attention(
self,
out: torch.Tensor,
B: int,
M: int,
L: int,
dtype: torch.dtype,
valid_token_indices: torch.Tensor,
):
"""
Args:
out: (total <= B * (N + L), local_dim)
valid_token_indices: (total <= B * (N + L),)
B: Batch size
M: Number of visual tokens per context parallel rank
L: Number of text tokens
dtype: Data type of the input and output tensors
Returns:
x: (B, N, dim_x) tensor of visual tokens where N = M * cp_size
y: (B, L, dim_y) tensor of text token features
"""
_, cp_size = cp.get_cp_rank_size()
local_heads = self.num_heads // cp_size
local_dim = local_heads * self.head_dim
N = M * cp_size
# Split sequence into visual and text tokens, adding back padding.
if B == 1:
out = out.view(B, -1, local_dim)
if out.size(1) > N:
x, y = torch.tensor_split(out, (N,), dim=1) # (B, N, local_dim), (B, <= L, local_dim)
y = F.pad(y, (0, 0, 0, L - y.size(1))) # (B, L, local_dim)
else:
# Empty prompt.
x, y = out, out.new_zeros(B, L, local_dim)
else:
x, y = pad_and_split_xy(out, valid_token_indices, B, N, L, dtype)
assert x.size() == (B, N, local_dim)
assert y.size() == (B, L, local_dim)
# Communicate across context parallel ranks.
x = x.view(B, N, local_heads, self.head_dim)
x = cp.all_to_all_collect_heads(x) # (B, M, dim_x = num_heads * head_dim)
if cp.is_cp_active():
y = cp.all_gather(y) # (cp_size * B, L, local_heads * head_dim)
y = rearrange(y, "(G B) L D -> B L (G D)", G=cp_size, D=local_dim) # (B, L, dim_x)
x = self.proj_x(x)
y = self.proj_y(y)
return x, y
def forward(
self,
x: torch.Tensor, # (B, M, dim_x)
y: torch.Tensor, # (B, L, dim_y)
*,
scale_x: torch.Tensor, # (B, dim_x), modulation for pre-RMSNorm.
scale_y: torch.Tensor, # (B, dim_y), modulation for pre-RMSNorm.
packed_indices: Dict[str, torch.Tensor] = None,
checkpoint_qkv: bool = False,
checkpoint_post_attn: bool = False,
**rope_rotation,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Forward pass of asymmetric multi-modal attention.
Args:
x: (B, M, dim_x) tensor of visual tokens
y: (B, L, dim_y) tensor of text token features
packed_indices: Dict with keys for Flash Attention
num_frames: Number of frames in the video. N = num_frames * num_spatial_tokens
Returns:
x: (B, M, dim_x) tensor of visual tokens after multi-modal attention
y: (B, L, dim_y) tensor of text token features after multi-modal attention
"""
B, L, _ = y.shape
_, M, _ = x.shape
# Predict a packed QKV tensor from visual and text features.
q, k, v = ck(self.prepare_qkv,
x=x,
y=y,
scale_x=scale_x,
scale_y=scale_y,
rope_cos=rope_rotation.get("rope_cos"),
rope_sin=rope_rotation.get("rope_sin"),
valid_token_indices=packed_indices["valid_token_indices_kv"],
max_seqlen_in_batch=packed_indices["max_seqlen_in_batch_kv"],
enabled=checkpoint_qkv,
) # (total <= B * (N + L), 3, local_heads, head_dim)
# Self-attention is expensive, so don't checkpoint it.
out = self.run_attention(
q, k, v, B=B,
cu_seqlens=packed_indices["cu_seqlens_kv"],
max_seqlen_in_batch=packed_indices["max_seqlen_in_batch_kv"],
)
x, y = ck(self.post_attention,
out,
B=B, M=M, L=L,
dtype=v.dtype,
valid_token_indices=packed_indices["valid_token_indices_kv"],
enabled=checkpoint_post_attn,
)
return x, y
@torch.compile(disable=not COMPILE_MMDIT_BLOCK)
class AsymmetricJointBlock(nn.Module):
def __init__(
self,
hidden_size_x: int,
hidden_size_y: int,
num_heads: int,
*,
mlp_ratio_x: float = 8.0, # Ratio of hidden size to d_model for MLP for visual tokens.
mlp_ratio_y: float = 4.0, # Ratio of hidden size to d_model for MLP for text tokens.
update_y: bool = True, # Whether to update text tokens in this block.
device: Optional[torch.device] = None,
**block_kwargs,
):
super().__init__()
self.update_y = update_y
self.hidden_size_x = hidden_size_x
self.hidden_size_y = hidden_size_y
self.mod_x = nn.Linear(hidden_size_x, 4 * hidden_size_x, device=device)
if self.update_y:
self.mod_y = nn.Linear(hidden_size_x, 4 * hidden_size_y, device=device)
else:
self.mod_y = nn.Linear(hidden_size_x, hidden_size_y, device=device)
# Self-attention:
self.attn = AsymmetricAttention(
hidden_size_x,
hidden_size_y,
num_heads=num_heads,
update_y=update_y,
device=device,
**block_kwargs,
)
# MLP.
mlp_hidden_dim_x = int(hidden_size_x * mlp_ratio_x)
assert mlp_hidden_dim_x == int(1536 * 8)
self.mlp_x = FeedForward(
in_features=hidden_size_x,
hidden_size=mlp_hidden_dim_x,
multiple_of=256,
ffn_dim_multiplier=None,
device=device,
)
# MLP for text not needed in last block.
if self.update_y:
mlp_hidden_dim_y = int(hidden_size_y * mlp_ratio_y)
self.mlp_y = FeedForward(
in_features=hidden_size_y,
hidden_size=mlp_hidden_dim_y,
multiple_of=256,
ffn_dim_multiplier=None,
device=device,
)
def forward(
self,
x: torch.Tensor,
c: torch.Tensor,
y: torch.Tensor,
# TODO: These could probably just go into attn_kwargs
checkpoint_ff: bool = False,
checkpoint_qkv: bool = False,
checkpoint_post_attn: bool = False,
**attn_kwargs,
):
"""Forward pass of a block.
Args:
x: (B, N, dim) tensor of visual tokens
c: (B, dim) tensor of conditioned features
y: (B, L, dim) tensor of text tokens
num_frames: Number of frames in the video. N = num_frames * num_spatial_tokens
Returns:
x: (B, N, dim) tensor of visual tokens after block
y: (B, L, dim) tensor of text tokens after block
"""
N = x.size(1)
c = F.silu(c)
mod_x = self.mod_x(c)
scale_msa_x, gate_msa_x, scale_mlp_x, gate_mlp_x = mod_x.chunk(4, dim=1)
mod_y = self.mod_y(c)
if self.update_y:
scale_msa_y, gate_msa_y, scale_mlp_y, gate_mlp_y = mod_y.chunk(4, dim=1)
else:
scale_msa_y = mod_y
# Self-attention block.
x_attn, y_attn = self.attn(
x,
y,
scale_x=scale_msa_x,
scale_y=scale_msa_y,
checkpoint_qkv=checkpoint_qkv,
checkpoint_post_attn=checkpoint_post_attn,
**attn_kwargs,
)
assert x_attn.size(1) == N
x = residual_tanh_gated_rmsnorm(x, x_attn, gate_msa_x)
if self.update_y:
y = residual_tanh_gated_rmsnorm(y, y_attn, gate_msa_y)
# MLP block.
x = ck(self.ff_block_x, x, scale_mlp_x, gate_mlp_x, enabled=checkpoint_ff)
if self.update_y:
y = ck(self.ff_block_y, y, scale_mlp_y, gate_mlp_y, enabled=checkpoint_ff) # type: ignore
return x, y
def ff_block_x(self, x, scale_x, gate_x):
x_mod = modulated_rmsnorm(x, scale_x)
x_res = self.mlp_x(x_mod)
x = residual_tanh_gated_rmsnorm(x, x_res, gate_x) # Sandwich norm
return x
def ff_block_y(self, y, scale_y, gate_y):
y_mod = modulated_rmsnorm(y, scale_y)
y_res = self.mlp_y(y_mod)
y = residual_tanh_gated_rmsnorm(y, y_res, gate_y) # Sandwich norm
return y
@torch.compile(disable=not COMPILE_FINAL_LAYER)
class FinalLayer(nn.Module):
"""
The final layer of DiT.
"""
def __init__(
self,
hidden_size,
patch_size,
out_channels,
device: Optional[torch.device] = None,
):
super().__init__()
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, device=device)
self.mod = nn.Linear(hidden_size, 2 * hidden_size, device=device)
self.linear = nn.Linear(hidden_size, patch_size * patch_size * out_channels, device=device)
def forward(self, x, c):
c = F.silu(c)
shift, scale = self.mod(c).chunk(2, dim=1)
x = modulate(self.norm_final(x), shift, scale)
x = self.linear(x)
return x
class AsymmDiTJoint(nn.Module):
"""
Diffusion model with a Transformer backbone.
Ingests text embeddings instead of a label.
"""
def __init__(
self,
*,
patch_size=2,
in_channels=4,
hidden_size_x=1152,
hidden_size_y=1152,
depth=48,
num_heads=16,
mlp_ratio_x=8.0,
mlp_ratio_y=4.0,
t5_feat_dim: int = 4096,
t5_token_length: int = 256,
patch_embed_bias: bool = True,
timestep_mlp_bias: bool = True,
timestep_scale: Optional[float] = None,
use_extended_posenc: bool = False,
rope_theta: float = 10000.0,
device: Optional[torch.device] = None,
**block_kwargs,
):
super().__init__()
self.in_channels = in_channels
self.out_channels = in_channels
self.patch_size = patch_size
self.num_heads = num_heads
self.hidden_size_x = hidden_size_x
self.hidden_size_y = hidden_size_y
self.head_dim = hidden_size_x // num_heads # Head dimension and count is determined by visual.
self.use_extended_posenc = use_extended_posenc
self.t5_token_length = t5_token_length
self.t5_feat_dim = t5_feat_dim
self.rope_theta = rope_theta # Scaling factor for frequency computation for temporal RoPE.
self.x_embedder = PatchEmbed(
patch_size=patch_size,
in_chans=in_channels,
embed_dim=hidden_size_x,
bias=patch_embed_bias,
device=device,
)
# Conditionings
# Timestep
self.t_embedder = TimestepEmbedder(hidden_size_x, bias=timestep_mlp_bias, timestep_scale=timestep_scale)
# Caption Pooling (T5)
self.t5_y_embedder = AttentionPool(t5_feat_dim, num_heads=8, output_dim=hidden_size_x, device=device)
# Dense Embedding Projection (T5)
self.t5_yproj = nn.Linear(t5_feat_dim, hidden_size_y, bias=True, device=device)
# Initialize pos_frequencies as an empty parameter.
self.pos_frequencies = nn.Parameter(torch.empty(3, self.num_heads, self.head_dim // 2, device=device))
# for depth 48:
# b = 0: AsymmetricJointBlock, update_y=True
# b = 1: AsymmetricJointBlock, update_y=True
# ...
# b = 46: AsymmetricJointBlock, update_y=True
# b = 47: AsymmetricJointBlock, update_y=False. No need to update text features.
blocks = []
for b in range(depth):
# Joint multi-modal block
update_y = b < depth - 1
block = AsymmetricJointBlock(
hidden_size_x,
hidden_size_y,
num_heads,
mlp_ratio_x=mlp_ratio_x,
mlp_ratio_y=mlp_ratio_y,
update_y=update_y,
device=device,
**block_kwargs,
)
blocks.append(block)
self.blocks = nn.ModuleList(blocks)
self.final_layer = FinalLayer(hidden_size_x, patch_size, self.out_channels, device=device)
def embed_x(self, x: torch.Tensor) -> torch.Tensor:
"""
Args:
x: (B, C=12, T, H, W) tensor of visual tokens
Returns:
x: (B, C=3072, N) tensor of visual tokens with positional embedding.
"""
return self.x_embedder(x) # Convert BcTHW to BCN
@torch.compile(disable=not COMPILE_MMDIT_BLOCK)
def prepare(
self,
x: torch.Tensor,
sigma: torch.Tensor,
t5_feat: torch.Tensor,
t5_mask: torch.Tensor,
):
"""Prepare input and conditioning embeddings."""
# Visual patch embeddings with positional encoding.
T, H, W = x.shape[-3:]
pH, pW = H // self.patch_size, W // self.patch_size
x = self.embed_x(x) # (B, N, D), where N = T * H * W / patch_size ** 2
assert x.ndim == 3
B = x.size(0)
# Construct position array of size [N, 3].
# pos[:, 0] is the frame index for each location,
# pos[:, 1] is the row index for each location, and
# pos[:, 2] is the column index for each location.
N = T * pH * pW
assert x.size(1) == N
pos = create_position_matrix(T, pH=pH, pW=pW, device=x.device, dtype=torch.float32) # (N, 3)
rope_cos, rope_sin = compute_mixed_rotation(
freqs=self.pos_frequencies, pos=pos
) # Each are (N, num_heads, dim // 2)
# Global vector embedding for conditionings.
c_t = self.t_embedder(1 - sigma) # (B, D)
# Pool T5 tokens using attention pooler
# Note y_feat[1] contains T5 token features.
assert (
t5_feat.size(1) == self.t5_token_length
), f"Expected L={self.t5_token_length}, got {t5_feat.shape} for y_feat."
t5_y_pool = self.t5_y_embedder(t5_feat, t5_mask) # (B, D)
assert t5_y_pool.size(0) == B, f"Expected B={B}, got {t5_y_pool.shape} for t5_y_pool."
c = c_t + t5_y_pool
y_feat = self.t5_yproj(t5_feat) # (B, L, t5_feat_dim) --> (B, L, D)
return x, c, y_feat, rope_cos, rope_sin
def forward(
self,
x: torch.Tensor, # [1, 12, 7, 60, 106]
sigma: torch.Tensor, # [1]
y_feat: List[torch.Tensor], # [0][1, 256, 4096]
y_mask: List[torch.Tensor], # [0][1, 256]
packed_indices: Dict[str, torch.Tensor] = None,
rope_cos: torch.Tensor = None,
rope_sin: torch.Tensor = None,
num_ff_checkpoint: int = 0, # 48
num_qkv_checkpoint: int = 0, # 48
num_post_attn_checkpoint: int = 0, # 0
):
"""Forward pass of DiT.
Args:
x: (B, C, T, H, W) tensor of spatial inputs (images or latent representations of images)
sigma: (B,) tensor of noise standard deviations
y_feat: List((B, L, y_feat_dim) tensor of caption token features. For SDXL text encoders: L=77, y_feat_dim=2048)
y_mask: List((B, L) boolean tensor indicating which tokens are not padding)
packed_indices: Dict with keys for Flash Attention. Result of compute_packed_indices.
"""
_, _, T, H, W = x.shape
if self.pos_frequencies.dtype != torch.float32:
warnings.warn(f"pos_frequencies dtype {self.pos_frequencies.dtype} != torch.float32")
# Use EFFICIENT_ATTENTION backend for T5 pooling, since we have a mask.
# Have to call sdpa_kernel outside of a torch.compile region.
with sdpa_kernel(torch.nn.attention.SDPBackend.EFFICIENT_ATTENTION):
x, c, y_feat, rope_cos, rope_sin = self.prepare(x, sigma, y_feat[0], y_mask[0]) # [1, 11130, 3072], [1, 3072], [1, 256, 1536], [11130, 24, 64]
del y_mask
cp_rank, cp_size = cp.get_cp_rank_size()
N = x.size(1)
M = N // cp_size
assert N % cp_size == 0, f"Visual sequence length ({x.shape[1]}) must be divisible by cp_size ({cp_size})."
if cp_size > 1:
x = x.narrow(1, cp_rank * M, M)
assert self.num_heads % cp_size == 0
local_heads = self.num_heads // cp_size
rope_cos = rope_cos.narrow(1, cp_rank * local_heads, local_heads)
rope_sin = rope_sin.narrow(1, cp_rank * local_heads, local_heads)
for i, block in enumerate(self.blocks):
x, y_feat = block( # [1, 11130, 3072], [1, 256, 1536]
x,
c,
y_feat,
rope_cos=rope_cos,
rope_sin=rope_sin,
packed_indices=packed_indices,
checkpoint_ff=i < num_ff_checkpoint,
checkpoint_qkv=i < num_qkv_checkpoint,
checkpoint_post_attn=i < num_post_attn_checkpoint,
) # (B, M, D), (B, L, D)
del y_feat # Final layers don't use dense text features.
x = self.final_layer(x, c) # (B, M, patch_size ** 2 * out_channels) [1, 11130, 48]
patch = x.size(2)
x = cp.all_gather(x)
x = rearrange(x, "(G B) M P -> B (G M) P", G=cp_size, P=patch) # [1, 11130, 48]
x = rearrange( # [1, 12, 7, 60, 106]
x,
"B (T hp wp) (p1 p2 c) -> B c T (hp p1) (wp p2)",
T=T,
hp=H // self.patch_size,
wp=W // self.patch_size,
p1=self.patch_size,
p2=self.patch_size,
c=self.out_channels,
)
return x
@@ -0,0 +1,158 @@
from typing import Tuple
import torch
import torch.distributed as dist
from einops import rearrange
_CONTEXT_PARALLEL_GROUP = None
_CONTEXT_PARALLEL_RANK = None
_CONTEXT_PARALLEL_GROUP_SIZE = None
_CONTEXT_PARALLEL_GROUP_RANKS = None
def get_cp_rank_size() -> Tuple[int, int]:
if _CONTEXT_PARALLEL_GROUP:
assert isinstance(_CONTEXT_PARALLEL_RANK, int) and isinstance(_CONTEXT_PARALLEL_GROUP_SIZE, int)
return _CONTEXT_PARALLEL_RANK, _CONTEXT_PARALLEL_GROUP_SIZE
else:
return 0, 1
def local_shard(x: torch.Tensor, dim: int = 2) -> torch.Tensor:
if not _CONTEXT_PARALLEL_GROUP:
return x
cp_rank, cp_size = get_cp_rank_size()
return x.tensor_split(cp_size, dim=dim)[cp_rank]
def set_cp_group(cp_group, ranks, global_rank):
global _CONTEXT_PARALLEL_GROUP, _CONTEXT_PARALLEL_RANK, _CONTEXT_PARALLEL_GROUP_SIZE, _CONTEXT_PARALLEL_GROUP_RANKS
if _CONTEXT_PARALLEL_GROUP is not None:
raise RuntimeError("CP group already initialized.")
_CONTEXT_PARALLEL_GROUP = cp_group
_CONTEXT_PARALLEL_RANK = dist.get_rank(cp_group)
_CONTEXT_PARALLEL_GROUP_SIZE = dist.get_world_size(cp_group)
_CONTEXT_PARALLEL_GROUP_RANKS = ranks
assert _CONTEXT_PARALLEL_RANK == ranks.index(
global_rank
), f"Rank mismatch: {global_rank} in {ranks} does not have position {_CONTEXT_PARALLEL_RANK} "
assert _CONTEXT_PARALLEL_GROUP_SIZE == len(
ranks
), f"Group size mismatch: {_CONTEXT_PARALLEL_GROUP_SIZE} != len({ranks})"
def get_cp_group():
if _CONTEXT_PARALLEL_GROUP is None:
raise RuntimeError("CP group not initialized")
return _CONTEXT_PARALLEL_GROUP
def is_cp_active():
return _CONTEXT_PARALLEL_GROUP is not None
class AllGatherIntoTensorFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, x: torch.Tensor, reduce_dtype, group: dist.ProcessGroup):
ctx.reduce_dtype = reduce_dtype
ctx.group = group
ctx.batch_size = x.size(0)
group_size = dist.get_world_size(group)
x = x.contiguous()
output = torch.empty(group_size * x.size(0), *x.shape[1:], dtype=x.dtype, device=x.device)
dist.all_gather_into_tensor(output, x, group=group)
return output
def all_gather(tensor: torch.Tensor) -> torch.Tensor:
if not _CONTEXT_PARALLEL_GROUP:
return tensor
return AllGatherIntoTensorFunction.apply(tensor, torch.float32, _CONTEXT_PARALLEL_GROUP)
@torch.compiler.disable()
def _all_to_all_single(output, input, group):
# Disable compilation since torch compile changes contiguity.
assert input.is_contiguous(), "Input tensor must be contiguous."
assert output.is_contiguous(), "Output tensor must be contiguous."
return dist.all_to_all_single(output, input, group=group)
class CollectTokens(torch.autograd.Function):
@staticmethod
def forward(ctx, qkv: torch.Tensor, group: dist.ProcessGroup, num_heads: int):
"""Redistribute heads and receive tokens.
Args:
qkv: query, key or value. Shape: [B, M, 3 * num_heads * head_dim]
Returns:
qkv: shape: [3, B, N, local_heads, head_dim]
where M is the number of local tokens,
N = cp_size * M is the number of global tokens,
local_heads = num_heads // cp_size is the number of local heads.
"""
ctx.group = group
ctx.num_heads = num_heads
cp_size = dist.get_world_size(group)
assert num_heads % cp_size == 0
ctx.local_heads = num_heads // cp_size
qkv = rearrange(
qkv,
"B M (qkv G h d) -> G M h B (qkv d)",
qkv=3,
G=cp_size,
h=ctx.local_heads,
).contiguous()
output_chunks = torch.empty_like(qkv)
_all_to_all_single(output_chunks, qkv, group=group)
return rearrange(output_chunks, "G M h B (qkv d) -> qkv B (G M) h d", qkv=3)
def all_to_all_collect_tokens(x: torch.Tensor, num_heads: int) -> torch.Tensor:
if not _CONTEXT_PARALLEL_GROUP:
# Move QKV dimension to the front.
# B M (3 H d) -> 3 B M H d
B, M, _ = x.size()
x = x.view(B, M, 3, num_heads, -1)
return x.permute(2, 0, 1, 3, 4)
return CollectTokens.apply(x, _CONTEXT_PARALLEL_GROUP, num_heads)
class CollectHeads(torch.autograd.Function):
@staticmethod
def forward(ctx, x: torch.Tensor, group: dist.ProcessGroup):
"""Redistribute tokens and receive heads.
Args:
x: Output of attention. Shape: [B, N, local_heads, head_dim]
Returns:
Shape: [B, M, num_heads * head_dim]
"""
ctx.group = group
ctx.local_heads = x.size(2)
ctx.head_dim = x.size(3)
group_size = dist.get_world_size(group)
x = rearrange(x, "B (G M) h D -> G h M B D", G=group_size).contiguous()
output = torch.empty_like(x)
_all_to_all_single(output, x, group=group)
del x
return rearrange(output, "G h M B D -> B M (G h D)")
def all_to_all_collect_heads(x: torch.Tensor) -> torch.Tensor:
if not _CONTEXT_PARALLEL_GROUP:
# Merge heads.
return x.view(x.size(0), x.size(1), x.size(2) * x.size(3))
return CollectHeads.apply(x, _CONTEXT_PARALLEL_GROUP)
@@ -0,0 +1,179 @@
import collections.abc
import math
from itertools import repeat
from typing import Callable, Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
# From PyTorch internals
def _ntuple(n):
def parse(x):
if isinstance(x, collections.abc.Iterable) and not isinstance(x, str):
return tuple(x)
return tuple(repeat(x, n))
return parse
to_2tuple = _ntuple(2)
class TimestepEmbedder(nn.Module):
def __init__(
self,
hidden_size: int,
frequency_embedding_size: int = 256,
*,
bias: bool = True,
timestep_scale: Optional[float] = None,
device: Optional[torch.device] = None,
):
super().__init__()
self.mlp = nn.Sequential(
nn.Linear(frequency_embedding_size, hidden_size, bias=bias, device=device),
nn.SiLU(),
nn.Linear(hidden_size, hidden_size, bias=bias, device=device),
)
self.frequency_embedding_size = frequency_embedding_size
self.timestep_scale = timestep_scale
@staticmethod
def timestep_embedding(t, dim, max_period=10000):
half = dim // 2
freqs = torch.arange(start=0, end=half, dtype=torch.float32, device=t.device)
freqs.mul_(-math.log(max_period) / half).exp_()
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
def forward(self, t):
if self.timestep_scale is not None:
t = t * self.timestep_scale
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
t_emb = self.mlp(t_freq)
return t_emb
class PooledCaptionEmbedder(nn.Module):
def __init__(
self,
caption_feature_dim: int,
hidden_size: int,
*,
bias: bool = True,
device: Optional[torch.device] = None,
):
super().__init__()
self.caption_feature_dim = caption_feature_dim
self.hidden_size = hidden_size
self.mlp = nn.Sequential(
nn.Linear(caption_feature_dim, hidden_size, bias=bias, device=device),
nn.SiLU(),
nn.Linear(hidden_size, hidden_size, bias=bias, device=device),
)
def forward(self, x):
return self.mlp(x)
class FeedForward(nn.Module):
def __init__(
self,
in_features: int,
hidden_size: int,
multiple_of: int,
ffn_dim_multiplier: Optional[float],
device: Optional[torch.device] = None,
):
super().__init__()
# keep parameter count and computation constant compared to standard FFN
hidden_size = int(2 * hidden_size / 3)
# custom dim factor multiplier
if ffn_dim_multiplier is not None:
hidden_size = int(ffn_dim_multiplier * hidden_size)
hidden_size = multiple_of * ((hidden_size + multiple_of - 1) // multiple_of)
self.hidden_dim = hidden_size
self.w1 = nn.Linear(in_features, 2 * hidden_size, bias=False, device=device)
self.w2 = nn.Linear(hidden_size, in_features, bias=False, device=device)
def forward(self, x):
# assert self.w1.weight.dtype == torch.bfloat16, f"FFN weight dtype {self.w1.weight.dtype} != bfloat16"
x, gate = self.w1(x).chunk(2, dim=-1)
x = self.w2(F.silu(x) * gate)
return x
class PatchEmbed(nn.Module):
def __init__(
self,
patch_size: int = 16,
in_chans: int = 3,
embed_dim: int = 768,
norm_layer: Optional[Callable] = None,
flatten: bool = True,
bias: bool = True,
dynamic_img_pad: bool = False,
device: Optional[torch.device] = None,
):
super().__init__()
self.patch_size = to_2tuple(patch_size)
self.flatten = flatten
self.dynamic_img_pad = dynamic_img_pad
self.proj = nn.Conv2d(
in_chans,
embed_dim,
kernel_size=patch_size,
stride=patch_size,
bias=bias,
device=device,
)
assert norm_layer is None
self.norm = norm_layer(embed_dim, device=device) if norm_layer else nn.Identity()
def forward(self, x):
B, _C, T, H, W = x.shape
if not self.dynamic_img_pad:
assert (
H % self.patch_size[0] == 0
), f"Input height ({H}) should be divisible by patch size ({self.patch_size[0]})."
assert (
W % self.patch_size[1] == 0
), f"Input width ({W}) should be divisible by patch size ({self.patch_size[1]})."
else:
pad_h = (self.patch_size[0] - H % self.patch_size[0]) % self.patch_size[0]
pad_w = (self.patch_size[1] - W % self.patch_size[1]) % self.patch_size[1]
x = F.pad(x, (0, pad_w, 0, pad_h))
x = rearrange(x, "B C T H W -> (B T) C H W", B=B, T=T)
x = self.proj(x)
# Flatten temporal and spatial dimensions.
if not self.flatten:
raise NotImplementedError("Must flatten output.")
x = rearrange(x, "(B T) C H W -> B (T H W) C", B=B, T=T)
x = self.norm(x)
return x
class RMSNorm(torch.nn.Module):
def __init__(self, hidden_size, eps=1e-5, device=None):
super().__init__()
self.eps = eps
self.weight = torch.nn.Parameter(torch.empty(hidden_size, device=device))
self.register_parameter("bias", None)
def forward(self, x):
# assert self.weight.dtype == torch.float32, f"RMSNorm weight dtype {self.weight.dtype} != float32"
x_fp32 = x.float()
x_normed = x_fp32 * torch.rsqrt(x_fp32.pow(2).mean(-1, keepdim=True) + self.eps)
return (x_normed * self.weight).type_as(x)
@@ -0,0 +1,112 @@
#! /usr/bin/env python3
import math
from typing import Dict, List, Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
class LoRALayer:
def __init__(
self,
r: int,
lora_alpha: int,
lora_dropout: float,
merge_weights: bool,
):
self.r = r
self.lora_alpha = lora_alpha
if lora_dropout > 0.0:
self.lora_dropout = nn.Dropout(p=lora_dropout)
else:
self.lora_dropout = lambda x: x
self.merged = False
self.merge_weights = merge_weights
def mark_only_lora_as_trainable(model: nn.Module, bias: str = "none") -> None:
assert bias == "none", f"Only bias='none' is supported"
for n, p in model.named_parameters():
if "lora_" not in n:
p.requires_grad = False
def lora_state_dict(model: nn.Module, bias: str = "none") -> Dict[str, torch.Tensor]:
assert bias == "none", f"Only bias='none' is supported"
my_state_dict = model.state_dict()
return {k: my_state_dict[k] for k in my_state_dict if "lora_" in k}
class LoraLinear(nn.Linear, LoRALayer):
# LoRA implemented in a dense layer
def __init__(
self,
in_features: int,
out_features: int,
r: int = 0,
lora_alpha: int = 1,
lora_dropout: float = 0.0,
fan_in_fan_out: bool = False, # Set this to True if the layer to replace stores weight like (fan_in, fan_out)
merge_weights: bool = True,
**kwargs,
):
nn.Linear.__init__(self, in_features, out_features, **kwargs)
LoRALayer.__init__(self, r=r, lora_alpha=lora_alpha, lora_dropout=lora_dropout, merge_weights=merge_weights)
self.fan_in_fan_out = fan_in_fan_out
# Actual trainable parameters
if r > 0:
self.lora_A = nn.Parameter(self.weight.new_zeros((r, in_features)).to(torch.float32))
self.lora_B = nn.Parameter(self.weight.new_zeros((out_features, r)).to(torch.float32))
self.scaling = self.lora_alpha / self.r
# Freezing the pre-trained weight matrix
self.weight.requires_grad = False
self.reset_parameters()
if fan_in_fan_out:
self.weight.data = self.weight.data.transpose(0, 1)
def reset_parameters(self):
nn.Linear.reset_parameters(self)
if hasattr(self, "lora_A"):
# initialize B the same way as the default for nn.Linear and A to zero
# this is different than what is described in the paper but should not affect performance
nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5))
nn.init.zeros_(self.lora_B)
def train(self, mode: bool = True):
def T(w):
return w.transpose(0, 1) if self.fan_in_fan_out else w
nn.Linear.train(self, mode)
if mode:
if self.merge_weights and self.merged:
# Make sure that the weights are not merged
if self.r > 0:
self.weight.data -= T(self.lora_B @ self.lora_A) * self.scaling
self.merged = False
else:
if self.merge_weights and not self.merged:
# Merge the weights and mark it
if self.r > 0:
self.weight.data += T(self.lora_B @ self.lora_A) * self.scaling
self.merged = True
def forward(self, x: torch.Tensor):
def T(w):
return w.transpose(0, 1) if self.fan_in_fan_out else w
if self.r > 0 and not self.merged:
result = F.linear(x, T(self.weight), bias=self.bias)
x = self.lora_dropout(x)
x = x @ self.lora_A.transpose(0, 1)
x = x @ self.lora_B.transpose(0, 1)
x = x * self.scaling
return result + x
else:
return F.linear(x, T(self.weight), bias=self.bias)
@@ -0,0 +1,15 @@
import torch
def modulated_rmsnorm(x, scale, eps=1e-6):
dtype = x.dtype
x = x.float()
# Compute RMS
mean_square = x.pow(2).mean(-1, keepdim=True)
inv_rms = torch.rsqrt(mean_square + eps)
# Normalize and modulate
x_normed = x * inv_rms
x_modulated = x_normed * (1 + scale.unsqueeze(1).float())
return x_modulated.to(dtype)
@@ -0,0 +1,20 @@
import torch
def residual_tanh_gated_rmsnorm(x, x_res, gate, eps=1e-6):
# Convert to fp32 for precision
x_res = x_res.float()
# Compute RMS
mean_square = x_res.pow(2).mean(-1, keepdim=True)
scale = torch.rsqrt(mean_square + eps)
# Apply tanh to gate
tanh_gate = torch.tanh(gate).unsqueeze(1)
# Normalize and apply gated scaling
x_normed = x_res * scale * tanh_gate
# Apply residual connection
output = x + x_normed.type_as(x)
return output
@@ -0,0 +1,88 @@
import functools
import math
import torch
def centers(start: float, stop, num, dtype=None, device=None):
"""linspace through bin centers.
Args:
start (float): Start of the range.
stop (float): End of the range.
num (int): Number of points.
dtype (torch.dtype): Data type of the points.
device (torch.device): Device of the points.
Returns:
centers (Tensor): Centers of the bins. Shape: (num,).
"""
edges = torch.linspace(start, stop, num + 1, dtype=dtype, device=device)
return (edges[:-1] + edges[1:]) / 2
@functools.lru_cache(maxsize=1)
def create_position_matrix(
T: int,
pH: int,
pW: int,
device: torch.device,
dtype: torch.dtype,
*,
target_area: float = 36864,
):
"""
Args:
T: int - Temporal dimension
pH: int - Height dimension after patchify
pW: int - Width dimension after patchify
Returns:
pos: [T * pH * pW, 3] - position matrix
"""
with torch.no_grad():
# Create 1D tensors for each dimension
t = torch.arange(T, dtype=dtype)
# Positionally interpolate to area 36864.
# (3072x3072 frame with 16x16 patches = 192x192 latents).
# This automatically scales rope positions when the resolution changes.
# We use a large target area so the model is more sensitive
# to changes in the learned pos_frequencies matrix.
scale = math.sqrt(target_area / (pW * pH))
w = centers(-pW * scale / 2, pW * scale / 2, pW)
h = centers(-pH * scale / 2, pH * scale / 2, pH)
# Use meshgrid to create 3D grids
grid_t, grid_h, grid_w = torch.meshgrid(t, h, w, indexing="ij")
# Stack and reshape the grids.
pos = torch.stack([grid_t, grid_h, grid_w], dim=-1) # [T, pH, pW, 3]
pos = pos.view(-1, 3) # [T * pH * pW, 3]
pos = pos.to(dtype=dtype, device=device)
return pos
def compute_mixed_rotation(
freqs: torch.Tensor,
pos: torch.Tensor,
):
"""
Project each 3-dim position into per-head, per-head-dim 1D frequencies.
Args:
freqs: [3, num_heads, num_freqs] - learned rotation frequency (for t, row, col) for each head position
pos: [N, 3] - position of each token
num_heads: int
Returns:
freqs_cos: [N, num_heads, num_freqs] - cosine components
freqs_sin: [N, num_heads, num_freqs] - sine components
"""
with torch.autocast("cuda", enabled=False):
assert freqs.ndim == 3
freqs_sum = torch.einsum("Nd,dhf->Nhf", pos.to(freqs), freqs)
freqs_cos = torch.cos(freqs_sum)
freqs_sin = torch.sin(freqs_sum)
return freqs_cos, freqs_sin
@@ -0,0 +1,34 @@
# Based on Llama3 Implementation.
import torch
def apply_rotary_emb_qk_real(
xqk: torch.Tensor,
freqs_cos: torch.Tensor,
freqs_sin: torch.Tensor,
) -> torch.Tensor:
"""
Apply rotary embeddings to input tensors using the given frequency tensor without complex numbers.
Args:
xqk (torch.Tensor): Query and/or Key tensors to apply rotary embeddings. Shape: (B, S, *, num_heads, D)
Can be either just query or just key, or both stacked along some batch or * dim.
freqs_cos (torch.Tensor): Precomputed cosine frequency tensor.
freqs_sin (torch.Tensor): Precomputed sine frequency tensor.
Returns:
torch.Tensor: The input tensor with rotary embeddings applied.
"""
assert xqk.dtype == torch.bfloat16
# Split the last dimension into even and odd parts
xqk_even = xqk[..., 0::2]
xqk_odd = xqk[..., 1::2]
# Apply rotation
cos_part = (xqk_even * freqs_cos - xqk_odd * freqs_sin).type_as(xqk)
sin_part = (xqk_even * freqs_sin + xqk_odd * freqs_cos).type_as(xqk)
# Interleave the results back into the original shape
out = torch.stack([cos_part, sin_part], dim=-1).flatten(-2)
assert out.dtype == torch.bfloat16
return out
@@ -0,0 +1,109 @@
from typing import Optional, Tuple
import torch
import torch.nn as nn
import torch.nn.functional as F
def modulate(x, shift, scale):
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
def pool_tokens(x: torch.Tensor, mask: torch.Tensor, *, keepdim=False) -> torch.Tensor:
"""
Pool tokens in x using mask.
NOTE: We assume x does not require gradients.
Args:
x: (B, L, D) tensor of tokens.
mask: (B, L) boolean tensor indicating which tokens are not padding.
Returns:
pooled: (B, D) tensor of pooled tokens.
"""
assert x.size(1) == mask.size(1) # Expected mask to have same length as tokens.
assert x.size(0) == mask.size(0) # Expected mask to have same batch size as tokens.
mask = mask[:, :, None].to(dtype=x.dtype)
mask = mask / mask.sum(dim=1, keepdim=True).clamp(min=1)
pooled = (x * mask).sum(dim=1, keepdim=keepdim)
return pooled
class AttentionPool(nn.Module):
def __init__(
self,
embed_dim: int,
num_heads: int,
output_dim: int = None,
device: Optional[torch.device] = None,
):
"""
Args:
spatial_dim (int): Number of tokens in sequence length.
embed_dim (int): Dimensionality of input tokens.
num_heads (int): Number of attention heads.
output_dim (int): Dimensionality of output tokens. Defaults to embed_dim.
"""
super().__init__()
self.num_heads = num_heads
self.to_kv = nn.Linear(embed_dim, 2 * embed_dim, device=device)
self.to_q = nn.Linear(embed_dim, embed_dim, device=device)
self.to_out = nn.Linear(embed_dim, output_dim or embed_dim, device=device)
def forward(self, x, mask):
"""
Args:
x (torch.Tensor): (B, L, D) tensor of input tokens.
mask (torch.Tensor): (B, L) boolean tensor indicating which tokens are not padding.
NOTE: We assume x does not require gradients.
Returns:
x (torch.Tensor): (B, D) tensor of pooled tokens.
"""
D = x.size(2)
# Construct attention mask, shape: (B, 1, num_queries=1, num_keys=1+L).
attn_mask = mask[:, None, None, :].bool() # (B, 1, 1, L).
attn_mask = F.pad(attn_mask, (1, 0), value=True) # (B, 1, 1, 1+L).
# Average non-padding token features. These will be used as the query.
x_pool = pool_tokens(x, mask, keepdim=True) # (B, 1, D)
# Concat pooled features to input sequence.
x = torch.cat([x_pool, x], dim=1) # (B, L+1, D)
# Compute queries, keys, values. Only the mean token is used to create a query.
kv = self.to_kv(x) # (B, L+1, 2 * D)
q = self.to_q(x[:, 0]) # (B, D)
# Extract heads.
head_dim = D // self.num_heads
kv = kv.unflatten(2, (2, self.num_heads, head_dim)) # (B, 1+L, 2, H, head_dim)
kv = kv.transpose(1, 3) # (B, H, 2, 1+L, head_dim)
k, v = kv.unbind(2) # (B, H, 1+L, head_dim)
q = q.unflatten(1, (self.num_heads, head_dim)) # (B, H, head_dim)
q = q.unsqueeze(2) # (B, H, 1, head_dim)
# Compute attention.
x = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask, dropout_p=0.0) # (B, H, 1, head_dim)
# Concatenate heads and run output.
x = x.squeeze(2).flatten(1, 2) # (B, D = H * head_dim)
x = self.to_out(x)
return x
def pad_and_split_xy(xy, indices, B, N, L, dtype) -> Tuple[torch.Tensor, torch.Tensor]:
D = xy.size(1)
# Pad sequences to (B, N + L, dim).
assert indices.ndim == 1
indices = indices.unsqueeze(1).expand(-1, D) # (total,) -> (total, num_heads * head_dim)
output = torch.zeros(B * (N + L), D, device=xy.device, dtype=dtype)
output = torch.scatter(output, 0, indices, xy)
xy = output.view(B, N + L, D)
# Split visual and text tokens along the sequence length.
return torch.tensor_split(xy, (N,), dim=1)
@@ -0,0 +1,682 @@
import json
import os
import random
from abc import ABC, abstractmethod
from contextlib import contextmanager
from functools import partial
from typing import Any, Dict, List, Literal, Optional, Union, cast
import numpy as np
import ray
import torch
import torch.distributed as dist
import torch.nn as nn
import torch.nn.functional as F
from einops import repeat
from safetensors import safe_open
from safetensors.torch import load_file
from torch import nn
from torch.distributed.fsdp import (
BackwardPrefetch,
MixedPrecision,
ShardingStrategy,
)
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp.wrap import (
lambda_auto_wrap_policy,
transformer_auto_wrap_policy,
)
from transformers import T5EncoderModel, T5Tokenizer
from transformers.models.t5.modeling_t5 import T5Block
import genmo.mochi_preview.dit.joint_model.context_parallel as cp
from genmo.lib.progress import get_new_progress_bar, progress_bar
from genmo.lib.utils import Timer
from genmo.mochi_preview.vae.models import (
Decoder,
Encoder,
decode_latents,
decode_latents_tiled_full,
decode_latents_tiled_spatial,
)
from genmo.mochi_preview.vae.vae_stats import dit_latents_to_vae_latents
def load_to_cpu(p, weights_only=True):
if p.endswith(".safetensors"):
return load_file(p)
else:
assert p.endswith(".pt")
return torch.load(p, map_location="cpu", weights_only=weights_only)
def linear_quadratic_schedule(num_steps, threshold_noise, linear_steps=None):
if linear_steps is None:
linear_steps = num_steps // 2
linear_sigma_schedule = [i * threshold_noise / linear_steps for i in range(linear_steps)]
threshold_noise_step_diff = linear_steps - threshold_noise * num_steps
quadratic_steps = num_steps - linear_steps
quadratic_coef = threshold_noise_step_diff / (linear_steps * quadratic_steps**2)
linear_coef = threshold_noise / linear_steps - 2 * threshold_noise_step_diff / (quadratic_steps**2)
const = quadratic_coef * (linear_steps**2)
quadratic_sigma_schedule = [
quadratic_coef * (i**2) + linear_coef * i + const for i in range(linear_steps, num_steps)
]
sigma_schedule = linear_sigma_schedule + quadratic_sigma_schedule + [1.0]
sigma_schedule = [1.0 - x for x in sigma_schedule]
return sigma_schedule
T5_MODEL = "google/t5-v1_1-xxl"
MAX_T5_TOKEN_LENGTH = 256
def setup_fsdp_sync(model, device_id, *, param_dtype, auto_wrap_policy) -> FSDP:
model = FSDP(
model,
sharding_strategy=ShardingStrategy.FULL_SHARD,
mixed_precision=MixedPrecision(
param_dtype=param_dtype,
reduce_dtype=torch.float32,
buffer_dtype=torch.float32,
),
auto_wrap_policy=auto_wrap_policy,
backward_prefetch=BackwardPrefetch.BACKWARD_PRE,
limit_all_gathers=True,
device_id=device_id,
sync_module_states=True,
use_orig_params=True,
)
torch.cuda.synchronize()
return model
class ModelFactory(ABC):
def __init__(self, **kwargs):
self.kwargs = kwargs
@abstractmethod
def get_model(self, *, local_rank: int, device_id: Union[int, Literal["cpu"]], world_size: int) -> Any:
assert isinstance(device_id, int) or device_id == "cpu", "device_id must be an integer or 'cpu'"
# FSDP does not work when the model is on the CPU
if device_id == "cpu":
assert world_size == 1, "CPU offload only supports single-GPU inference"
class T5ModelFactory(ModelFactory):
def __init__(self, model_dir=None):
super().__init__()
self.model_dir = model_dir or T5_MODEL
def get_model(self, *, local_rank, device_id, world_size):
super().get_model(local_rank=local_rank, device_id=device_id, world_size=world_size)
model = T5EncoderModel.from_pretrained(self.model_dir)
if world_size > 1:
model = setup_fsdp_sync(
model,
device_id=device_id,
param_dtype=torch.float32,
auto_wrap_policy=partial(
transformer_auto_wrap_policy,
transformer_layer_cls={
T5Block,
},
),
)
elif isinstance(device_id, int):
model = model.to(torch.device(f"cuda:{device_id}")) # type: ignore
return model.eval()
class DitModelFactory(ModelFactory):
def __init__(
self, *,
model_path: str,
model_dtype: str,
lora_path: Optional[str] = None,
attention_mode: Optional[str] = None
):
# Infer attention mode if not specified
if attention_mode is None:
from genmo.lib.attn_imports import flash_varlen_attn # type: ignore
attention_mode = "sdpa" if flash_varlen_attn is None else "flash"
print(f"Attention mode: {attention_mode}")
super().__init__(
model_path=model_path,
lora_path=lora_path,
model_dtype=model_dtype,
attention_mode=attention_mode
)
def get_model(
self,
*,
local_rank,
device_id,
world_size,
model_kwargs=None,
patch_model_fns=None,
strict_load=True,
load_checkpoint=True,
fast_init=True,
):
from genmo.mochi_preview.dit.joint_model.asymm_models_joint import AsymmDiTJoint
if not model_kwargs:
model_kwargs = {}
lora_sd = None
lora_path = self.kwargs["lora_path"]
if lora_path is not None:
if lora_path.endswith(".safetensors"):
lora_sd = {}
with safe_open(lora_path, framework="pt") as f:
for k in f.keys():
lora_sd[k] = f.get_tensor(k)
lora_kwargs = json.loads(f.metadata()["kwargs"])
print(f"Loaded LoRA kwargs: {lora_kwargs}")
else:
lora = load_to_cpu(lora_path, weights_only=False)
lora_sd, lora_kwargs = lora["state_dict"], lora["kwargs"]
model_kwargs.update(cast(dict, lora_kwargs))
model_args = dict(
depth=48,
patch_size=2,
num_heads=24,
hidden_size_x=3072,
hidden_size_y=1536,
mlp_ratio_x=4.0,
mlp_ratio_y=4.0,
in_channels=12,
qk_norm=True,
qkv_bias=False,
out_bias=True,
patch_embed_bias=True,
timestep_mlp_bias=True,
timestep_scale=1000.0,
t5_feat_dim=4096,
t5_token_length=256,
rope_theta=10000.0,
attention_mode=self.kwargs["attention_mode"],
**model_kwargs,
)
if fast_init:
model: nn.Module = torch.nn.utils.skip_init(AsymmDiTJoint, **model_args)
else:
model: nn.Module = AsymmDiTJoint(**model_args)
for fn in patch_model_fns or []:
model = fn(model)
# FSDP syncs weights from rank 0 to all other ranks
if local_rank == 0 and load_checkpoint:
model_path = self.kwargs["model_path"]
sd = load_to_cpu(model_path)
# Load the state dictionary and capture the return value
load_result = model.load_state_dict(sd, strict=strict_load)
if not strict_load:
# Print mismatched keys
missing_keys = [k for k in load_result.missing_keys if ".lora_" not in k]
if missing_keys:
print(f"Missing keys from {model_path}: {missing_keys}")
if load_result.unexpected_keys:
print(f"Unexpected keys from {model_path}: {load_result.unexpected_keys}")
if lora_sd:
model.load_state_dict(lora_sd, strict=strict_load) # type: ignore
if world_size > 1:
assert self.kwargs["model_dtype"] == "bf16", "FP8 is not supported for multi-GPU inference"
model = setup_fsdp_sync(
model,
device_id=device_id,
param_dtype=torch.float32,
auto_wrap_policy=partial(
lambda_auto_wrap_policy,
lambda_fn=lambda m: m in model.blocks,
),
)
elif isinstance(device_id, int):
model = model.to(torch.device(f"cuda:{device_id}"))
return model.eval()
class DecoderModelFactory(ModelFactory):
def __init__(self, *, model_path: str):
super().__init__(model_path=model_path)
def get_model(self, *, local_rank=0, device_id=0, world_size=1):
# TODO(ved): Set flag for torch.compile
# TODO(ved): Use skip_init
decoder = Decoder(
out_channels=3,
base_channels=128,
channel_multipliers=[1, 2, 4, 6],
temporal_expansions=[1, 2, 3],
spatial_expansions=[2, 2, 2],
num_res_blocks=[3, 3, 4, 6, 3],
latent_dim=12,
has_attention=[False, False, False, False, False],
output_norm=False,
nonlinearity="silu",
output_nonlinearity="silu",
causal=True,
)
# VAE is not FSDP-wrapped
state_dict = load_file(self.kwargs["model_path"])
decoder.load_state_dict(state_dict, strict=True)
device = torch.device(f"cuda:{device_id}") if isinstance(device_id, int) else "cpu"
decoder.eval().to(device)
return decoder
class EncoderModelFactory(ModelFactory):
def __init__(self, *, model_path: str):
super().__init__(model_path=model_path)
def get_model(self, *, local_rank=0, device_id=0, world_size=1):
# TODO(ved): Set flag for torch.compile
# TODO(ved): Use skip_init
# We don't FSDP the encoder b/c it is small
encoder = Encoder(
in_channels=15,
base_channels=64,
channel_multipliers=[1, 2, 4, 6],
num_res_blocks=[3, 3, 4, 6, 3],
latent_dim=12,
temporal_reductions=[1, 2, 3],
spatial_reductions=[2, 2, 2],
prune_bottlenecks=[False, False, False, False, False],
has_attentions=[False, True, True, True, True],
affine=True,
bias=True,
input_is_conv_1x1=True,
padding_mode="replicate",
)
state_dict = load_file(self.kwargs["model_path"])
encoder.load_state_dict(state_dict, strict=True)
device = torch.device(f"cuda:{device_id}") if isinstance(device_id, int) else "cpu"
encoder.eval().to(device)
return encoder
def get_conditioning(
tokenizer: T5Tokenizer,
encoder: Encoder,
device: torch.device,
batch_inputs: bool,
*,
prompt: str,
negative_prompt: str,
):
if batch_inputs:
return dict(
batched=get_conditioning_for_prompts(
tokenizer, encoder, device, [prompt, negative_prompt]
)
)
else:
cond_input = get_conditioning_for_prompts(tokenizer, encoder, device, [prompt])
null_input = get_conditioning_for_prompts(tokenizer, encoder, device, [negative_prompt])
return dict(cond=cond_input, null=null_input)
def get_conditioning_for_prompts(tokenizer, encoder, device, prompts: List[str]):
assert len(prompts) in [1, 2] # [neg] or [pos] or [pos, neg]
B = len(prompts)
t5_toks = tokenizer(
prompts,
padding="max_length",
truncation=True,
max_length=MAX_T5_TOKEN_LENGTH,
return_tensors="pt",
return_attention_mask=True,
)
caption_input_ids_t5 = t5_toks["input_ids"]
caption_attention_mask_t5 = t5_toks["attention_mask"].bool()
del t5_toks
assert caption_input_ids_t5.shape == (B, MAX_T5_TOKEN_LENGTH)
assert caption_attention_mask_t5.shape == (B, MAX_T5_TOKEN_LENGTH)
# Special-case empty negative prompt by zero-ing it
if prompts[-1] == "":
caption_input_ids_t5[-1] = 0
caption_attention_mask_t5[-1] = False
caption_input_ids_t5 = caption_input_ids_t5.to(device, non_blocking=True)
caption_attention_mask_t5 = caption_attention_mask_t5.to(device, non_blocking=True)
y_mask = [caption_attention_mask_t5]
y_feat = [encoder(caption_input_ids_t5, caption_attention_mask_t5).last_hidden_state.detach()]
# Sometimes returns a tensor, othertimes a tuple, not sure why
# See: https://huggingface.co/genmo/mochi-1-preview/discussions/3
assert tuple(y_feat[-1].shape) == (B, MAX_T5_TOKEN_LENGTH, 4096)
assert y_feat[-1].dtype == torch.float32
return dict(y_mask=y_mask, y_feat=y_feat)
def compute_packed_indices(
device: torch.device, text_mask: torch.Tensor, num_latents: int
) -> Dict[str, Union[torch.Tensor, int]]:
"""
Based on https://github.com/Dao-AILab/flash-attention/blob/765741c1eeb86c96ee71a3291ad6968cfbf4e4a1/flash_attn/bert_padding.py#L60-L80
Args:
num_latents: Number of latent tokens
text_mask: (B, L) List of boolean tensor indicating which text tokens are not padding.
Returns:
packed_indices: Dict with keys for Flash Attention:
- valid_token_indices_kv: up to (B * (N + L),) tensor of valid token indices (non-padding)
in the packed sequence.
- cu_seqlens_kv: (B + 1,) tensor of cumulative sequence lengths in the packed sequence.
- max_seqlen_in_batch_kv: int of the maximum sequence length in the batch.
"""
# Create an expanded token mask saying which tokens are valid across both visual and text tokens.
PATCH_SIZE = 2
num_visual_tokens = num_latents // (PATCH_SIZE**2)
assert num_visual_tokens > 0
mask = F.pad(text_mask, (num_visual_tokens, 0), value=True) # (B, N + L)
seqlens_in_batch = mask.sum(dim=-1, dtype=torch.int32) # (B,)
valid_token_indices = torch.nonzero(mask.flatten(), as_tuple=False).flatten() # up to (B * (N + L),)
assert valid_token_indices.size(0) >= text_mask.size(0) * num_visual_tokens # At least (B * N,)
cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0))
max_seqlen_in_batch = seqlens_in_batch.max().item()
return {
"cu_seqlens_kv": cu_seqlens.to(device, non_blocking=True),
"max_seqlen_in_batch_kv": cast(int, max_seqlen_in_batch),
"valid_token_indices_kv": valid_token_indices.to(device, non_blocking=True),
}
def assert_eq(x, y, msg=None):
assert x == y, f"{msg or 'Assertion failed'}: {x} != {y}"
def sample_model(device, dit, conditioning, **args):
random.seed(args["seed"])
np.random.seed(args["seed"])
torch.manual_seed(args["seed"])
generator = torch.Generator(device=device)
generator.manual_seed(args["seed"])
w, h, t = args["width"], args["height"], args["num_frames"]
sample_steps = args["num_inference_steps"]
cfg_schedule = args["cfg_schedule"]
sigma_schedule = args["sigma_schedule"]
assert_eq(len(cfg_schedule), sample_steps, "cfg_schedule must have length sample_steps")
assert_eq((t - 1) % 6, 0, "t - 1 must be divisible by 6")
assert_eq(
len(sigma_schedule),
sample_steps + 1,
"sigma_schedule must have length sample_steps + 1",
)
B = 1
SPATIAL_DOWNSAMPLE = 8
TEMPORAL_DOWNSAMPLE = 6
IN_CHANNELS = 12
latent_t = ((t - 1) // TEMPORAL_DOWNSAMPLE) + 1
latent_w, latent_h = w // SPATIAL_DOWNSAMPLE, h // SPATIAL_DOWNSAMPLE
z = torch.randn(
(B, IN_CHANNELS, latent_t, latent_h, latent_w),
device=device,
dtype=torch.float32,
)
num_latents = latent_t * latent_h * latent_w
cond_batched = cond_text = cond_null = None
if "cond" in conditioning:
cond_text = conditioning["cond"]
cond_null = conditioning["null"]
cond_text["packed_indices"] = compute_packed_indices(device, cond_text["y_mask"][0], num_latents)
cond_null["packed_indices"] = compute_packed_indices(device, cond_null["y_mask"][0], num_latents)
else:
cond_batched = conditioning["batched"]
cond_batched["packed_indices"] = compute_packed_indices(device, cond_batched["y_mask"][0], num_latents)
z = repeat(z, "b ... -> (repeat b) ...", repeat=2)
def model_fn(*, z, sigma, cfg_scale):
if cond_batched:
with torch.autocast("cuda", dtype=torch.bfloat16):
out = dit(z, sigma, **cond_batched)
out_cond, out_uncond = torch.chunk(out, chunks=2, dim=0)
else:
nonlocal cond_text, cond_null
with torch.autocast("cuda", dtype=torch.bfloat16):
out_cond = dit(z, sigma, **cond_text)
out_uncond = dit(z, sigma, **cond_null)
assert out_cond.shape == out_uncond.shape
out_uncond = out_uncond.to(z)
out_cond = out_cond.to(z)
return out_uncond + cfg_scale * (out_cond - out_uncond)
# Euler sampler w/ customizable sigma schedule & cfg scale
for i in get_new_progress_bar(range(0, sample_steps), desc="Sampling"):
sigma = sigma_schedule[i]
dsigma = sigma - sigma_schedule[i + 1]
# `pred` estimates `z_0 - eps`.
pred = model_fn(
z=z,
sigma=torch.full([B] if cond_text else [B * 2], sigma, device=z.device),
cfg_scale=cfg_schedule[i],
)
assert pred.dtype == torch.float32
z = z + dsigma * pred
z = z[:B] if cond_batched else z
return dit_latents_to_vae_latents(z)
@contextmanager
def move_to_device(model: nn.Module, target_device, *, enabled=True):
if not enabled:
yield
return
og_device = next(model.parameters()).device
if og_device == target_device:
print(f"move_to_device is a no-op model is already on {target_device}")
else:
print(f"moving model from {og_device} -> {target_device}")
model.to(target_device)
yield
if og_device != target_device:
print(f"moving model from {target_device} -> {og_device}")
model.to(og_device)
def t5_tokenizer(model_dir=None):
return T5Tokenizer.from_pretrained(model_dir or T5_MODEL, legacy=False)
class MochiSingleGPUPipeline:
def __init__(
self,
*,
text_encoder_factory: ModelFactory,
dit_factory: ModelFactory,
decoder_factory: ModelFactory,
cpu_offload: Optional[bool] = False,
decode_type: str = "full",
decode_args: Optional[Dict[str, Any]] = None,
fast_init=True,
strict_load=True
):
self.device = torch.device("cuda:0")
self.tokenizer = t5_tokenizer(text_encoder_factory.model_dir)
t = Timer()
self.cpu_offload = cpu_offload
self.decode_args = decode_args or {}
self.decode_type = decode_type
init_id = "cpu" if cpu_offload else 0
with t("load_text_encoder"):
self.text_encoder = text_encoder_factory.get_model(
local_rank=0,
device_id=init_id,
world_size=1,
)
with t("load_dit"):
self.dit = dit_factory.get_model(local_rank=0, device_id=init_id, world_size=1, fast_init=fast_init, strict_load=strict_load) # type: ignore
with t("load_vae"):
self.decoder = decoder_factory.get_model(local_rank=0, device_id=init_id, world_size=1)
t.print_stats()
def __call__(self, batch_cfg, prompt, negative_prompt, **kwargs):
with torch.inference_mode():
print_max_memory = lambda: print(
f"Max memory reserved: {torch.cuda.max_memory_reserved() / 1024**3:.2f} GB"
)
print_max_memory()
with move_to_device(self.text_encoder, self.device):
conditioning = get_conditioning(
tokenizer=self.tokenizer,
encoder=self.text_encoder,
device=self.device,
batch_inputs=batch_cfg,
prompt=prompt,
negative_prompt=negative_prompt,
)
print_max_memory()
with move_to_device(self.dit, self.device):
latents = sample_model(self.device, self.dit, conditioning, **kwargs)
print_max_memory()
with move_to_device(self.decoder, self.device):
if self.decode_type == "tiled_full":
frames = decode_latents_tiled_full(
self.decoder, latents, **self.decode_args)
elif self.decode_type == "tiled_spatial":
frames = decode_latents_tiled_spatial(
self.decoder, latents, **self.decode_args,
num_tiles_w=4, num_tiles_h=2)
else:
frames = decode_latents(self.decoder, latents)
print_max_memory()
return frames.cpu().numpy()
def cast_dit(model, dtype):
for name, module in model.named_modules():
if isinstance(module, nn.Linear):
assert any(
n in name for n in ["mlp", "t5", "mod_", "attn.qkv_", "attn.proj_", "final_layer"]
), f"Unexpected linear layer: {name}"
module.to(dtype=dtype)
elif isinstance(module, nn.Conv2d):
assert "x_embedder.proj" in name, f"Unexpected conv2d layer: {name}"
module.to(dtype=dtype)
return model
### ALL CODE BELOW HERE IS FOR MULTI-GPU MODE ###
# In multi-gpu mode, all models must belong to a device which has a predefined context parallel group
# So it doesn't make sense to work with models individually
class MultiGPUContext:
def __init__(
self,
*,
text_encoder_factory,
dit_factory,
decoder_factory,
device_id,
local_rank,
world_size,
):
t = Timer()
self.device = torch.device(f"cuda:{device_id}")
print(f"Initializing rank {local_rank+1}/{world_size}")
assert world_size > 1, f"Multi-GPU mode requires world_size > 1, got {world_size}"
os.environ["MASTER_ADDR"] = "127.0.0.1"
os.environ["MASTER_PORT"] = "29500"
with t("init_process_group"):
dist.init_process_group(
"nccl",
rank=local_rank,
world_size=world_size,
device_id=self.device, # force non-lazy init
)
pg = dist.group.WORLD
cp.set_cp_group(pg, list(range(world_size)), local_rank)
distributed_kwargs = dict(local_rank=local_rank, device_id=device_id, world_size=world_size)
self.world_size = world_size
self.tokenizer = t5_tokenizer(text_encoder_factory.model_dir)
with t("load_text_encoder"):
self.text_encoder = text_encoder_factory.get_model(**distributed_kwargs)
with t("load_dit"):
self.dit = dit_factory.get_model(**distributed_kwargs)
with t("load_vae"):
self.decoder = decoder_factory.get_model(**distributed_kwargs)
self.local_rank = local_rank
t.print_stats()
def run(self, *, fn, **kwargs):
return fn(self, **kwargs)
class MochiMultiGPUPipeline:
def __init__(
self,
*,
text_encoder_factory: ModelFactory,
dit_factory: ModelFactory,
decoder_factory: ModelFactory,
world_size: int,
):
ray.init()
RemoteClass = ray.remote(MultiGPUContext)
self.ctxs = [
RemoteClass.options(num_gpus=1).remote(
text_encoder_factory=text_encoder_factory,
dit_factory=dit_factory,
decoder_factory=decoder_factory,
world_size=world_size,
device_id=0,
local_rank=i,
)
for i in range(world_size)
]
for ctx in self.ctxs:
ray.get(ctx.__ray_ready__.remote())
def __call__(self, **kwargs):
def sample(ctx, *, batch_cfg, prompt, negative_prompt, **kwargs):
with progress_bar(type="ray_tqdm", enabled=ctx.local_rank == 0), torch.inference_mode():
conditioning = get_conditioning(
ctx.tokenizer,
ctx.text_encoder,
ctx.device,
batch_cfg,
prompt=prompt,
negative_prompt=negative_prompt,
)
latents = sample_model(ctx.device, ctx.dit, conditioning=conditioning, **kwargs)
if ctx.local_rank == 0:
torch.save(latents, "latents.pt")
frames = decode_latents(ctx.decoder, latents)
return frames.cpu().numpy()
return ray.get([ctx.run.remote(fn=sample, **kwargs, show_progress=i == 0) for i, ctx in enumerate(self.ctxs)])[
0
]
@@ -0,0 +1,155 @@
from typing import Tuple, Union
import torch
import torch.distributed as dist
import torch.nn.functional as F
import genmo.mochi_preview.dit.joint_model.context_parallel as cp
def cast_tuple(t, length=1):
return t if isinstance(t, tuple) else ((t,) * length)
def cp_pass_frames(x: torch.Tensor, frames_to_send: int) -> torch.Tensor:
"""
Forward pass that handles communication between ranks for inference.
Args:
x: Tensor of shape (B, C, T, H, W)
frames_to_send: int, number of frames to communicate between ranks
Returns:
output: Tensor of shape (B, C, T', H, W)
"""
cp_rank, cp_world_size = cp.get_cp_rank_size()
if frames_to_send == 0 or cp_world_size == 1:
return x
group = cp.get_cp_group()
global_rank = dist.get_rank()
# Send to next rank
if cp_rank < cp_world_size - 1:
assert x.size(2) >= frames_to_send
tail = x[:, :, -frames_to_send:].contiguous()
dist.send(tail, global_rank + 1, group=group)
# Receive from previous rank
if cp_rank > 0:
B, C, _, H, W = x.shape
recv_buffer = torch.empty(
(B, C, frames_to_send, H, W),
dtype=x.dtype,
device=x.device,
)
dist.recv(recv_buffer, global_rank - 1, group=group)
x = torch.cat([recv_buffer, x], dim=2)
return x
def _pad_to_max(x: torch.Tensor, max_T: int) -> torch.Tensor:
if max_T > x.size(2):
pad_T = max_T - x.size(2)
pad_dims = (0, 0, 0, 0, 0, pad_T)
return F.pad(x, pad_dims)
return x
def gather_all_frames(x: torch.Tensor) -> torch.Tensor:
"""
Gathers all frames from all processes for inference.
Args:
x: Tensor of shape (B, C, T, H, W)
Returns:
output: Tensor of shape (B, C, T_total, H, W)
"""
cp_rank, cp_size = cp.get_cp_rank_size()
if cp_size == 1:
return x
cp_group = cp.get_cp_group()
# Ensure the tensor is contiguous for collective operations
x = x.contiguous()
# Get the local time dimension size
local_T = x.size(2)
local_T_tensor = torch.tensor([local_T], device=x.device, dtype=torch.int64)
# Gather all T sizes from all processes
all_T = [torch.zeros(1, dtype=torch.int64, device=x.device) for _ in range(cp_size)]
dist.all_gather(all_T, local_T_tensor, group=cp_group)
all_T = [t.item() for t in all_T]
# Pad the tensor at the end of the time dimension to match max_T
max_T = max(all_T)
x = _pad_to_max(x, max_T).contiguous()
# Prepare a list to hold the gathered tensors
gathered_x = [torch.zeros_like(x).contiguous() for _ in range(cp_size)]
# Perform the all_gather operation
dist.all_gather(gathered_x, x, group=cp_group)
# Slice each gathered tensor back to its original T size
for idx, t_size in enumerate(all_T):
gathered_x[idx] = gathered_x[idx][:, :, :t_size]
return torch.cat(gathered_x, dim=2)
def excessive_memory_usage(input: torch.Tensor, max_gb: float = 2.0) -> bool:
"""Estimate memory usage based on input tensor size and data type."""
element_size = input.element_size() # Size in bytes of each element
memory_bytes = input.numel() * element_size
memory_gb = memory_bytes / 1024**3
return memory_gb > max_gb
class ContextParallelCausalConv3d(torch.nn.Conv3d):
def __init__(
self,
in_channels,
out_channels,
kernel_size: Union[int, Tuple[int, int, int]],
stride: Union[int, Tuple[int, int, int]],
**kwargs,
):
kernel_size = cast_tuple(kernel_size, 3)
stride = cast_tuple(stride, 3)
height_pad = (kernel_size[1] - 1) // 2
width_pad = (kernel_size[2] - 1) // 2
super().__init__(
in_channels=in_channels,
out_channels=out_channels,
kernel_size=kernel_size,
stride=stride,
dilation=(1, 1, 1),
padding=(0, height_pad, width_pad),
**kwargs,
)
def forward(self, x: torch.Tensor):
cp_rank, cp_world_size = cp.get_cp_rank_size()
context_size = self.kernel_size[0] - 1
if cp_rank == 0:
mode = "constant" if self.padding_mode == "zeros" else self.padding_mode
x = F.pad(x, (0, 0, 0, 0, context_size, 0), mode=mode)
if cp_world_size == 1:
return super().forward(x)
if all(s == 1 for s in self.stride):
# Receive some frames from previous rank.
x = cp_pass_frames(x, context_size)
return super().forward(x)
# Less efficient implementation for strided convs.
# All gather x, infer and chunk.
x = gather_all_frames(x) # [B, C, k - 1 + global_T, H, W]
x = super().forward(x)
x_chunks = x.tensor_split(cp_world_size, dim=2)
assert len(x_chunks) == cp_world_size
return x_chunks[cp_rank]
@@ -0,0 +1,35 @@
"""Container for latent space posterior."""
import torch
class LatentDistribution:
def __init__(self, mean: torch.Tensor, logvar: torch.Tensor):
"""Initialize latent distribution.
Args:
mean: Mean of the distribution. Shape: [B, C, T, H, W].
logvar: Logarithm of variance of the distribution. Shape: [B, C, T, H, W].
"""
assert mean.shape == logvar.shape
self.mean = mean
self.logvar = logvar
def sample(self, temperature=1.0, generator: torch.Generator = None, noise=None):
if temperature == 0.0:
return self.mean
if noise is None:
noise = torch.randn(self.mean.shape, device=self.mean.device, dtype=self.mean.dtype, generator=generator)
else:
assert noise.device == self.mean.device
noise = noise.to(self.mean.dtype)
if temperature != 1.0:
raise NotImplementedError(f"Temperature {temperature} is not supported.")
# Just Gaussian sample with no scaling of variance.
return noise * torch.exp(self.logvar * 0.5) + self.mean
def mode(self):
return self.mean
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,67 @@
import torch
# Channel-wise mean and standard deviation of VAE encoder latents
STATS = {
"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,
]
),
"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,
]
),
}
def dit_latents_to_vae_latents(dit_outputs: torch.Tensor) -> torch.Tensor:
"""Unnormalize latents output by Mochi's DiT to be compatible with VAE.
Run this on sampled latents before calling the VAE decoder.
Args:
latents (torch.Tensor): [B, C_z, T_z, H_z, W_z], float
Returns:
torch.Tensor: [B, C_z, T_z, H_z, W_z], float
"""
mean = STATS["mean"][:, None, None, None]
std = STATS["std"][:, None, None, None]
assert dit_outputs.ndim == 5
assert dit_outputs.size(1) == mean.size(0) == std.size(0)
return dit_outputs * std.to(dit_outputs) + mean.to(dit_outputs)
def vae_latents_to_dit_latents(vae_latents: torch.Tensor):
"""Normalize latents output by the VAE encoder to be compatible with Mochi's DiT.
E.g, for fine-tuning or video-to-video.
"""
mean = STATS["mean"][:, None, None, None]
std = STATS["std"][:, None, None, None]
assert vae_latents.ndim == 5
assert vae_latents.size(1) == mean.size(0) == std.size(0)
return (vae_latents - mean.to(vae_latents)) / std.to(vae_latents)
@@ -0,0 +1,431 @@
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)
@@ -0,0 +1,42 @@
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,13 +19,29 @@ 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 fastvideo.model.norm import MochiLayerNormContinuous, MochiRMSNormZero, MochiModulatedRMSNorm, MochiRMSNorm
from diffusers.loaders import PeftAdapterMixin
from fastvideo.models.mochi_hf.norm import (
MochiLayerNormContinuous,
MochiRMSNormZero,
MochiModulatedRMSNorm,
MochiRMSNorm,
)
from diffusers.models.normalization import AdaLayerNormContinuous
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
@@ -39,6 +55,9 @@ 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,
@@ -51,37 +70,50 @@ 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)
)
def flash_attn_no_pad(qkv, key_padding_mask, causal=False, dropout_p=0.0, softmax_scale=None):
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 = 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)
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
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,
)
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,
@@ -115,17 +147,25 @@ 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
@@ -143,15 +183,15 @@ class MochiAttention(nn.Module):
attention_mask=attention_mask,
**kwargs,
)
class MochiAttnProcessor2_0:
"""Attention processor used in Mochi."""
def __init__(self):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError("MochiAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.")
raise ImportError(
"MochiAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0."
)
def __call__(
self,
@@ -172,12 +212,11 @@ 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)
@@ -186,37 +225,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()
@@ -224,9 +263,10 @@ 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),
@@ -237,6 +277,7 @@ class MochiAttnProcessor2_0:
sequence_length = query.size(1)
encoder_sequence_length = encoder_query.size(1)
# H
query = torch.cat([query, encoder_query], dim=1).unsqueeze(2)
key = torch.cat([key, encoder_key], dim=1).unsqueeze(2)
@@ -246,14 +287,15 @@ 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 = 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(
@@ -261,7 +303,9 @@ 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)
@@ -273,8 +317,6 @@ 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)
@@ -286,6 +328,7 @@ class MochiAttnProcessor2_0:
return hidden_states, encoder_hidden_states
@maybe_allow_in_graph
class MochiTransformerBlock(nn.Module):
r"""
@@ -328,7 +371,9 @@ 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,
@@ -352,12 +397,18 @@ 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(
@@ -377,14 +428,19 @@ 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)
@@ -392,20 +448,27 @@ 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(
@@ -447,18 +510,22 @@ 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
@@ -478,7 +545,7 @@ class MochiRoPE(nn.Module):
@maybe_allow_in_graph
class MochiTransformer3DModel(ModelMixin, ConfigMixin):
class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
r"""
A Transformer model for video-like data introduced in [Mochi](https://huggingface.co/genmo/mochi-1-preview).
@@ -545,7 +612,9 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin):
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(
@@ -564,7 +633,11 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin):
)
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)
@@ -576,32 +649,57 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin):
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
timestep: torch.LongTensor,
encoder_attention_mask: torch.Tensor,
hidden_states: torch.Tensor, # [2, 12, 28, 60, 106]
encoder_hidden_states: torch.Tensor, # [2, 256, 4096]
timestep: torch.LongTensor, # [2]
encoder_attention_mask: torch.Tensor, #[2, 256]
output_attn = False,
attention_kwargs: Optional[Dict[str, Any]] = None,
return_dict: bool = False,
) -> torch.Tensor:
assert return_dict is False, "return_dict is not supported in MochiTransformer3DModel"
batch_size, num_channels, num_frames, height, width = hidden_states.shape
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."
)
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
temb, encoder_hidden_states = self.time_embed( # [2, 3072], [2, 256, 1536]
timestep,
encoder_hidden_states,
encoder_attention_mask,
hidden_dtype=hidden_states.dtype,
)
hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1)
hidden_states = self.patch_embed(hidden_states)
hidden_states = hidden_states.unflatten(0, (batch_size, -1)).flatten(1, 2)
hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1) # [56, 12, 60, 106]
hidden_states = self.patch_embed(hidden_states) # [56, 1590, 3072]
hidden_states = hidden_states.unflatten(0, (batch_size, -1)).flatten(1, 2) # [2, 44520, 3072]
image_rotary_emb = self.rope(
self.pos_frequencies,
num_frames,
image_rotary_emb = self.rope( #[0][44520, 24, 64]
self.pos_frequencies, #[3, 24, 64]
num_frames, # 28
post_patch_height,
post_patch_width,
device=hidden_states.device,
@@ -617,8 +715,14 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin):
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,
@@ -629,26 +733,30 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin):
**ckpt_kwargs,
)
else:
hidden_states, encoder_hidden_states, attn_outputs = block(
hidden_states, encoder_hidden_states, attn_outputs = block( # [2, 44520, 3072], [2, 256, 1536],
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
encoder_attention_mask=encoder_attention_mask,
temb=temb,
image_rotary_emb=image_rotary_emb,
output_attn = output_attn,
output_attn=output_attn,
)
attn_outputs_list.append(attn_outputs)
hidden_states = self.norm_out(hidden_states, temb)
hidden_states = self.proj_out(hidden_states)
hidden_states = self.proj_out(hidden_states) #[2, 44520, 48]
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)
hidden_states = hidden_states.reshape(batch_size, num_frames, post_patch_height, post_patch_width, p, p, -1) # [2, 28, 30, 53, 2, 2, 12]
hidden_states = hidden_states.permute(0, 6, 1, 2, 4, 3, 5) # [2, 12, 28, 30, 2, 53, 2]
output = hidden_states.reshape(batch_size, -1, num_frames, height, width) # [2, 12, 28, 60, 106]
if not output_attn :
attn_outputs_list = None
if USE_PEFT_BACKEND:
# remove `lora_scale` from each PEFT layer
unscale_lora_layers(self, lora_scale)
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,7 +63,7 @@ class MochiRMSNorm(nn.Module):
hidden_states = hidden_states.to(hidden_states_dtype)
return hidden_states
class MochiLayerNormContinuous(nn.Module):
def __init__(
@@ -92,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"""
@@ -102,7 +102,11 @@ 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__()
@@ -118,7 +122,9 @@ 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
from typing import Callable, Dict, List, Optional, Union, Any
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.model.modeling_mochi import MochiTransformer3DModel
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import (
@@ -35,7 +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 fastvideo.utils.communications import all_gather
from diffusers.loaders import Mochi1LoraLoaderMixin
if is_torch_xla_available():
import torch_xla.core.xla_model as xm
@@ -80,14 +81,19 @@ 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)
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]
@@ -127,9 +133,13 @@ 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"
@@ -139,7 +149,9 @@ 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"
@@ -154,7 +166,7 @@ def retrieve_timesteps(
return timesteps, num_inference_steps
class MochiPipeline(DiffusionPipeline):
class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
r"""
The mochi pipeline for text-to-video generation.
@@ -199,14 +211,17 @@ class MochiPipeline(DiffusionPipeline):
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
@@ -238,22 +253,32 @@ class MochiPipeline(DiffusionPipeline):
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)
@@ -320,7 +345,11 @@ class MochiPipeline(DiffusionPipeline):
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(
@@ -334,7 +363,10 @@ class MochiPipeline(DiffusionPipeline):
" 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,
@@ -342,7 +374,12 @@ class MochiPipeline(DiffusionPipeline):
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,
@@ -356,10 +393,13 @@ class MochiPipeline(DiffusionPipeline):
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]}"
@@ -374,14 +414,25 @@ class MochiPipeline(DiffusionPipeline):
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:
@@ -467,6 +518,10 @@ class MochiPipeline(DiffusionPipeline):
def num_timesteps(self):
return self._num_timesteps
@property
def attention_kwargs(self):
return self._attention_kwargs
@property
def interrupt(self):
return self._interrupt
@@ -492,10 +547,11 @@ class MochiPipeline(DiffusionPipeline):
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.
@@ -547,6 +603,10 @@ class MochiPipeline(DiffusionPipeline):
[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,
@@ -586,6 +646,7 @@ class MochiPipeline(DiffusionPipeline):
)
self._guidance_scale = guidance_scale
self._attention_kwargs = attention_kwargs
self._interrupt = False
# 2. Define call parameters
@@ -618,7 +679,9 @@ class MochiPipeline(DiffusionPipeline):
)
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
@@ -635,9 +698,10 @@ class MochiPipeline(DiffusionPipeline):
)
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
@@ -646,7 +710,7 @@ class MochiPipeline(DiffusionPipeline):
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,
@@ -660,7 +724,9 @@ class MochiPipeline(DiffusionPipeline):
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
@@ -669,26 +735,36 @@ class MochiPipeline(DiffusionPipeline):
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:
@@ -706,7 +782,9 @@ class MochiPipeline(DiffusionPipeline):
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:
@@ -714,34 +792,49 @@ class MochiPipeline(DiffusionPipeline):
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)
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
)
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:
+55 -29
View File
@@ -2,12 +2,15 @@ import json
import torch.distributed as dist
import torch
from fastvideo.model.pipeline_mochi import MochiPipeline
from fastvideo.models.mochi_hf.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
@@ -19,17 +22,16 @@ def generate_video_and_latent(pipe, prompt, height, width, num_frames, num_infer
generator=generator,
num_inference_steps=num_inference_steps,
guidance_scale=guidance_scale,
return_all_states=True,
output_type="latent_and_video",
)
# 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)
@@ -37,47 +39,71 @@ 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)
torch.cuda.set_device(local_rank)
dist.init_process_group(backend='nccl', init_method='env://', world_size=world_size, 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
)
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)
@@ -85,7 +111,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"
@@ -97,11 +123,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)
+96 -114
View File
@@ -1,12 +1,15 @@
import torch
from fastvideo.model.pipeline_mochi import MochiPipeline
from fastvideo.models.mochi_hf.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.model.modeling_mochi import MochiTransformer3DModel
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
import json
from typing import Optional
from safetensors.torch import save_file, load_file
@@ -17,88 +20,21 @@ 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
from fastvideo.models.mochi_genmo.mochi_preview.dit.joint_model.asymm_models_joint import AsymmDiTJoint
from safetensors.torch import load_file
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()
@@ -110,54 +46,92 @@ def main(args):
scheduler = FlowMatchEulerDiscreteScheduler()
else:
linear_quadratic = True if "linear_quadratic" in args.scheduler_type else False
scheduler = PCMFMScheduler(1000, args.shift, args.num_euler_timesteps, linear_quadratic,args.linear_threshold, args.linear_range)
if args.transformer_path is not None:
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
else:
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder = 'transformer/')
if args.lora_checkpoint_dir is not None:
# Load and merge LoRA weights
transformer = load_lora_checkpoint(
transformer=transformer,
optimizer=None, # No optimizer needed for inference
output_dir=args.lora_checkpoint_dir
scheduler = PCMFMScheduler(
1000,
args.shift,
args.num_euler_timesteps,
linear_quadratic,
args.linear_threshold,
args.linear_range,
)
print(f"Loaded and merged LoRA weights from {args.lora_checkpoint_dir}")
pipe = MochiPipeline.from_pretrained(args.model_path, transformer = transformer,scheduler=scheduler)
mochi_genmo = False
if mochi_genmo:
model_path = "/root/weights/dit.safetensors"
state_dcit = load_file(model_path)
transformer = AsymmDiTJoint()
transformer.load_state_dict(state_dcit)
# from IPython import embed
# embed()
transformer.config.in_channels = 12
print("load gennmo mochi successfully")
else:
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()
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}")
# 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:
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])
for prompt in prompts:
video = pipe(
prompt=[prompt],
height=args.height,
width=args.width,
num_frames=args.num_frames,
num_inference_steps=args.num_inference_steps,
guidance_scale=args.guidance_scale,
generator=generator,
).frames
videos.append(video[0])
else:
with torch.autocast("cuda", dtype=torch.bfloat16):
videos = pipe(
@@ -173,18 +147,21 @@ def main(args):
if nccl_info.global_rank <= 0:
if prompts is not None:
# mkdir
# 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)
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)
@@ -198,7 +175,12 @@ 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)
+15 -8
View File
@@ -1,9 +1,11 @@
import torch
from fastvideo.model.pipeline_mochi import MochiPipeline
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
from fastvideo.models.mochi_hf.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)
@@ -12,8 +14,12 @@ 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()
@@ -29,14 +35,15 @@ 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)
+445 -187
View File
@@ -5,10 +5,14 @@ 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.model.mochi_latents_utils import normalize_mochi_dit_input
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_mochi_dit_input
from fastvideo.utils.validation import log_validation
import time
from torch.utils.data import DataLoader
@@ -25,32 +29,40 @@ import wandb
from accelerate.utils import set_seed
from tqdm.auto import tqdm
from fastvideo.fsdp_util import get_dit_fsdp_kwargs, apply_fsdp_checkpointing
import diffusers
from diffusers.utils import convert_unet_state_dict_to_peft
from diffusers import (
FlowMatchEulerDiscreteScheduler,
)
from diffusers.optimization import get_scheduler
from fastvideo.model.modeling_mochi import MochiTransformer3DModel
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from diffusers.utils import check_min_version
from fastvideo.dataset.latent_datasets import LatentDataset, latent_collate_function
import torch.distributed as dist
from safetensors.torch import save_file, load_file
from peft import LoraConfig, inject_adapter_in_model
from 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_training
from fastvideo.utils.checkpoint import (
save_checkpoint,
save_lora_checkpoint,
resume_lora_optimizer,
)
from fastvideo.utils.logging import main_print
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
# 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.
@@ -61,7 +73,13 @@ 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)
@@ -70,6 +88,7 @@ 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)
@@ -82,16 +101,35 @@ def get_sigmas(noise_scheduler, device, timesteps, n_dim=4, dtype=torch.float32)
return sigma
def train_one_step_mochi(transformer, optimizer, lr_scheduler,loader, noise_scheduler, noise_random_generator, gradient_accumulation_steps, sp_size, precondition_outputs, max_grad_norm, weighting_scheme, logit_mean, logit_std, mode_scale):
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,
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,
@@ -104,56 +142,58 @@ def train_one_step_mochi(transformer, optimizer, lr_scheduler,loader, noise_sche
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
target = noise - latents
loss = (
torch.mean((model_pred.float() - target.float()) ** 2)
/ gradient_accumulation_steps
)
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()
@@ -167,45 +207,77 @@ 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:
f
main_print(f"--> loading model from {args.pretrained_model_name_or_path}")
# keep the master weight to float32
weight_type = torch.float32 if args.master_weight_type == 'fp32' else torch.bfloat16
transformer = MochiTransformer3DModel.from_pretrained(
args.pretrained_model_name_or_path,
subfolder="transformer",
torch_dtype = torch.float32,
#torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
torch_dtype=torch.float32
if args.master_weight_type == "fp32"
else torch.bfloat16,
)
if args.use_lora:
lora_config = LoraConfig(
transformer.requires_grad_(False)
transformer_lora_config = LoraConfig(
r=args.lora_rank,
lora_alpha=args.lora_alpha,
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
init_lora_weights=True,
target_modules=["to_k", "to_q", "to_v", "to_out.0"],
)
transformer = get_lora_model(transformer, lora_config)
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 = get_dit_fsdp_kwargs(
args.fsdp_sharding_startegy,
args.use_lora,
args.use_cpu_offload,
args.master_weight_type,
)
main_print(f" Total training parameters = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e6} M")
main_print(f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}")
fsdp_kwargs = get_dit_fsdp_kwargs(args.fsdp_sharding_startegy, args.use_lora, args.use_cpu_offload)
if args.use_lora:
transformer.config.lora_rank = args.lora_rank
transformer.config.lora_alpha = args.lora_alpha
transformer.config.lora_target_modules = ["to_k", "to_q", "to_v", "to_out.0"]
transformer._no_split_modules = ["MochiTransformerBlock"]
fsdp_kwargs['auto_wrap_policy'] = fsdp_kwargs['auto_wrap_policy'](transformer)
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
transformer = FSDP(
transformer,
**fsdp_kwargs,
@@ -226,39 +298,44 @@ 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_training(
transformer, optimizer, init_steps = resume_lora_optimizer(
transformer, args.resume_from_lora_checkpoint, optimizer
)
)
main_print(f"optimizer: {optimizer}")
#todo add lr scheduler
lr_scheduler = get_scheduler(
args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=args.lr_warmup_steps * world_size,
num_training_steps=args.max_train_steps * world_size,
num_cycles=args.lr_num_cycles,
power=args.lr_power,
last_epoch=init_steps - 1,
)
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,
)
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,
@@ -266,93 +343,136 @@ 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,
disable=local_rank > 0,
)
loader = sp_parallel_dataloader_wrapper(
train_dataloader,
device,
args.train_batch_size,
args.sp_size,
args.train_sp_batch_size,
)
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_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)
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()
@@ -363,91 +483,195 @@ if __name__ == "__main__":
# 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("--cache_dir", type=str, default='./cache_dir')
parser.add_argument('--enable_stable_fp32', action='store_true') # TODO
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
# diffusion setting
parser.add_argument("--ema_decay", type=float, default=0.999)
parser.add_argument("--ema_start_step", type=int, default=0)
parser.add_argument('--cfg', type=float, default=0.1)
parser.add_argument("--precondition_outputs", action="store_true", help="Whether to precondition the outputs of the model.")
parser.add_argument("--cfg", type=float, default=0.1)
parser.add_argument(
"--precondition_outputs",
action="store_true",
help="Whether to precondition the outputs of the model.",
)
# validation & logs
parser.add_argument("--validation_prompt_dir", type=str)
parser.add_argument("--uncond_prompt_dir", type=str)
parser.add_argument("--validation_sampling_steps", type=int, default=64)
parser.add_argument('--validation_guidance_scale', type=float, default=4.5)
parser.add_argument('--validation_steps', type=float, default=4.5)
parser.add_argument(
"--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("--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(
@@ -457,10 +681,16 @@ 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",
@@ -469,14 +699,42 @@ 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(
"--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(
"--Mochi_type",
type=str,
default="hf",
help="Choose Mochi model between hf and genmo(original mochi).",
)
args = parser.parse_args()
main(args)
main(args)
+135 -97
View File
@@ -1,29 +1,42 @@
# 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 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
def save_checkpoint(model, optimizer, rank, output_dir, step, discriminator=False):
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,
model,
optimizer,
)
#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)
# 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)
@@ -39,21 +52,30 @@ 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")
@@ -62,8 +84,7 @@ def save_checkpoint_generator_discriminator(model, optimizer, discriminator, dis
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)
@@ -74,44 +95,53 @@ def save_checkpoint_generator_discriminator(model, optimizer, discriminator, dis
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(),
)
@@ -120,38 +150,62 @@ 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)
@@ -162,83 +216,67 @@ 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
):
main_print(f"--> saving LoRA checkpoint at step {step}")
def save_lora_checkpoint(transformer, optimizer, rank, output_dir, 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_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,
transformer,
optimizer,
)
if rank <= 0:
save_dir = os.path.join(output_dir, f"lora-checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
weight_path = os.path.join(save_dir, "lora_weights.safetensors")
save_file(lora_state_dict, weight_path)
# save optimizer
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_training(
transformer,
checkpoint_dir,
optimizer
):
weight_path = os.path.join(checkpoint_dir, "lora_weights.safetensors")
lora_weights = load_file(weight_path)
def resume_lora_optimizer(transformer, checkpoint_dir, optimizer):
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 training from step {step}")
return transformer, optimizer, step
step = config_dict["step"]
main_print(f"--> Successfully resuming LoRA optimizer from step {step}")
return transformer, optimizer, step
+81 -45
View File
@@ -10,11 +10,12 @@ 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:
@@ -112,7 +113,6 @@ 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,18 +129,16 @@ class SeqAllToAll4D(torch.autograd.Function):
None,
None,
)
def all_to_all_4D(
input_: torch.Tensor,
scatter_dim: int = 2,
gather_dim: int = 1,
):
return SeqAllToAll4D.apply( nccl_info.group,input_, scatter_dim, gather_dim)
return SeqAllToAll4D.apply(nccl_info.group, input_, scatter_dim, gather_dim)
def _all_to_all(
input_: torch.Tensor,
world_size: int,
@@ -148,7 +146,9 @@ 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()
@@ -170,7 +170,9 @@ 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
@@ -198,7 +200,6 @@ def all_to_all(
return _AllToAll.apply(input_, nccl_info.group, scatter_dim, gather_dim)
class _AllGather(torch.autograd.Function):
"""All-gather communication with autograd support.
@@ -237,6 +238,7 @@ 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.
@@ -250,49 +252,83 @@ 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,
)
@@ -1,16 +1,18 @@
import argparse
import torch
from accelerate.logging import get_logger
from fastvideo.model.pipeline_mochi import MochiPipeline
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
class T5dataset(Dataset):
def __init__(
self,
@@ -21,31 +23,40 @@ class T5dataset(Dataset):
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'])
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']
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")
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)
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)
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)
@@ -53,32 +64,40 @@ def main(args):
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)
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'],
prompt=data["caption"],
)
if args.vae_debug:
latents = data['latents']
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")
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)
@@ -86,11 +105,11 @@ def main(args):
if args.vae_debug:
export_to_video(video[idx], video_path, fps=30)
item = {}
item['length'] = int(data['length'][idx])
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]
item["caption"] = data["caption"][idx]
json_data.append(item)
dist.barrier()
local_data = json_data
@@ -99,19 +118,35 @@ def main(args):
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:
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")
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)
@@ -14,21 +14,30 @@ 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)
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)
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)
sampler = DistributedSampler(
train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True
)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
@@ -36,29 +45,36 @@ def main(args):
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")
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']):
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")
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]
item["caption"] = data["text"][idx]
json_data.append(item)
print(f"{video_name} processed")
dist.barrier()
@@ -67,40 +83,61 @@ def main(args):
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:
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(
"--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("--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***."
),
)
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)
+151 -67
View File
@@ -13,11 +13,13 @@ 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."""
@@ -31,17 +33,20 @@ 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:
@@ -50,6 +55,7 @@ def pad_to_multiple(number, ds_stride):
padding = ds_stride - remainder
return number + padding
class Collate:
def __init__(self, args):
self.batch_size = args.train_batch_size
@@ -71,9 +77,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):
@@ -81,13 +87,29 @@ 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
@@ -98,13 +120,30 @@ 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]
@@ -115,50 +154,61 @@ 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.
@@ -184,13 +234,16 @@ 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)
@@ -204,48 +257,70 @@ 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):
return aligned_magabatches
def get_length_grouped_indices(
lengths,
batch_size,
world_size,
generator=None,
group_frame=False,
group_resolution=False,
seed=42,
):
# We need to use torch for the random part as a distributed sampler will set the random seed for torch.
if generator is None:
generator = torch.Generator().manual_seed(seed) # every rank will generate a fixed order but random index
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):
@@ -259,9 +334,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:
@@ -279,15 +354,24 @@ 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
@@ -1,328 +0,0 @@
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")
+6 -5
View File
@@ -1,23 +1,24 @@
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
@@ -0,0 +1,77 @@
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
+16 -5
View File
@@ -2,6 +2,7 @@ import torch
import torch.distributed as dist
import os
class COMM_INFO:
def __init__(self):
self.group = None
@@ -10,8 +11,11 @@ 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:
@@ -19,22 +23,29 @@ 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
@@ -1,471 +0,0 @@
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))
+136 -69
View File
@@ -1,13 +1,14 @@
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.model.pipeline_mochi import linear_quadratic_schedule, retrieve_timesteps
from fastvideo.models.mochi_hf.pipeline_mochi import (
linear_quadratic_schedule,
retrieve_timesteps,
)
from tqdm import tqdm
from diffusers.video_processor import VideoProcessor
from diffusers import (
@@ -20,6 +21,8 @@ from diffusers.utils import export_to_video
import os
import wandb
import gc
def prepare_latents(
batch_size,
num_channels_latents,
@@ -38,10 +41,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,
@@ -60,8 +63,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,
vae_spatial_scale_factor=8,
vae_temporal_scale_factor=6,
):
device = vae.device
@@ -70,7 +73,9 @@ 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
@@ -85,13 +90,14 @@ 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
@@ -118,12 +124,16 @@ 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
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(
@@ -133,16 +143,20 @@ def sample_validation_video(
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:
@@ -150,28 +164,35 @@ 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:
@@ -180,87 +201,134 @@ def sample_validation_video(
video = vae.decode(latents, return_dict=False)[0]
video_processor = VideoProcessor(vae_scale_factor=vae_spatial_scale_factor)
video = video_processor.postprocess_video(video, output_type=output_type)
return (video,)
@torch.no_grad()
@torch.autocast("cuda", dtype=torch.bfloat16)
def log_validation(args, transformer, device, weight_dtype, global_step, scheduler_type="euler",shift=1.0, num_euler_timesteps=100, linear_quadratic_threshold=0.025, linear_range=0.5, ema=False):
#TODO
def log_validation(
args,
transformer,
device,
weight_dtype,
global_step,
scheduler_type="euler",
shift=1.0,
num_euler_timesteps=100,
linear_quadratic_threshold=0.025,
linear_range=0.5,
ema=False,
):
# TODO
print(f"Running validation....\n")
vae = AutoencoderKLMochi.from_pretrained(args.pretrained_model_name_or_path, subfolder="vae", torch_dtype=weight_dtype).to("cuda")
vae = 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
num_embeds = len([f for f in os.listdir(args.validation_prompt_dir) if "embed" in f])
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
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(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)
prompt_embed_path = os.path.join(
args.validation_prompt_dir, f"embed{i}.pt"
)
prompt_mask_path = os.path.join(
args.validation_prompt_dir, f"mask{i}.pt"
)
prompt_embeds = (
torch.load(prompt_embed_path, map_location="cpu", weights_only=True)
.to(device)
.to(weight_dtype)
.unsqueeze(0)
)
prompt_attention_mask = (
torch.load(prompt_mask_path, map_location="cpu", weights_only=True)
.to(device)
.to(weight_dtype)
.unsqueeze(0)
)
negative_prompt_embeds = (
torch.zeros(256, 4096).to(device).to(weight_dtype).unsqueeze(0)
)
negative_prompt_attention_mask = (
torch.zeros(256).bool().to(device).unsqueeze(0)
)
generator = torch.Generator(device="cuda").manual_seed(12345)
video = sample_validation_video(
transformer,
vae,
scheduler,
scheduler_type=scheduler_type,
num_frames=args.num_frames,
# Peiyuan TODO: remove hardcode
height=480,
width=848,
num_inference_steps=validation_sampling_step,
guidance_scale=validation_guidance_scale,
generator=generator,
prompt_embeds = prompt_embeds,
prompt_attention_mask = prompt_attention_mask,
negative_prompt_embeds = negative_prompt_embeds,
negative_prompt_attention_mask = negative_prompt_attention_mask,
)[0]
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")
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)
@@ -271,4 +339,3 @@ def log_validation(args, transformer, device, weight_dtype, global_step, schedu
]
}
wandb.log(logs, step=global_step)
-19
View File
@@ -1,24 +1,5 @@
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
-11
View File
@@ -1,11 +0,0 @@
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)
+1 -5
View File
@@ -21,12 +21,8 @@ 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"
]
"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]
+34 -14
View File
@@ -1,19 +1,39 @@
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')
# 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(di
"--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)
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
@@ -1,56 +0,0 @@
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
@@ -1,49 +0,0 @@
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
@@ -1,45 +0,0 @@
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
-44
View File
@@ -1,44 +0,0 @@
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-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/shift1_euler_50"\
--tracker_project_name PCM \
--num_frames 163 \
--shift 1 \
--validation_guidance_scale "2.5,3.5,4.5" \
--num_euler_timesteps 50
-51
View File
@@ -1,51 +0,0 @@
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 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 125\
--validation_sampling_steps 8 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_lrg_0.75_phase1"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--multi_phased_distill_schedule "4000-1"
-52
View File
@@ -1,52 +0,0 @@
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 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 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_thres0.1_lrg_0.75_cfg_4.5"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--distill_cfg 4.5
-51
View File
@@ -1,51 +0,0 @@
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 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 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_thres0.15_lrg_0.75"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.15 \
--linear_range 0.75
-54
View File
@@ -1,54 +0,0 @@
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 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 125\
--validation_sampling_steps 8 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_lrg_0.75_ema_0.95_decay0.0"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--use_ema \
--ema_decay 0.95 \
--weight_decay 0.0
-57
View File
@@ -1,57 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=offline
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
CACHE_DIR=/data/.cache
EXPERIMENT=lq_euler_50_thres0.1_lrg_0.75_cfg0.0
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
export WANDB_DIR=$DATA_DIR/wandb/
torchrun --nnodes 1 --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 $CACHE_DIR \
--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 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=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=$OUTPUT_DIR \
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "1.5,2.5,4.5,6.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--not_apply_cfg_solver
-55
View File
@@ -1,55 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=offline
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
CACHE_DIR=/data/.cache
EXPERIMENT=lq_euler_50_thres0.1_lrg_0.75_cfg0.0
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
export WANDB_DIR=$DATA_DIR/wandb/
torchrun --nnodes 8 --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 $CACHE_DIR \
--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 24\
--sp_size 8\
--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=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=$OUTPUT_DIR \
--tracker_project_name PCM \
--num_frames 139 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "1.5,2.5,4.5,6.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--not_apply_cfg_solver
-52
View File
@@ -1,52 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=offline
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
CACHE_DIR=/data/.cache
EXPERIMENT=shift16_euler_50
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
export WANDB_DIR=$DATA_DIR/wandb/
torchrun --nnodes 1 --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 $CACHE_DIR \
--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 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=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=$OUTPUT_DIR \
--tracker_project_name PCM \
--num_frames 163 \
--shift 16 \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50
-52
View File
@@ -1,52 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=offline
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
CACHE_DIR=/data/.cache
EXPERIMENT=shift16_euler_50
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
export WANDB_DIR=$DATA_DIR/wandb/
torchrun --nnodes 8 --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 $CACHE_DIR \
--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 24\
--sp_size 8\
--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=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=$OUTPUT_DIR \
--tracker_project_name PCM \
--num_frames 139 \
--shift 16 \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50
-52
View File
@@ -1,52 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=offline
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
CACHE_DIR=/data/.cache
EXPERIMENT=4step_infer_shift16_euler_50
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
export WANDB_DIR=$DATA_DIR/wandb/
torchrun --nnodes 1 --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 $CACHE_DIR \
--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 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=500\
--validation_steps 125\
--validation_sampling_steps 4 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--ema_decay 0.999\
--log_validation\
--output_dir=$OUTPUT_DIR \
--tracker_project_name PCM \
--num_frames 163 \
--shift 16 \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50
-52
View File
@@ -1,52 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=offline
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
CACHE_DIR=/data/.cache
EXPERIMENT=4step_infer_shift16_euler_50
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
export WANDB_DIR=$DATA_DIR/wandb/
torchrun --nnodes 8 --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 $CACHE_DIR \
--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 24\
--sp_size 8\
--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=500\
--validation_steps 125\
--validation_sampling_steps 4 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--ema_decay 0.999\
--log_validation\
--output_dir=$OUTPUT_DIR \
--tracker_project_name PCM \
--num_frames 139 \
--shift 16 \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50
-48
View File
@@ -1,48 +0,0 @@
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_thresh0.05"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale 4.5 \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.05
gsutil cp data/outputs/lq_euler_50_thresh0.05/checkpoint-4000 gs://vid_gen/runlong_temp_folder_for_pandas70m_debugging/fastvid/lq_euler_50_thresh0.05/checkpoint-4000
-52
View File
@@ -1,52 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=offline
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
CACHE_DIR=/data/.cache
EXPERIMENT=4step_infer_shift12_euler_50
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
export WANDB_DIR=$DATA_DIR/wandb/
torchrun --nnodes 1 --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 $CACHE_DIR \
--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 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=500\
--validation_steps 125\
--validation_sampling_steps 4 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--ema_decay 0.999\
--log_validation\
--output_dir=$OUTPUT_DIR \
--tracker_project_name PCM \
--num_frames 163 \
--shift 12 \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50
-52
View File
@@ -1,52 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=offline
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
CACHE_DIR=/data/.cache
EXPERIMENT=4step_infer_shift12_euler_50
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
export WANDB_DIR=$DATA_DIR/wandb/
torchrun --nnodes 8 --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 $CACHE_DIR \
--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 24\
--sp_size 8\
--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=500\
--validation_steps 125\
--validation_sampling_steps 4 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--ema_decay 0.999\
--log_validation\
--output_dir=$OUTPUT_DIR \
--tracker_project_name PCM \
--num_frames 139 \
--shift 12 \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50
-54
View File
@@ -1,54 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=offline
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
CACHE_DIR=/data/.cache
EXPERIMENT=4step_infer_lq_euler_50_thresh0.1_lrg_0.75
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
export WANDB_DIR=$DATA_DIR/wandb/
torchrun --nnodes 1 --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 $CACHE_DIR \
--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 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=500\
--validation_steps 125\
--validation_sampling_steps 4 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--ema_decay 0.999\
--log_validation\
--output_dir=$OUTPUT_DIR \
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "1.5,2.5,4.5,6.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75
-54
View File
@@ -1,54 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=offline
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
CACHE_DIR=/data/.cache
EXPERIMENT=4step_infer_lq_euler_50_thresh0.1_lrg_0.75
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
export WANDB_DIR=$DATA_DIR/wandb/
torchrun --nnodes 8 --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 $CACHE_DIR \
--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 24\
--sp_size 8\
--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=500\
--validation_steps 125\
--validation_sampling_steps 4 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--ema_decay 0.999\
--log_validation\
--output_dir=$OUTPUT_DIR \
--tracker_project_name PCM \
--num_frames 139 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "1.5,2.5,4.5,6.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75
-55
View File
@@ -1,55 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=offline
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
CACHE_DIR=/data/.cache
EXPERIMENT=lq_euler_50_thresh0.1_lrg_0.75_phase1
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
export WANDB_DIR=$DATA_DIR/wandb/
torchrun --nnodes 1 --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 $CACHE_DIR \
--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 1\
--train_sp_batch_size 1\
--dataloader_num_workers 8\
--gradient_accumulation_steps=1\
--max_train_steps=2000\
--learning_rate=1e-7\
--mixed_precision="bf16"\
--checkpointing_steps=500\
--validation_steps 125\
--validation_sampling_steps 4 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--ema_decay 0.999\
--log_validation\
--output_dir=$OUTPUT_DIR \
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "1.5,2.5,4.5,6.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--multi_phased_distill_schedule "4000-1"
-55
View File
@@ -1,55 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=offline
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
CACHE_DIR=/data/.cache
EXPERIMENT=lq_euler_50_thresh0.1_lrg_0.75_phase1
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
export WANDB_DIR=$DATA_DIR/wandb/
torchrun --nnodes 8 --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 $CACHE_DIR \
--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 24\
--sp_size 8\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=2000\
--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=$OUTPUT_DIR \
--tracker_project_name PCM \
--num_frames 139 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "1.5,2.5,4.5,6.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--multi_phased_distill_schedule "4000-1"
-52
View File
@@ -1,52 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=offline
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
CACHE_DIR=/data/.cache
EXPERIMENT=lq_euler_50_thres0.1_lrg_0.75_bs_64
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
export WANDB_DIR=$DATA_DIR/wandb/
export WANDB_DIR=$DATA_DIR/wandb/
torchrun --nnodes 8 --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 $CACHE_DIR \
--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 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=500\
--validation_steps 125\
--validation_sampling_steps 8 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir=$OUTPUT_DIR \
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1
-52
View File
@@ -1,52 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=offline
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
CACHE_DIR=/data/.cache
EXPERIMENT=4step_infer_shift16_euler_50
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
export WANDB_DIR=$DATA_DIR/wandb/
torchrun --nnodes 1 --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 $CACHE_DIR \
--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 1\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=2000\
--learning_rate=5e-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\
--log_validation\
--output_dir=$OUTPUT_DIR \
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1
-51
View File
@@ -1,51 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_DIR="$HOME"
export WANDB_MODE=offline
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
CACHE_DIR=/data/.cache
EXPERIMENT=4step_infer_shift16_euler_50
OUTPUT_DIR=$DATA_DIR/outputs/$EXPERIMENT
export WANDB_DIR=$DATA_DIR/wandb/
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 $CACHE_DIR \
--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 24\
--sp_size 8\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=2000\
--learning_rate=5e-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\
--log_validation\
--output_dir=$OUTPUT_DIR \
--tracker_project_name PCM \
--num_frames 139 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1
-49
View File
@@ -1,49 +0,0 @@
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=172.23.30.16
torchrun --nnodes 4 --nproc_per_node 4\
--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 1\
--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 125\
--validation_sampling_steps 8 \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/shift1_euler_50_0.75_phase1"\
--tracker_project_name PCM \
--num_frames 163 \
--shift 1 \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1"
-43
View File
@@ -1,43 +0,0 @@
export WANDB_MODE=online
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
DATA_DIR=data
torchrun --nnodes 1 --nproc_per_node 8\
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 1\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--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\
--log_validation\
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_lrg_0.75_phase_ema0.95"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--multi_phased_distill_schedule "4000-1" \
--use_ema \
--ema_decay 0.95
-49
View File
@@ -1,49 +0,0 @@
export WANDB_MODE=online
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
DATA_DIR=data
IP=172.23.30.16
torchrun --nnodes 4 --nproc_per_node 4\
--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 2\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--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\
--log_validation\
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_lrg_0.75_phase_ema0.95_cfg4.5"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--multi_phased_distill_schedule "4000-1" \
--use_ema \
--ema_decay 0.95 \
--distill_cfg 4.5
-53
View File
@@ -1,53 +0,0 @@
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 2\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=4000\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=500\
--validation_steps 125\
--validation_sampling_steps "4,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/lq_euler_50_thresh0.1_lrg_0.75_phase1_ema0.95"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--multi_phased_distill_schedule "4000-1" \
--use_ema \
--ema_decay 0.95
-54
View File
@@ -1,54 +0,0 @@
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 2\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=4000\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=500\
--validation_steps 125\
--validation_sampling_steps "4,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_lrg_0.75_phase_ema0.95_cfg6.0"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "0.5,1.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--multi_phased_distill_schedule "4000-1" \
--use_ema \
--ema_decay 0.95 \
--distill_cfg 6.0
-47
View File
@@ -1,47 +0,0 @@
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 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/shift8_euler_100"\
--tracker_project_name PCM \
--num_frames 163 \
--shift 8.0 \
--validation_guidance_scale 4.5
gsutil cp data/outputs/shift8_euler_100/checkpoint-4000 gs://vid_gen/runlong_temp_folder_for_pandas70m_debugging/fastvid/shift8_euler_100/checkpoint-4000
-54
View File
@@ -1,54 +0,0 @@
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 2\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=4000\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=500\
--validation_steps 125\
--validation_sampling_steps "4,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/lq_euler_50_thresh0.1_lrg_0.75_phase1_ema0.98_cfg4.5"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "0.5,1.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--multi_phased_distill_schedule "4000-1" \
--use_ema \
--ema_decay 0.98 \
--distill_cfg 4.5
-51
View File
@@ -1,51 +0,0 @@
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 2\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=4000\
--learning_rate=3e-7\
--mixed_precision="bf16"\
--checkpointing_steps=500\
--validation_steps 125\
--validation_sampling_steps "4,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/lq_euler_50_thresh0.1_lrg_0.75_phase1_lr_3e-7"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "0.5,1.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--multi_phased_distill_schedule "4000-1"
-54
View File
@@ -1,54 +0,0 @@
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 2\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=4000\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=500\
--validation_steps 125\
--validation_sampling_steps "4,8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/lq_euler_50_thresh0.15_lrg_0.75_phase1_ema0.95_cfg4.5"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "0.5,1.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.15 \
--linear_range 0.75 \
--multi_phased_distill_schedule "4000-1" \
--use_ema \
--ema_decay 0.95 \
--distill_cfg 4.5
-53
View File
@@ -1,53 +0,0 @@
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.142.161
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 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 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_thres0.1_linear_range_0.75_cfg_7.0"\
--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 \
--distill_cfg 6.0 \
--multi_phased_distill_schedule "4000-8" \
-46
View File
@@ -1,46 +0,0 @@
export WANDB_MODE=online
export WANDB_API_KEY=4f6de3765d6464f43e0506ec7d785641af645e73
DATA_DIR=data
IP=172.23.30.16
torchrun --nnodes 4 --nproc_per_node 4\
--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 2\
--train_sp_batch_size 1\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=4000\
--learning_rate=1e-6\
--mixed_precision="bf16"\
--checkpointing_steps=125\
--validation_steps 125\
--validation_sampling_steps "8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_lrg_0.75_reproduce"\
--tracker_project_name PCM \
--num_frames 163 \
--scheduler_type pcm_linear_quadratic \
--validation_guidance_scale "0.5,1.5,2.5,4.5" \
--num_euler_timesteps 50 \
--linear_quadratic_threshold 0.1 \
--linear_range 0.75 \
--multi_phased_distill_schedule "4000-8" \
-51
View File
@@ -1,51 +0,0 @@
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 2\
--dataloader_num_workers 4\
--gradient_accumulation_steps=1\
--max_train_steps=4000\
--learning_rate=5e-6\
--mixed_precision="bf16"\
--checkpointing_steps=64\
--validation_steps 64\
--validation_sampling_steps "8" \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--log_validation\
--output_dir="$DATA_DIR/outputs/lq_euler_50_thres0.1_lrg_0.75_phase1_lr5e-6"\
--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"

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