Compare commits

...
Author SHA1 Message Date
William Lin d03ffcb5f7 [misc] update wechat group link 2026-02-13 01:18:55 -08:00
Kaiqin Kong 31c0f1b341 [feat] Port LingBot-World-Base (Cam) (#1081) 2026-02-10 11:12:33 -08:00
William Lin 4bee0fa199 [misc] cleanup assets/ and demo/ (#1091) 2026-02-10 02:26:09 -08:00
530e6b8363 [Model] LTX 2 Base (#1064)
Co-authored-by: Davids048 <jundasu@ucsd.edu>
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2026-02-10 01:11:17 -08:00
Jinzhe Pan 9ab2725db1 [ci] CI Transformer Tests (#1089) 2026-02-10 01:08:59 -08:00
IshanandJinzhe Pan 0aff68f51d [Feat] Add Stable Diffusion 3.5 (#1075)
Co-authored-by: Jinzhe Pan <eigensystem1318@gmail.com>
2026-02-10 14:31:36 +08:00
ad58f802f3 [Feat] Port LTX2 trainer (#1074)
Co-authored-by: Davids048 <jundasu@ucsd.edu>
Co-authored-by: Matthew Noto <99706358+RandNMR73@users.noreply.github.com>
2026-02-09 17:32:57 -08:00
Wei Zhou 04fa356ee3 [Misc] [Training] Fixed a bunch of bugs in current training pipeline (#1084) 2026-02-09 16:01:05 -08:00
Matthew Notoandgemini-code-assist[bot] f9c076fe2b [misc] add AGENTS.md file (#1085)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-02-09 01:03:49 -08:00
XOR-op f76efe798e [bugfix]: _compile_conditions regression (#1077) 2026-02-06 18:57:43 -08:00
Zhang Peiyuan 09f455233e [misc] readme small fix (#1076) 2026-02-06 16:57:48 -08:00
XOR-op aea300f690 [perf]: use CUDA IPC in multiproc executor to avoid serialization overhead (#1061) 2026-02-06 19:43:18 -05:00
Hao Zhang b92219f6a6 more fix and relocate STA arguments to pipeline config (#1073) 2026-02-06 13:38:35 -08:00
Jinzhe Pan a321b95a8a [Fix] remove video ratio limitation (#1069) 2026-02-05 20:33:08 -08:00
Wei Zhou 98308db7e0 [Feature] [Hy1.5] Support HY1.5 super-resolution pipeline for 1080p videos (#1046) 2026-02-05 16:39:20 -08:00
Hao Zhang c1e18f6722 Some minor fixes (#1068) 2026-02-05 16:28:00 -08:00
William Lin 75e193a2c9 [core] Refactor and centralize our registry for models, pipelines, and sampling params (#1066) 2026-02-05 14:40:30 -08:00
XOR-op d6e0a7d0dd [refactor] Action module (#1065) 2026-02-05 13:56:35 -08:00
William Linandgemini-code-assist[bot] 7fc5f241da [misc] Fix naming instruction in runpod.md (#1067)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-02-05 13:45:49 -08:00
Mingjia Huo aae48a7e90 [feat] HYworld VAE with cache (#1057) 2026-02-05 04:18:17 -08:00
William Lin 88f38eb0f4 [misc] upgrade torch to 2.10 (#1048) 2026-02-05 04:15:22 -08:00
Shao Duan d750b463dc Added Sequence Parallelism for LTX-2 Distilled (#1036) 2026-02-04 15:25:30 -08:00
KyleShaoandWill Lin e10b26a3d8 [feat] Add Cosmos 2.5 I2W/V2W support (staged pipeline + examples) (#1021)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2026-02-04 14:19:30 -08:00
XOR-op 74636ba246 [chore]: use higher precision timestamp in logging (#1062) 2026-02-04 13:22:59 -08:00
Wei Zhou 7e2f3f14e7 [Bugfix] [Wan I2V] Fix CLIP Image encoder config (#1063) 2026-02-04 13:20:19 -08:00
Kaiqin Kong caa1c402ba [bugfix] Double Normalization in Preprocessing Dataset (#1055) 2026-01-31 11:11:59 -08:00
XOR-op 38a6bd93d3 [chore]: update sageattn3 installation instructions (#1050) 2026-01-29 15:45:52 -08:00
alexzms b867ef7e7c [SP Sharding] Fix SP loss sharding on token axis (thw) with padding; add distributed correctness tests (#1045)
Fixes the sequence parallel sharding on t, now SP shards on t*h*w
2026-01-27 22:53:04 -08:00
William Lin 3ae58c277a [docs] Update design overview and add agents tutorial (#1044) 2026-01-27 15:56:58 -08:00
Kaiqin Kong 0c6862ca55 [feature] Add Matrix Game 2.0 training (#1017)
The CI tests are quite unstable, but since multiple CI tests indicates that each individual tests are passed, I think we can merge this.
2026-01-26 19:21:43 -08:00
XOR-opandWill Lin e8c854bcf1 [docs] Offloading instruction (#1022)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2026-01-26 13:40:28 -08:00
William Lin 06860e96fe [docs] Update runpod instructions (#1043) 2026-01-26 13:13:54 -08:00
Matthew Noto 1b503554d1 [bugfix] fix torchvision import (#1039) 2026-01-24 22:37:14 -08:00
Shreejith SGandgemini-code-assist[bot] 351ceb7c59 [bugfix]: handle architectural differences while lora extraction (#1035)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-01-24 15:42:28 -08:00
KyleShao 10875e0d7b [bugfix] Fix NCCL all_gather contiguity + correct ParallelTiledVAE decode tiling threshold (#1037) 2026-01-24 15:37:35 -08:00
alexzms 1eaae8a10b [ci] Increase ci test error threshold (#1038) 2026-01-24 15:36:10 -08:00
Mingjia Huo 59e00f6164 [feat] Add HY-World1.5-Bidirectional-480P-I2V (#1027)
VAE requires further improvement, will raise PR in near future.
2026-01-23 14:18:04 -08:00
745cc05b10 [bugfix] Allow update timesteps for hy1.5 model. (#1033)
Co-authored-by: Davids048 <jundasu@ucsd.edu>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-01-22 22:04:53 -08:00
William Lin c5dc244871 [bugfix] add omegaconf as dep. (#1032) 2026-01-22 11:59:28 -08:00
alexzms dbf3917bf4 [fastvideo-kernel] replace map to index with Triton implementation + add vsa benchmark (#1029) 2026-01-22 11:35:02 -08:00
XOR-op 050f189c95 fix: SP for hunyuanvideo 1.5 (#1026) 2026-01-21 14:40:06 -08:00
Shao DuanandWill Lin 029216029f Added LTX-2 Distilled T2V Generation (#1016)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2026-01-21 14:11:39 -08:00
alexzmsandWilliam Lin 31f44110b5 [kernel] [bugfix] [ci] bump v0.2.4. Fix STA output handling, TurboDiffusion CUDA norm dtypes for fastvideo-kernel unit tests. (#1020)
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
2026-01-19 17:42:42 -08:00
William Lin 21f3ce6577 [kernel] Fix fastvideo-kernel release workflow (#1019) 2026-01-17 15:57:46 -08:00
XOR-op 785d123e36 [feat] Hooks API and layerwise offloading for all DiTs (#1006) 2026-01-17 11:22:02 -08:00
William Lin d58c551c11 [chore] release fastvideo-kernel 0.2.3 (#1018) 2026-01-17 02:24:23 -08:00
alexzms 560628709c [Bug Fix] Add autograd wrapper for block-sparse attention in fastvideo-kernel + fix CMake extension linking (#1015) 2026-01-16 21:16:43 -08:00
William Lin 0f53b51e6c [CI] Fix OOM issues in ssim tests (#1011) 2026-01-16 21:15:20 -08:00
alexzmsandWill Lin 06093a9c4e [CI] SSIM tests optimization: load all model weights from Modal persistent Volume (#958)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2026-01-16 11:50:43 -08:00
KyleShao dbddfab6d2 [feat] Introduce Cosmos 2.5 Text2World pipeline (#974) 2026-01-15 15:09:05 -08:00
William Lin 7188170277 [misc] [bugfix] unpin 'av' in pyproject (#1009) 2026-01-13 15:40:46 -08:00
XOR-op b7f69c2c1d [feat!] Disable FSDP inference by default (#1001) 2026-01-13 14:20:05 -08:00
Loay Rashid 23a4531491 [CI] Fixed Turbodiffusion I2V CI (#1002) 2026-01-13 01:08:58 -08:00
William Linandgemini-code-assist[bot] 7d52ad0118 [ci] temporarily disable turbodiffusion ssim test (#1000)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-01-08 15:43:50 -08:00
Will Lin 4d7bf35fa3 Revert "dit"
This reverts commit a6a9c9ca07.
2026-01-07 03:22:48 -08:00
Will Lin a6a9c9ca07 dit 2026-01-07 03:18:48 -08:00
f4704847c2 [bugfix] Add configs for TurboDiffusion T2V/I2V models (#993)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2026-01-06 16:36:45 -06:00
Shreejith SGandWill Lin d9c996310b [docs]: add LoRA extraction utilities documentation (#992)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2026-01-06 16:36:26 -06:00
Shao Duan d6651afd2e [examples] Added longcat-video python api examples (#994) 2026-01-06 15:03:42 -06:00
William Lin cf67618cad [chore] release 0.1.7 (real) (#980) 2026-01-05 15:47:05 -06:00
William Lin 2f0a2b3c57 [misc] add pin_cpu_memory false for RTX 4090 (#990) 2026-01-05 15:45:35 -06:00
Loay Rashid e7748d9952 [feat] add Turbodiffusion I2V pipeline (#984) 2026-01-05 15:41:23 -06:00
William Lin 8eb3140b2f [misc] pin fastvideo-kernel in .toml file (#989) 2026-01-05 13:42:32 -06:00
Shao Duan d6ddcea682 Add LongCat-Video I2V and Video Continuation (Base, Distillation and Refinement) Support to FastVideo (#953) 2026-01-04 22:20:09 -06:00
William Lin 3559ba2377 [chore] update wechat QR code (#988) 2026-01-04 21:59:38 -06:00
William Lin 61e63ea0d7 [chore] release fastvideo-kernel 0.2.2 (#986) 2026-01-04 21:21:06 -06:00
William Lin 4ce4ac4734 [ci] increase ssim and lora inference test timeout (#985) 2026-01-04 15:08:20 -06:00
William Lin e7f6db9bd1 [docs] Update docs and README (#975) 2026-01-04 14:59:19 -06:00
Ohm-Rishabh d83f45a6a0 Layer offloading (#966) 2026-01-03 21:46:00 -08:00
XOR-op dd91542cd1 [feat] Support text encoder weight override and quantization (#983) 2026-01-03 15:33:33 -06:00
Kaiqin Kong 581e8115fe [feat] support Matrix-Game 2.0 streaming generation (#957) 2026-01-02 19:14:38 -06:00
Loay Rashid dea69cf651 [New Model] Turbodiffusion (#971) 2026-01-02 17:55:56 -06:00
XOR-op 60ac6537df [feat] Support absmax style quantization for FP8 (#981) 2026-01-02 16:00:18 -06:00
Qi Jia 5285116e73 [docs]: fix various broken links across the documentation (#979) 2026-01-01 20:02:39 -06:00
William Lin 40ce2d72f5 [kernel] add turbodiffusion kernels (#972) 2025-12-30 04:23:10 -06:00
William Lin 704bc9aaf9 [misc] Add util script to create diffuser HF repo from custom component weights (#970) 2025-12-29 19:38:30 -06:00
RoyWangandroywang de264fcc99 [fix]: fix STA trition kernel for AMD RDNA archs (#969)
Co-authored-by: roywang <roywang@amd.com>
2025-12-29 14:25:32 -06:00
RoyWangandroywang 7b952e4673 [fix]: fix fastvideo-kernel Rocm build and Dockerfile for Rocm (#968)
Co-authored-by: roywang <roywang@amd.com>
2025-12-29 14:24:46 -06:00
551b2d2048 [fix]: fix sliding_tile_attn with sdpa(without flash_attn) (#967)
Co-authored-by: roywang <roywang@amd.com>
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-12-29 14:19:40 -06:00
Ketaki Tank 7bfaf82fd7 [feat] Add new feature extractors for fvd (#954) 2025-12-27 05:08:26 -06:00
William Lin 9cd6a86b95 [chore] release v0.1.7 (#955) 2025-12-27 05:05:37 -06:00
William Lin 16e9552778 [kernel] Fix docker release build for kernel (#965) 2025-12-26 21:42:55 -06:00
William Lin 87f8a2782d [docs] refactor attention docs (#964) 2025-12-26 15:21:49 -06:00
William Lin cbbb09d7b8 [kernel] Release fastvideo-kernel v0.2.1 (#963) 2025-12-26 13:59:03 -06:00
William LinandShreejithSG 2f6230abcf [kernel] Reorg and fix fastvideo-kernel (#962)
Co-authored-by: ShreejithSG <shreejithsg@gmail.com>
2025-12-26 01:50:24 -06:00
Shreejith SGandWilliam Lin f8bfc76015 feat: consolidate attention kernels into unified fastvideo-kernel package (#946)
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
2025-12-24 01:39:51 -06:00
alexzmsandShao Duan 8f1e6c3336 Add LongCat T2V (Base, Distillation and Refinement) Support to FastVideo (#883)
Co-authored-by: Shao Duan <shaoxiongduan@gmail.com>
2025-12-23 01:11:18 -06:00
William Lin 8e7d2e7879 [bugfix] [dmd2] allow dmd2 simulate_student_forward to use text-only dataset (#951) 2025-12-23 00:41:31 -06:00
William Lin 6ab2870942 [rocm] Add rocm fastvideo docker image (#952) 2025-12-22 18:21:04 -06:00
RoyWang e0ad145152 [feat] add sliding_tile attention triton kernel and ROCM support (#916) 2025-12-22 18:02:51 -06:00
Matthew Noto da04d08426 [docs] small fixes (#947) 2025-12-22 15:23:55 -06:00
Wei Zhou 1f70032af5 [New Model] Hunyuan1.5 (#943) 2025-12-21 00:57:52 -06:00
William Lin 7f71994653 [misc] Allow manual override of Pipeline class through override_pipeline_cls_name (#945) 2025-12-20 14:39:17 -06:00
Loay Rashid 2bb3349da1 [bugfix] Added VSA Padding logic (#944) 2025-12-20 14:29:11 -06:00
Kaiqin Kong 8fe1689968 [feat] Add Matrix-Game 2.0 (#938) 2025-12-20 14:09:12 -06:00
Loay Rashid e53730f324 [docs] Minor Fixes (#942) 2025-12-19 16:48:16 -06:00
Loay Rashid 7a4fe9086a [feat] Support sequence packing and shard after pachification for USP (#894) 2025-12-19 16:19:46 -06:00
Ohm-Rishabh d277361aae [misc] add schedule configurations to pytorch profiler (#934) 2025-12-18 01:45:23 -06:00
alexzms 734a54e7a9 [ci]: Use pre-built docker image & skip VSA compilation (#939) 2025-12-16 23:14:11 -08:00
alexzms 91364982df [Feature] Support for Variable Q/KV Sequence Lengths in VSA ThunderKittens kernel (#911) 2025-12-16 20:08:15 -08:00
William Lin 50145e4fcb [CI] Fix CI tests (#935) 2025-12-16 04:59:43 -08:00
William Lin 4112507e99 [misc] upgrade pytorch version to 2.9.0 (#928) 2025-12-15 04:12:43 -08:00
William Lin 424fc2b4ae [bugfix] [lora] [distillation] Fix lora distillation bug (#933) 2025-12-15 04:12:02 -08:00
William Lin e6066223e6 [bugfix] [VSA] [distillation] Various bugfixes for VSA and distillation and nightly tests (#932) 2025-12-12 16:51:54 -08:00
William Lin b6fa3d24d8 [misc] update wechat image (#931) 2025-12-11 21:22:21 -08:00
Ketaki Tank 55c2e7cd76 [feat] Add fvd implementation (#923) 2025-12-11 19:06:19 -08:00
Tuyabei 5a549af823 [bugfix] [VSA] Fix block_size computation in backward kernel (#925) 2025-12-10 14:36:40 -08:00
Shreejith SG 92fb660c2e Add LoRA extraction, verification, and comparison scripts (#865) 2025-12-08 16:07:58 -08:00
William Lin 3ff640b2e6 [bigfix] [distillation] Fix DMD inference pipeline noise initialization shape (#921) 2025-12-08 13:00:48 -08:00
William Lin c722429ab5 [docs] fix testing.md visibility (#920) 2025-12-08 00:44:53 -08:00
KyleShaoandKyleS1016 e04a192de6 [feat]: add COSMOS 2.5 DiT implementation (#897)
Co-authored-by: KyleS1016 <kyle.s@gmicloud.ai>
2025-12-07 21:48:32 -08:00
William Lin c9ca6d1298 [docs] add docs for ssim testing (#918) 2025-12-06 18:20:04 -08:00
Wenxuan TanandSolitaryThinker 754292c419 Use assert_close in tests (#429)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-12-06 18:18:25 -08:00
Qi Jia 0082bc66fc fix: correct mp backend GPU assignment on multi-GPU systems (#912) 2025-11-30 23:00:22 -08:00
Ohm-Rishabh 8b1937422e [feat] training mfu calculation scripts (#871) 2025-11-27 16:54:17 -08:00
fb6cbf23e6 Fix the docs (#905)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
2025-11-27 00:37:03 -08:00
Mihir Jagtap c8fdd5ed7b [docs] modified the .github/workflows/docs.yml file to include path filtering (#906) 2025-11-26 17:34:21 -08:00
Loay Rashid 1c19a6a00c [Bugfix] Minor bugfixes (#889) 2025-11-26 17:20:45 -08:00
William Lin d44409c704 [CI] fix VSA training CI (#900) 2025-11-24 17:47:59 -08:00
Zhang Peiyuan 5d1c7852b7 + Awesome work using FastVideo or our research projects (#898) 2025-11-23 22:22:27 -08:00
Wenxuan Tan 77a211d006 [misc] Update wechat link (#893) 2025-11-20 19:59:05 -08:00
Wei Zhou bef8169bb1 [Feat] [I2V] resize all image sizes to below 480*832 (#890) 2025-11-20 00:08:36 -08:00
William Lin 681f1583f9 [readme] update link to inference code (#887) 2025-11-19 13:24:13 -08:00
e3b4564d5a [feat] Add inference for MoE SF (#880)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-11-19 13:16:24 -08:00
Shao Duan c0d03fc43d [bugfix] [lora] [CI] Fix LoRA alpha scaling factor & Fix LoRA Inference CI (#870) 2025-11-19 01:02:01 -08:00
Wei Zhou 404ee8538e [Bugfix] [DMD Distillation] Each rank should have its own timestep sampled (#885) 2025-11-18 14:03:25 -08:00
Shao Duan e57ac59462 Fix mp worker busy loop to handle all string RPC methods (#881) 2025-11-16 13:26:44 -08:00
Mihir Jagtap 8c55fdaf7e [docs] add favicon (#878) 2025-11-15 13:44:16 -08:00
Y-aang c30779184f fix: incorrect dv in vsa Triton kernel causing test_vsa error (#879) 2025-11-14 22:00:39 -08:00
William Lin 9d188c0b6c [misc] update wechat and slack invite links (#875) 2025-11-12 23:03:56 -08:00
Mihir Jagtap 9dd7c54221 [docs] Update Home Readme.md with fixed links (#873) 2025-11-12 13:32:44 -08:00
William Lin 62b95d8287 [feat] prepare for wan2.2 SF (#861) 2025-11-04 18:06:48 -08:00
Kaiqin Kong fdf21702f5 [Docs] add diagrams to docs (#863) 2025-11-04 16:29:07 -08:00
Ohm-Rishabh 2972fc9449 Improve FSDP loading with size-based filtering (#853) 2025-11-04 15:31:07 -08:00
Mihir Jagtap 8f5712629f [docs] port to mkdocs (#855) 2025-11-04 14:31:56 -08:00
Kevin Lin 436c701b9f [bugfix] Add Cosmos2 sampling params to registry (#862) 2025-11-02 00:09:17 -07:00
Kevin Lin 543fea88e3 [Feature] Add Cosmos2 i2v pipeline (#837) 2025-10-30 20:03:57 -07:00
Kaiqin Kong bdec816b31 move STA_configuration.py to fastvideo/attention/backends (#856) 2025-10-29 13:54:13 -07:00
William Lin 2cd2e57d2e [ci] fix causal ssim test (#848) 2025-10-26 19:33:07 -07:00
William Linandainsley 9370234294 [feat] Add gradio local inference demo (#847)
Co-authored-by: ainsley <jzhang2765@wisc.edu>
2025-10-26 07:01:33 -07:00
Jinzhe Pan 50da62e722 [bugfix] always force spawn instead of fork (#852) 2025-10-23 16:36:50 -07:00
William Lin 4f3e8751db [bugfix] [misc] Use training_state_checkpointing_steps in scripts/ (#846) 2025-10-19 20:20:53 -07:00
Jinzhe PanandXingyu Long f4c58894d9 [Feat] add ray support (#838)
Co-authored-by: Xingyu Long <xingyulong97@gmail.com>
2025-10-16 23:17:54 -07:00
Ohm-Rishabh 01c94ef385 [feat] unified trainer logging (#841) 2025-10-16 23:16:16 -07:00
Zhang Peiyuan 2415226d25 Update WeChat Link 2025-10-13 21:02:46 -07:00
Jiali Chen 404314d00f [Feature]Add video-to-video (V2V) pipeline (#829) 2025-10-12 21:53:05 -07:00
zyang6andkiritorl 87489f0872 Add wan2.1 functionality support for Ascend NPU platform (#810)
Co-authored-by: kiritorl <1021709528@qq.com>
2025-10-09 16:25:08 -07:00
Zhang Peiyuan 9ce7c8039e Update Wechat link 2025-10-06 15:01:19 -07:00
William Lin e1e25e95f9 [feature] Add torch profiler (#827) 2025-10-06 07:59:46 -07:00
William Lin 490bde90e1 [bugfix] Allow overriding dit checkpoint for inference and Lower VSA LR in example scripts (#831) 2025-10-05 01:44:49 -07:00
dc7596b973 [self-forcing][8/n] Self-Forcing For Wan2.2-A14B + torch.compile training and distillation support (#818)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2025-10-02 15:01:45 -07:00
William Lin 335afa4457 [bugfix] Use training_state_checkpointing_steps instead of checkpointing_steps (#821) 2025-09-28 15:22:43 -07:00
Yongqi Chen 3f77a6805a [Feature]Update count trainable param for FSDP2 (#820) 2025-09-28 15:22:04 -07:00
RandNMR73 13d0aae706 Add Sage Attention 3 Backend (#815) 2025-09-24 15:11:38 -07:00
William Lin 404cbf4f3c [self-forcing] [6/n] Add Ode Init training (#811) 2025-09-22 17:58:19 -07:00
William Lin 958ffec844 [bugfix] Update learning rates for sparse distillation recipe (#812) 2025-09-22 12:07:03 -07:00
31f000d1cc [self-forcing] [5/n] Add Self-Forcing distillation pipeline (#808)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-09-20 19:32:10 -07:00
Yongqi Chen cd32b3e02f Update example files and readme (#809) 2025-09-20 18:15:59 -07:00
Zhang Peiyuan bf27908095 Update WeChat Link 2025-09-20 14:16:20 -07:00
William Lin c5f9ea53b2 [self-forcing] [4/n] Preprocessing for collecting ODE trajectory (#788) 2025-09-15 17:54:42 -07:00
William Lin d32a7184da [bugfix] Wan2.2 Boundary ratio (#804) 2025-09-15 11:17:35 -07:00
Wenxuan Tanandgemini-code-assist[bot] 2930abe456 [Bugfix] Fix VMoba requirements (#802)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2025-09-14 18:28:52 -07:00
William Lin b93ef4289d [bugfix] Fix empty PipelineConfigs for Wan2.2 A14B (#800) 2025-09-13 17:31:38 -07:00
401bdbd316 [self-forcing] [3/n] Text embed only preprocessing (#797)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
Co-authored-by: kevin314 <kevin.lin.cs1@gmail.com>
2025-09-13 14:03:53 -07:00
William Lin 1048d79cf8 [bugfix] pin gradio version and set current_vsa_sparsity in TrainingPipeline (#798) 2025-09-11 17:04:47 -07:00
1e8406162d [bugfix] Fix delta calculation (#796)
Co-authored-by: zbchu2 <zbchu2@iflytek.com>
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
2025-09-11 16:31:23 -07:00
William Lin 03edd35c83 [preprocessing] [self-forcing] [2/n] Improve preprocessing and add ode trajectory dataset schema (#794) 2025-09-10 17:33:57 -07:00
William LinandRandNMR73 ac11127397 [Self-forcing] [1/n] Handle extra dim in time embedding and add timestep warping (#792)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
2025-09-09 02:52:02 -07:00
Eric LiangandEricLiang e028dcc7c0 [Backend][Vmoba] Add implementation of VMoba (#778)
Co-authored-by: EricLiang <https://github.com/EricLina>
2025-09-08 23:53:25 -07:00
Wenxuan Tanandgemini-code-assist[bot] 076f45c1ee [Feature] Support Lora for DMD (#755)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2025-09-08 14:18:21 -07:00
85eb7265db fix: lora_B init zeros (#781)
Co-authored-by: zbchu2 <zbchu2@iflytek.com>
Co-authored-by: Wenxuan Tan <wenxuan.tan@wisc.edu>
2025-09-05 22:56:52 -07:00
William Lin d3ceb67e66 [misc] Update Slack invite link (#786) 2025-09-05 12:16:18 -07:00
Zhang Peiyuan 7ac153a5ca Update WeChat Link 2025-09-05 11:40:47 -07:00
William Lin d1e7aa0abd [CI] Add ssim test for causal inference (#784) 2025-09-05 01:23:01 -07:00
William Lin 2d846c55a1 [misc] Improve text encoding stage (#774) 2025-09-04 17:51:27 -07:00
Jinzhe Pan b318063c0a [Preprocess][Fix] video quality issue (#773) 2025-09-03 20:47:33 -07:00
Jinzhe Pan 4aa307be55 [Preprocess][Feat] support torchvision to load video in new preprocessing (#761) 2025-09-01 23:37:01 -07:00
William Lin 055e52e5ea [misc] [VSA] [STA] fix tk_root in setup.py for VSA and STA (#772) 2025-08-29 01:13:37 -07:00
William Lin 7d2069596b [bugfix] [VSA] [STA] Fix MANIFEST.in for VSA and STA; Move tk into both directories (#771) 2025-08-29 00:51:05 -07:00
William Lin c45009c9a4 [bugfix] fix STA install setup.py import (#770) 2025-08-28 23:02:53 -07:00
William LinandPeiyuan Zhang b91020b407 [VSA] [STA] Fix directory structure for pypi publishing (#769)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-08-28 22:34:03 -07:00
William Lin 2dcc5ea4f6 [chore] Release 0.1.6 (#768) 2025-08-28 20:56:21 -07:00
Wei ZhouandSolitaryThinker 359151d9a0 [Feature] Add wan2.2 5b i2v (#760)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-08-28 18:15:59 -07:00
Wei ZhouandSolitaryThinker ce67cd3729 [Feat] Support Self-Forcing's Causal Inference for Wan2.1 T2V 1.3B (#766)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-08-28 16:47:49 -07:00
Zhang Peiyuan 7c554e5da8 Update Community Link (#765) 2025-08-27 16:12:47 -07:00
William Lin 663ea33ff1 [bugfix] Fix wrong HF model string for FastWan2.2 5B (#763) 2025-08-26 22:05:40 -07:00
William Lin 3ef04f1654 [misc] [docs] Various fixes for logging and docs (#758) 2025-08-23 21:13:50 -07:00
Jinzhe Pan 0eced76a41 [Feat][Preprocess] support multi-gpus (#753) 2025-08-23 11:34:42 +08:00
Jinzhe Pan 3ab6470d1a [Feat][Preprocess] support merged dataset (#752) 2025-08-22 15:29:33 -07:00
Wenxuan Tan 989a03532c Optionally use unmerged weights for inference (#745) 2025-08-22 15:20:31 -07:00
William Lin fa15369a02 [bugfix] Check that model_index.json module is in required_modules list before removing (#756) 2025-08-22 14:36:44 -07:00
Zhang Peiyuan 78a9cb88d8 [Fix] fix seed in dmd denoising loop (#736) 2025-08-21 18:06:16 -07:00
Peng Xiaoand肖鹏 a0bff12746 [bugfix] [dmd] Align backward simulation with dmd2 sample back (#744)
Co-authored-by: 肖鹏 <xiaopeng1@aishi.ai>
2025-08-20 22:25:33 -07:00
William Lin 98f2af94e5 [bugfix] Missing Docker file for cuda12.9 (#750) 2025-08-20 15:34:31 -07:00
William Lin 46f7b6d574 [Docker] add 12.9 docker image and also fix py3.10 and py3.11 dockerfile (#749) 2025-08-20 15:31:15 -07:00
Jinzhe Pan 911a6a6a35 [Feat][Preprocessing] i2v preprocessing workflow (#737) 2025-08-14 20:47:25 -07:00
Zhang Peiyuan 38c7949d5c Update WeChat group link (#739) 2025-08-14 15:03:35 -07:00
Jinzhe Pan 7e7a0dba9d feat: preprocess validation dataset only when exist (#734) 2025-08-12 02:16:31 -07:00
Zhang Peiyuan f62e210ae6 Fix vsa backward gQ (#735) 2025-08-11 21:43:13 -07:00
William Lin 6ceb4942a0 [bugfix] [dmd] Fix backward simulation and also naming in wan_i2v_dmd_pipeline (#731) 2025-08-10 21:13:30 -07:00
William LinandRandNMR73 8cae5e4708 [feature] add Gradio live serving demo code (#727)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
2025-08-10 15:34:03 -07:00
William Lin 2a773fa34e [bugfix] [distill] remove i2v validation schema import in distill (#728) 2025-08-09 20:47:42 -07:00
Wenxuan Tan 5357f63327 Fix LoRA load from training checkpoint (#719) 2025-08-09 20:46:00 -05:00
William Lin 60f61c8101 [bugfix] fix pyproject install and VSA precision test (#726) 2025-08-08 18:45:03 -07:00
Jiali Chen 3d75ba8251 update version selection for VSA workflow (#725) 2025-08-08 13:05:16 -07:00
Wenxuan Tan 6c6bcd914d Remove all empty_cache (#713) 2025-08-07 22:50:38 -07:00
Jiali Chen f79b08de81 add cicd workflow for publishing VSA kernel (#723) 2025-08-07 18:53:05 -07:00
Jinzhe Pan f2bc037fff [Fix] training pipeline pin_cpu_memory issue (#692) 2025-08-07 02:31:20 -07:00
Jinzhe Pan 86604a684b [3/3][Preprocess] add preprocessing workflows (#645) 2025-08-07 01:49:07 -07:00
Zhang Peiyuan 47bd1e0178 [Misc] change installation logic of vsa (#721) 2025-08-06 21:54:09 -07:00
Wei ZhouandSolitaryThinker c41305ad18 [Feat] Add Wan2.2 14B MoE (#688)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-08-06 20:31:03 -07:00
Zhang Peiyuan 98ce9034f0 [Chore] Include our demo in the readme. (#720) 2025-08-06 19:29:40 -07:00
William Lin 0ceff110da [chore] Release 0.1.5 (#717) 2025-08-06 13:07:52 -07:00
Yongqi Chen 1d018acb3e [Feature]Add Data-free distillation readme (#710) 2025-08-05 14:27:39 -04:00
Yongqi Chen 7d8cf38dbe Fix typo (#709) 2025-08-04 20:21:21 -07:00
Yongqi Chen 8d483fe4aa [Bugfix] Fix neg_prompt bug when training from local cp (#708) 2025-08-04 15:54:06 -07:00
Zhang Peiyuan c1191250bf Add WeChat group link (#707) 2025-08-04 15:19:01 -07:00
Wenxuan Tan 4b7266349a [misc] Remove allow_tf32 in scripts (#705) 2025-08-04 15:37:56 -05:00
Yongqi Chen 22f9b7681f [Feature]Update Wan2.2+DMD doc example (#706) 2025-08-04 16:14:22 -04:00
Yongqi Chen 589d32cc39 [Feature] Update Readme and scripts (#703) 2025-08-04 15:02:32 -04:00
Hao Zhang 89199837db Update readme pre-release (#704) 2025-08-04 11:54:31 -07:00
William Lin d6ebaf1b49 [Docs] Fix README (#701) 2025-08-04 11:27:27 -07:00
Yongqi Chen fac927777c [Feature] Update readme (#702) 2025-08-04 14:27:18 -04:00
William Lin ecbd697dae [misc] Readme fixes (#699) 2025-08-04 10:20:57 -07:00
Yongqi Chen 7d4acef64d [Feature] Update sparse distill readme and doc (#700) 2025-08-04 10:16:20 -07:00
William Lin 9f0ce517cf [Docs] Update README and docs for FastWan (#698) 2025-08-04 09:05:18 -07:00
Yongqi Chen c718e56b0d [Feature] Remove unused args (#695) 2025-08-03 23:01:54 -04:00
Yongqi ChenandSolitaryThinker b65f0316d1 [Feature] Add Wan2.2 DMD example files; Update lr scheduler (#694)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-08-03 22:55:01 -04:00
William Lin 8d8bcb76b0 [config] Add config for FastWan2.2 ti2v 5B (#693) 2025-08-03 19:09:30 -07:00
Yongqi ChenandSolitaryThinker 5f42748ed1 [Feature] Add Wan2.2-TI2V-5B Sparse Distill (#690)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-08-03 01:29:57 -04:00
Yongqi Chen c9005045dc [Feature[[Readme] Add VSA/DMD doc (#673) 2025-08-02 02:35:46 -04:00
Wenxuan Tan 6c81befc87 [Feature] Optionally enable torch compile (#684) 2025-08-01 20:17:40 -07:00
Yongqi Chen dfe0b288e1 [Bugfix] Add i2v vae loading (#686) 2025-08-01 23:15:37 -04:00
Wenxuan Tanandgemini-code-assist[bot] 31200fbb83 [Misc] Fix training scripts (#683)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2025-08-01 15:54:31 -05:00
Yongqi Chen 9185978c55 [Bugfix] Fix multi-gpu training lr_scheduler (#682) 2025-08-01 15:54:03 -04:00
MartinPernus fcba463553 [Bugfix] fix _normalize_dit_input (#681) 2025-08-01 05:10:27 -04:00
Yongqi Chen 2c53d3eecf [Feature]Add DMD visualization for debugging (#674) 2025-07-31 05:54:47 -04:00
Zhang Peiyuan 516ecd374a [Misc] Update examples/ and other misc (#672) 2025-07-30 19:10:27 -07:00
Wei Zhou 3b1b54a74d Modify args to make sure the scripts are runnable on 4090 (#671) 2025-07-30 14:55:08 -07:00
Yongqi Chen 6914e7c904 [Bugfix]Fix DMD pipeline registry (#670) 2025-07-30 13:21:35 -07:00
Sopiko Kurdadze 5452369749 [Feature] [Inference]Add ROCm platform support for single-gpu inference (#669) 2025-07-30 12:57:02 -07:00
Yongqi Chen a113311e77 [Bugfix][Training]Fix Wan2.2 training vae config issue (#668) 2025-07-30 12:24:45 -07:00
Kevin Lin 44da97da92 [chore] Release 0.1.4 (#667) 2025-07-30 01:02:36 -07:00
Yongqi Chen f759980a58 [Feature]Add VSA slurm training example scripts (#666) 2025-07-30 01:27:54 -04:00
Zhang Peiyuan 37e0f8c236 [BUG] Fix distillation + vsa (#665) 2025-07-29 19:48:11 -07:00
Kevin Lin 51711d5906 [ComfyUI] Add __init__.py for node discovery (#663) 2025-07-29 18:21:09 -07:00
William LinandJerryZhou54 6375223b16 [Feature] Add wan2.2 5B T2V (#658)
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
2025-07-29 17:16:03 -07:00
Yongqi Chen 4cb046768d [Feature]Add DMD distillation training resume checkpoint; Update DMD CI test (#662) 2025-07-29 19:06:05 -04:00
Yongqi Chen 3322542444 [Feature] Add DMD CI test (#661) 2025-07-29 03:14:00 -04:00
Yongqi Chen 65f707354b [Bugfix]Fix mdoel inference checkpoint saving when enabling HSDP (#660) 2025-07-28 22:39:06 -07:00
Yongqi Chen 109e2e7e9d [Bugfix]Fix DMD wan pipeline (#659) 2025-07-28 21:50:44 -07:00
Zhang Peiyuan cbc3a6bb9d [Feat] Support VSA with any resolution. (#650) 2025-07-28 20:14:40 -07:00
Yongqi Chen 2fa8d4ae6d [Feature][Distill]Add 14B 480p T2V distill example scripts (#655) 2025-07-28 18:36:31 -04:00
Jinzhe Pan 7b6c8aee99 [2/3][Preprocess] refactor pipeline registry & file structure (#639) 2025-07-27 23:30:17 -07:00
Yongqi Chen 6284eaa363 [Feature][Distill]Add DMD+VSA joint training example (#654) 2025-07-27 18:18:01 -04:00
Yongqi Chen 636524e87f [Feature] Add Wan-14B-T2V-VSA CLI inference; add master port args (#653) 2025-07-27 07:13:44 -04:00
Yongqi Chen 202b2f3972 [Feature] Ignore [union-attr] and [override] mypy check and remove from training (#652) 2025-07-27 04:34:38 -04:00
Yongqi Chen 247fe273d8 [Feature] Add DMD T2V training pipeline (#651) 2025-07-27 03:35:51 -04:00
William Lin cb320dfa3a [bugfix] VideoGenerator improperly extracts output_video_name (#649) 2025-07-26 19:46:26 -07:00
Kevin Lin d8bb5abc46 [CI] Fix ComfyUI publisher ID (#648) 2025-07-25 19:02:52 -07:00
Kevin Lin cc703eca51 [CI] Add publish workflow for ComfyUI (#647) 2025-07-25 18:39:05 -07:00
William Lin 81c9df629c [core] Add offloading for vae and image encoder and rename offloading args (#643) 2025-07-25 17:55:03 -07:00
Yongqi Chen d3c0c52208 [Feature] Add prompt_txt support for CLI inference; Add DMD CLI inference (#646) 2025-07-25 19:44:10 -04:00
William Lin 744e0555c0 [misc] Use FASTVIDEO_STAGE_LOGGING for perf timing of stage (#644) 2025-07-25 16:20:53 -07:00
Jinzhe Pan 3a38f7dfdc [1/3][Preprocess] refactor preprocessing configs (#638) 2025-07-25 14:27:12 -07:00
William Lin f572319bd9 [Feature] Remove V1 folder (#642) 2025-07-24 22:43:12 -07:00
Wenxuan TanandWei Feng 48528f468c [Feature] Multi-lora inference (#640)
Co-authored-by: Wei (Will) Feng <134637289+weifengpy@users.noreply.github.com>
2025-07-24 21:01:28 -07:00
Yongqi ChenandSolitaryThinker 4264a80ca9 [Feature] Add DMD inference pipeline (#637)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-07-24 21:00:47 -07:00
William Lin 8573d4f05e [Docs] Docs update for Training and MPS (#641) 2025-07-24 19:23:58 -07:00
Wenxuan Tan 210a733515 [Bugfix] Fix LoRA trainable params and training ckpt loading (#630) 2025-07-23 20:01:40 -07:00
William Lin 0aef0e6f63 [bugfix] Fix preprocessing pipelines and nightly tests (#633) 2025-07-22 22:44:25 -07:00
Kevin Lin dd022ad9be [CI] Fix CI for pull request targets other than main (#632) 2025-07-22 21:00:03 -07:00
William Lin 832ad61e5b [bugfix] fa3 no longer returns lse (#631) 2025-07-22 18:31:34 -07:00
Wenxuan Tan 9419c04ee3 Fix lora train steps (#627) 2025-07-21 23:29:30 -05:00
Zhang Peiyuanandroot a37b39d83c Py/add triton block sparse (#593)
Co-authored-by: root <a1286225768@gmail,com>
2025-07-17 16:50:40 -07:00
Wenxuan Tan bb8c769c8e [LoRA] Support v1 LoRA training (#576) 2025-07-17 15:52:09 -05:00
William Lin 576c214f28 [v0] Remove V0 code (#621) 2025-07-15 22:05:16 -07:00
RandNMR73 b79d1fc15b video gen working on apple silicon (addressed issues from prior pr) (#595) 2025-07-15 22:04:27 -07:00
William Lin eb66e1c18d [chore] release 0.1.2 (#622) 2025-07-15 14:42:54 -07:00
William Lin 616d43c1cf [bugfix] [training] use separate generator for validation (#610) 2025-07-15 13:19:38 -07:00
Wenxuan Tan 7244a4b27f [CI] Add LoRA inference tests (#546) 2025-07-15 15:06:44 -05:00
Yongqi Chen 7e5ebb4582 [Feature][Training]Update example fine-tuning scripts to enable gradient checkpointing (#618) 2025-07-15 11:22:56 -07:00
Wenxuan Tan 65ed588570 Set encoder TP size to 1 by default (#569) 2025-07-09 17:44:24 -05:00
Wenxuan Tan 14adfe2edc Remove all unnecessary torch.cuda.empty_cache (#606) 2025-07-09 16:45:20 -05:00
William Lin 6198c6a640 [docs] update dev guide runpod image to py3.12 (#602) 2025-07-07 19:56:07 -05:00
William Lin e6b71b531b [docs] Update slack invite (#601) 2025-07-07 14:20:42 -05:00
William Lin ae1d112c6a [bugfix] [training] fix deadlock in latent datasets and init error in multi-node training (#598) 2025-07-06 01:34:15 -05:00
William Lin bf4de1f38f [chore] Upgrade min Python version from 3.8 to 3.10 (#597) 2025-07-04 22:08:31 -05:00
William Lin 66fdcc8e76 [Training] Use inference pipeline for training validation (#585) 2025-07-04 17:23:55 -05:00
Wenxuan Tan ed1e8d6bad [Feature] Offload all text encoders by default (#594) 2025-07-03 19:30:14 -05:00
Kevin Lin b9423ca3f8 Add ComfyUI custom node for inference (#596) 2025-07-03 16:09:03 -05:00
Wenxuan Tan ad16289871 [LoRA] Fix lora merge weights (#579) 2025-07-02 00:12:52 -05:00
Wenxuan Tan 2a41da1e6b Fix VAE precisions (#588) 2025-07-01 14:49:04 -05:00
William Lin 32133171da [chore] Release 0.1.1 (#592) 2025-07-01 01:43:12 -05:00
Kevin Lin 19674c6f29 [CI] Fix fork builds (#590) 2025-07-01 01:03:34 -05:00
Yongqi Chen 508afb7002 [docs] Update Readme (#591) 2025-06-30 23:48:45 -05:00
Yongqi Chen 288ea88105 [Feat][Training] Rename weight conversion function and update gradient checkpoint in scripts (#589) 2025-07-01 00:20:02 -04:00
Jinzhe Pan eb0f1318f3 [Feat] activation checkpointing (#584) 2025-06-30 15:24:29 -05:00
William Lin ce9b5910cc [Training] add caption to validation log (#582) 2025-06-30 02:42:17 -05:00
William Lin d0e5a6214a [misc] [training] Add --video_length_tolerance_range 10 to preprocessing scripts (#581) 2025-06-30 02:21:22 -05:00
Wenxuan Tan 834562b2db [CI] Fix pre-commit CI (#578) 2025-06-29 16:52:29 -05:00
Wei (Will) Feng 060cc7b9ba fully_shard usage on RMSNorm (#577) 2025-06-29 16:35:24 -05:00
Yongqi Chen 6c58a5ba62 [Bugfix]Fix VSA sp for training/inference (#574) 2025-06-29 13:44:33 -05:00
William Lin 48d9f61f86 [ci] [misc] fix training test threshold (#573) 2025-06-28 22:17:18 -05:00
William Lin 5f938b5844 [Revert] "[Feature] Load weights from distributed" (#571) 2025-06-28 20:55:14 -05:00
Wenxuan Tan 74da2a7370 Fix CLIP config (#568) 2025-06-28 19:01:23 -05:00
Kevin Lin 580d6dfe1f [CI] Add tests to Modal (#562) 2025-06-28 14:02:16 -05:00
Wenxuan Tan 344e43006a [CI] Fix SSIM and transformers CI (#564) 2025-06-28 00:26:20 -05:00
Wenxuan Tan c5155b256e [Feature] Load weights from distributed (#470) 2025-06-27 22:52:40 -05:00
William Lin e005c7f3ac [Docs] [Training] add readme for example training (#563) 2025-06-27 14:50:42 -05:00
Yongqi Chen ff5a79ef60 [Feature][Inference] Add VSA inference script (#561) 2025-06-27 02:19:23 -05:00
William Lin ab01dc4ba5 [Feature] [Training] Add i2v training (#559) 2025-06-27 01:56:50 -05:00
William Lin 285a950c1b [CI] fix vae and ssim tests (#557) 2025-06-26 23:53:01 -05:00
William Lin 46a0a85d85 [Training] Fixes SP for training; Improve Datasets and schema (#555) 2025-06-26 21:13:28 -05:00
Yongqi Chen 4aeabbc629 [Feature][Training] Add cfg rate for dataset loader (#556) 2025-06-26 18:22:37 -04:00
Wenxuan Tan 949bb5c835 [CI] Fix CI checks (#553) 2025-06-25 14:07:51 -05:00
Wenxuan Tan aab74c1271 [Kernel] Remove all syncs from STA & VSA kernels (#517) 2025-06-23 13:13:09 -07:00
Yongqi Chen f89d86944f [Feature][Training]Add diffusers format checkpoint saving for inference (#542) 2025-06-22 01:23:41 -04:00
William Lin 8741d204a5 [Training] Refactor and improve validation datasets (#539) 2025-06-21 17:58:35 -07:00
Wenxuan Tan cdc85f58a8 [chore] Bump torch to 2.7.1 to support Blackwell (#483) 2025-06-20 22:10:56 -07:00
William Lin 0262d2f089 [misc] [training] Reorganize training pipeline (#533) 2025-06-20 20:42:25 -07:00
William Lin 62c0343465 [bugfix] [VSA] Fix layernorm type for VSA Wan2.1 TransformerBlock (#534) 2025-06-20 00:24:51 -07:00
William Lin 1e1a023fb0 [bugfix] Fix stage validator for multi text encoder models (#535) 2025-06-19 22:49:16 -07:00
William Lin 1d2517ad8e [misc] Remove gradient checking code (#532) 2025-06-18 23:29:25 -07:00
William Lin d41186cb4a [Feat] Add Stage input and output verification (#523) 2025-06-18 23:29:11 -07:00
78e0c7eec9 Specify cu128 Pytorch installation (#530)
Co-authored-by: Edenzzzz <wtan45@wisc.edu>
Co-authored-by: Wenxuan Tan <wenxuan.tan@wisc.edu>
2025-06-18 20:02:50 -05:00
Wenxuan Tan 1c41a94b62 [Refactor] Move dict_to_3d_list under utils (#507) 2025-06-18 13:34:37 -07:00
Yongqi Chen 2e66aafe20 [Bugfix][Readme]Fix readme website bugs and add VSA finetune docs (#531) 2025-06-17 22:48:29 -07:00
Yongqi ChenandWill Lin 55074bda76 [CI] Add STA-inference/VSA-training test (#527)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-06-17 21:13:06 -07:00
William Lin de65bec2b7 [Ci] add sta and vsa install to docker image (#528) 2025-06-17 18:09:48 -07:00
Yongqi Chen 7664dd0de3 [Bugfix][Inference]Fix envs.attn_backend (#525) 2025-06-17 18:38:06 -05:00
William Linandkevin314 019a88ced4 [CI][bugfix] Use new 3.12 docker image (#526)
Co-authored-by: kevin314 <kevin.lin.cs1@gmail.com>
2025-06-17 15:37:08 -07:00
Kevin Lin 72de11abcc [CI] Add current PR test workflow to Buildkite/Modal (#512) 2025-06-17 13:29:22 -07:00
Kevin Lin d71a4ebffc [CI] Update Docker image to flash-attn 2.8.0 / CUDA 12.8 (#524) 2025-06-16 17:48:23 -07:00
William Lin 1089ab43bf [bugfix] [Training] use diffusers fp32layernorm for wan2.1 (#490) 2025-06-15 22:45:48 -07:00
William Lin 97d4b984c9 [misc] [ci] fix e2e preprocess+training data path (#521) 2025-06-14 22:37:51 -07:00
Wenxuan Tan 2a8953d74d [Refactor] Fix attn backend selection not correctly setting env variable (#516) 2025-06-15 00:04:54 -05:00
Yongqi Chen 8801b10da7 [Bugfix][Preprocess]fix mini dataset name (#520) 2025-06-14 22:03:22 -07:00
William Lin 6b413f2ec4 [CI] [Training] drop negative prompt in validation dataset and CI test for preprocess + training overfit (#519) 2025-06-14 18:50:17 -07:00
Yongqi Chen 28b72694aa [Feature][Preprocess]Add Readme doc for preprocess (#518) 2025-06-14 20:41:13 -04:00
Yongqi Chen 4afb0cfe4f [Feature][Training]vsa for t2v training ready (#513) 2025-06-14 01:08:00 -04:00
Zhang Peiyuan 3eec1281cf [misc] Fix preprocessing and dataloader extra padding (#514) 2025-06-13 15:15:33 -07:00
Wenxuan Tan 0660489e38 [CI] Restrict training CI to v1 (#508) 2025-06-12 15:26:05 -07:00
Zhang Peiyuan dd871a17bf fix logging (#509) 2025-06-12 15:24:12 -07:00
dc11529862 [Refactor][Configurations] clean config orgnization (#505)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-06-12 13:27:08 -07:00
Zhang Peiyuan ffabf85e31 [feat] Add parquet iterable dataset. (#506) 2025-06-12 04:30:56 -04:00
William Lin c0026ca5ba [CI] [Training] Initial e2e small training test (#504) 2025-06-11 13:53:36 -07:00
Zhang Peiyuan 0f2bbe71ac [misc] rename dp_size to hdsp_replicate_dim (#491) 2025-06-10 16:36:56 -07:00
Yongqi ChenandJerryZhou54 2a46902ecb [Feature][VSA]Update STA publish workflow (#498)
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
2025-06-10 19:33:34 -04:00
Zhang Peiyuan 66012d3a4c [Feat][Dataloader] 1/n Refactor parquet map-style dataloader (#492) 2025-06-10 16:00:13 -07:00
William Lin f666b9de41 [misc] Add missing license headers (#499) 2025-06-10 14:25:32 -07:00
Yongqi Chen 7e3c073b55 [Feature] Adding VSA inference (#478) 2025-06-10 16:03:53 -04:00
Wei Zhou a6aa21bd07 [bugfix][Cli Inference] Resolve runtime errors when running fastvideo generate (#495) 2025-06-10 02:50:21 -04:00
Wenxuan Tan 6519b57aab [chore] Fix main pre-commit CI failure (#494) 2025-06-10 00:54:37 -05:00
Wei Zhou 675aea6ece [bugfix][Cli Inference] Resolve runtime errors when running fastvideo generate (#493) 2025-06-09 19:49:34 -07:00
Zhang Peiyuan 46e7a15e0d [misc] Improve distributed related env variables and setup (#487) 2025-06-08 09:14:48 -07:00
Yongqi ChenandJerryZhou54 e4f702d7ec [Bug] Fix multi gpus issues in v1 scripts (#489)
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
2025-06-07 21:55:49 -07:00
Wenxuan Tan bb68fcc809 Revert "Add torch.compile for all small ops" (#484) 2025-06-07 07:21:17 -05:00
Wenxuan Tan b392e6a874 Add torch.compile for all small ops (#432) 2025-06-06 21:10:42 -07:00
Zhang PeiyuanandWill Lin 0991003905 [bugfix] [misc] fix denoising stage init; rename distributed env function; fix logging. (#481)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-06-06 20:01:10 -07:00
Zhang PeiyuanandWill Lin 8f8ce6d9e1 [bugfix] [training] Add negative prompt to preprocessing and validation (#479)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-06-06 11:21:21 -07:00
1062 changed files with 125198 additions and 41841 deletions
+196
View File
@@ -0,0 +1,196 @@
env:
IMAGE_VERSION: "py3.12-latest"
BUILDKITE_CLEAN_CHECKOUT: true
steps:
- label: "pre-commit"
command: ".buildkite/scripts/pre_commit.sh"
agents:
queue: "default"
- wait
- label: "Trigger Tests"
plugins:
- monorepo-diff#v1.4.0:
diff: 'git fetch origin "$BUILDKITE_PULL_REQUEST_BASE_BRANCH" && git diff --name-only origin/"$BUILDKITE_PULL_REQUEST_BASE_BRANCH"...HEAD'
watch:
- path:
- "fastvideo/models/encoders/**"
- "fastvideo/models/loader/**"
- "fastvideo/tests/encoders/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 20m .buildkite/scripts/pr_test.sh"
label: "Encoder Tests"
env:
- TEST_TYPE=encoder
agents:
queue: "default"
- path:
- "fastvideo/models/vaes/**"
- "fastvideo/models/loader/**"
- "fastvideo/tests/vaes/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 20m .buildkite/scripts/pr_test.sh"
label: "VAE Tests"
env:
- TEST_TYPE=vae
agents:
queue: "default"
- path:
- "fastvideo/models/dits/**"
- "fastvideo/models/loader/**"
- "fastvideo/tests/transformers/**"
- "fastvideo/layers/**"
- "fastvideo/attention/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Transformer Tests"
env:
- TEST_TYPE=transformer
agents:
queue: "default"
- path:
- "fastvideo/**/*.py"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 90m .buildkite/scripts/pr_test.sh"
label: "SSIM Tests"
env:
- TEST_TYPE=ssim
agents:
queue: "default"
- path:
- "fastvideo/tests/lora/**"
- "fastvideo/models/loader/**"
- "fastvideo/tests/transformers/**"
- "fastvideo/pipelines/**"
- "fastvideo/layers/lora/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 20m .buildkite/scripts/pr_test.sh"
label: "LoRA Inference Tests"
env:
- TEST_TYPE=inference_lora
agents:
queue: "default"
- path:
- "fastvideo/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Training Tests"
env:
- TEST_TYPE=training
agents:
queue: "default"
- path:
- "fastvideo/training/*distillation_pipeline.py"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Distillation DMDTests"
env:
- TEST_TYPE=distillation_dmd
agents:
queue: "default"
- path:
- "fastvideo/training/*self_forcing_distillation_pipeline.py"
- "fastvideo/tests/training/self-forcing/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Self-Forcing Tests"
env:
- TEST_TYPE=self_forcing
agents:
queue: "default"
- path:
- "fastvideo/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "LoRA Training Tests"
env:
- TEST_TYPE=training_lora
agents:
queue: "default"
- path:
- "fastvideo/**"
- "fastvideo-kernel/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Training Tests VSA"
env:
- TEST_TYPE=training_vsa
agents:
queue: "default"
- path:
- "fastvideo/**"
- "fastvideo-kernel/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Inference Tests STA"
env:
- TEST_TYPE=inference_sta
agents:
queue: "default"
- path:
- "fastvideo-kernel/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Kernel Tests"
env:
- TEST_TYPE=kernel_tests
- path:
- "fastvideo-kernel/**"
- "fastvideo/attention/backends/vmoba.py"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Inference Tests VMoBA"
env:
- TEST_TYPE=inference_vmoba
agents:
queue: "default"
- path:
- "fastvideo/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Unit Tests"
env:
- TEST_TYPE=unit_test
agents:
queue: "default"
# - path:
# - "scripts/lora_extraction/**"
# - "pyproject.toml"
# - "docker/Dockerfile.python3.12"
# config:
# command: "timeout 90m .buildkite/scripts/pr_test.sh"
# label: "LoRA Extraction Tests"
# env:
# - TEST_TYPE=lora_extraction
# agents:
# queue: "default"
+142
View File
@@ -0,0 +1,142 @@
#!/bin/bash
set -uo pipefail
log() {
echo "[$(date '+%Y-%m-%d %H:%M:%S')] $1"
}
log "=== Starting Modal test execution ==="
# Change to the project directory
cd "$(dirname "$0")/../.."
PROJECT_ROOT=$(pwd)
log "Project root: $PROJECT_ROOT"
# Install Modal if not available
if ! python3 -m modal --version &> /dev/null; then
log "Modal not found, installing..."
python3 -m pip install modal
# Verify installation
if ! python3 -m modal --version &> /dev/null; then
log "Error: Failed to install modal. Please install it manually."
exit 1
fi
fi
log "modal version: $(python3 -m modal --version)"
# Set up Modal authentication using Buildkite secrets
log "Setting up Modal authentication from Buildkite secrets..."
MODAL_TOKEN_ID=$(buildkite-agent secret get modal_token_id)
MODAL_TOKEN_SECRET=$(buildkite-agent secret get modal_token_secret)
# Retrieve other secrets
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
HF_API_KEY=$(buildkite-agent secret get hf_api_key)
if [ -n "$MODAL_TOKEN_ID" ] && [ -n "$MODAL_TOKEN_SECRET" ]; then
log "Retrieved Modal credentials from Buildkite secrets"
python3 -m modal token set --token-id "$MODAL_TOKEN_ID" --token-secret "$MODAL_TOKEN_SECRET" --profile buildkite-ci --activate --verify
if [ $? -eq 0 ]; then
log "Modal authentication successful"
else
log "Error: Failed to set Modal credentials"
exit 1
fi
else
log "Error: Could not retrieve Modal credentials from Buildkite secrets."
log "Please ensure 'modal_token_id' and 'modal_token_secret' secrets are set in Buildkite."
exit 1
fi
MODAL_TEST_FILE="fastvideo/tests/modal/pr_test.py"
if [ -z "${TEST_TYPE:-}" ]; then
log "Error: TEST_TYPE environment variable is not set"
exit 1
fi
log "Test type: $TEST_TYPE"
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$BUILDKITE_PULL_REQUEST IMAGE_VERSION=$IMAGE_VERSION"
case "$TEST_TYPE" in
"encoder")
log "Running encoder tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
;;
"vae")
log "Running VAE tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
;;
"transformer")
log "Running transformer tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
;;
"ssim")
log "Running SSIM tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
;;
"training")
log "Running training tests..."
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests"
;;
"training_lora")
log "Running LoRA training tests..."
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_lora_tests"
;;
"training_vsa")
log "Running training VSA tests..."
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests_VSA"
;;
"inference_sta")
log "Running inference STA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_STA"
;;
"kernel_tests")
log "Running kernel tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_kernel_tests"
;;
"inference_lora")
log "Running LoRA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_lora_tests"
;;
"distillation_dmd")
log "Running distillation DMD tests..."
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_distill_dmd_tests"
;;
# run_inference_tests_vmoba
"self_forcing")
log "Running self-forcing tests..."
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_self_forcing_tests"
;;
"inference_vmoba")
log "Running V-MoBA inference tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_vmoba"
;;
"unit_test")
log "Running unit tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_unit_test"
;;
"lora_extraction")
log "Running LoRA extraction tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_lora_extraction_tests"
;;
*)
log "Error: Unknown test type: $TEST_TYPE"
exit 1
;;
esac
log "Executing: $MODAL_COMMAND"
eval "$MODAL_COMMAND"
TEST_EXIT_CODE=$?
if [ $TEST_EXIT_CODE -eq 0 ]; then
log "Modal test completed successfully"
else
log "Error: Modal test failed with exit code: $TEST_EXIT_CODE"
fi
log "=== Test execution completed with exit code: $TEST_EXIT_CODE ==="
exit $TEST_EXIT_CODE
+40
View File
@@ -0,0 +1,40 @@
#!/bin/bash
set -uo pipefail
log() {
echo "[$(date '+%Y-%m-%d %H:%M:%S')] $1"
}
log "=== Starting pre-commit checks ==="
cd "$(dirname "$0")/../.."
PROJECT_ROOT=$(pwd)
log "Project root: $PROJECT_ROOT"
if ! python3 -m pre_commit --version &> /dev/null; then
log "pre-commit not found, installing..."
python3 -m pip install --user pre-commit==4.0.1
if ! python3 -m pre_commit --version &> /dev/null; then
log "Error: Failed to install pre-commit."
exit 1
fi
fi
log "Pre-commit version: $(python3 -m pre_commit --version)"
log "Installing/updating pre-commit hooks..."
python3 -m pre_commit install --install-hooks
log "Running pre-commit checks on all files..."
python3 -m pre_commit run --all-files
PRE_COMMIT_EXIT_CODE=$?
if [ $PRE_COMMIT_EXIT_CODE -eq 0 ]; then
log "Pre-commit checks completed successfully"
else
log "Error: Pre-commit checks failed with exit code: $PRE_COMMIT_EXIT_CODE"
fi
log "=== Pre-commit checks completed with exit code: $PRE_COMMIT_EXIT_CODE ==="
exit $PRE_COMMIT_EXIT_CODE
+1 -1
View File
@@ -23,7 +23,7 @@ body:
attributes:
label: Environment
description: |
Please share your environment with us. You can run the command **python fastvideo/utils/collect_env.py** and copy-paste its output below.
Please share your environment with us. You can run the command **python collect_env.py** and copy-paste its output below.
placeholder: FastVideo version, platform, python version, cuda version...
validations:
required: true
+56
View File
@@ -0,0 +1,56 @@
name: 💬 Request for comments (RFC).
description: Ask for feedback on major architectural changes or design choices.
title: "[RFC]: "
labels: ["RFC"]
body:
- type: markdown
attributes:
value: >
#### Please take a look at previous [RFCs](https://github.com/hao-ai-lab/FastVideo/issues?q=label%3ARFC+sort%3Aupdated-desc) for reference.
- type: textarea
attributes:
label: Motivation.
description: >
The motivation of the RFC.
validations:
required: true
- type: textarea
attributes:
label: Proposed Change.
description: >
The proposed change of the RFC.
validations:
required: true
- type: textarea
attributes:
label: Feedback Period.
description: >
The feedback period of the RFC. Usually at least one week.
validations:
required: false
- type: textarea
attributes:
label: CC List.
description: >
The list of people you want to CC.
validations:
required: false
- type: textarea
attributes:
label: Any Other Things.
description: >
Any other things you would like to mention.
validations:
required: false
- type: markdown
attributes:
value: >
Thanks for contributing 🎉!
- type: checkboxes
id: askllm
attributes:
label: Before submitting a new issue...
options:
- label: Make sure you already searched for relevant issues.
required: true
+1 -2
View File
@@ -160,8 +160,7 @@ def execute_command(pod_id):
setup_steps = [
"tar -xzf /tmp/repo.tar.gz --no-same-owner -C /workspace/",
f"cd /workspace/{repo_name}",
"source /opt/conda/etc/profile.d/conda.sh",
"conda activate fastvideo-dev",
"source $HOME/.local/bin/env && source /opt/venv/bin/activate",
args.test_command
]
+15
View File
@@ -18,6 +18,12 @@ on:
required: false
default: false
type: boolean
python_3_12_cuda_12_9:
description: 'Build Python 3.12 image Cuda 12.9'
required: false
default: false
type: boolean
permissions:
contents: read
@@ -49,4 +55,13 @@ jobs:
python_version: '3.12'
dockerfile_path: docker/Dockerfile.python3.12
tag_suffix: py3.12
secrets: inherit
build-python-3-12-cuda-12-9:
if: ${{ github.event.inputs.python_3_12_cuda_12_9 == 'true' }}
uses: ./.github/workflows/build-image-template.yml
with:
python_version: '3.12'
dockerfile_path: docker/Dockerfile.python3.12.cuda12.9.1
tag_suffix: py3.12-cuda12.9.1
secrets: inherit
+26 -43
View File
@@ -1,82 +1,65 @@
# Sample workflow for building and deploying a Hugo site to GitHub Pages
name: Deploy FastVideo Docs to Pages
name: Deploy Documentation
on:
# Runs on pushes targeting the default branch
push:
branches:
- main
branches: [ main ]
paths:
- "docs/**/*.md"
- "fastvideo/v1/examples/**/*.py"
- 'docs/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.txt'
- '.github/workflows/docs.yml'
pull_request:
branches:
- main
types: [opened, ready_for_review, synchronize, reopened]
branches: [ main ]
paths:
- "docs/**/*.md"
- "fastvideo/v1/examples/**/*.py"
- 'docs/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.txt'
- '.github/workflows/docs.yml'
# Allows you to run this workflow manually from the Actions tab
workflow_dispatch:
# Sets permissions of the GITHUB_TOKEN to allow deployment to GitHub Pages
permissions:
contents: read
pages: write
id-token: write
# Allow only one concurrent deployment, skipping runs queued between the run in-progress and latest queued.
# However, do NOT cancel in-progress runs as we want to allow these production deployments to complete.
concurrency:
group: "pages"
cancel-in-progress: false
# Default to bash
defaults:
run:
shell: bash
jobs:
pre-commit:
uses: ./.github/workflows/pre-commit.yml
# Build job
build:
runs-on: ubuntu-latest
needs: pre-commit
steps:
- name: Checkout
uses: actions/checkout@v4
- name: Setup Pages
id: pages
uses: actions/configure-pages@v5
- name: Set up Python
- name: Setup Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
python-version: '3.12'
- name: Install dependencies
run: |
cd docs
pip install -r requirements-docs.txt
- name: Build docs
run: |
cd docs
make clean
make html
python -m pip install --upgrade pip
pip install -r requirements-mkdocs.txt
- name: Setup Pages
uses: actions/configure-pages@v4
- name: Build documentation
run: mkdocs build
- name: Upload artifact
uses: actions/upload-pages-artifact@v3
with:
path: ./docs/build/html
path: ./site
# Deployment job
deploy:
environment:
name: github-pages
url: ${{ steps.deployment.outputs.page_url }}
if: ${{ github.event_name == 'push' }}
runs-on: ubuntu-latest
needs: build
if: github.ref == 'refs/heads/main'
steps:
- name: Deploy to GitHub Pages
id: deployment
@@ -0,0 +1,222 @@
name: Publish FastVideo Kernel to PyPI on Version Change
on:
push:
branches:
- main
paths:
- "fastvideo-kernel/pyproject.toml"
workflow_dispatch:
jobs:
check-version-change:
runs-on: ubuntu-latest
outputs:
version-changed: ${{ steps.check-version.outputs.changed }}
new-version: ${{ steps.check-version.outputs.new-version }}
steps:
- name: Checkout code
uses: actions/checkout@v4
with:
fetch-depth: 2
- name: Check if version changed
id: check-version
run: |
cd fastvideo-kernel
# Get current commit's version from pyproject.toml
# Use ^ to match start of line to avoid matching minimum-version
NEW_VERSION=$(grep -oP '^version\s*=\s*"\K[^"]+' pyproject.toml)
echo "New version: $NEW_VERSION"
# Get previous version from git history
# Note: git show expects path relative to repo root
OLD_VERSION=$(git show HEAD~1:fastvideo-kernel/pyproject.toml | grep -oP '^version\s*=\s*"\K[^"]+' || echo "0.0.0")
echo "Old version: $OLD_VERSION"
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
echo "changed=true" >> $GITHUB_OUTPUT
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
else
echo "Version did not change"
echo "changed=false" >> $GITHUB_OUTPUT
fi
build_wheels:
name: Build Wheel
needs: check-version-change
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
runs-on: ${{ matrix.os }}
strategy:
fail-fast: false
matrix:
os: [ubuntu-22.04]
python-version: ['3.10', '3.11', '3.12']
torch-cuda:
# - torch-version: '2.5.1'
# cuda-version: '12.4.1'
# torch-cuda-short: 'cu124'
# - torch-version: '2.6.0'
# cuda-version: '12.6.3'
# torch-cuda-short: 'cu126'
# - torch-version: '2.7.1'
# cuda-version: '12.8.0'
# torch-cuda-short: 'cu128'
- torch-version: '2.9.1'
cuda-version: '12.8.0'
torch-cuda-short: 'cu128'
steps:
- name: Free up disk space
run: |
echo "Initial disk space:"
df -h
# Remove large directories
sudo rm -rf /usr/share/dotnet
sudo rm -rf /usr/local/lib/android
sudo rm -rf /opt/ghc
sudo rm -rf /usr/local/share/boost
sudo rm -rf /usr/share/swift
sudo rm -rf /usr/local/lib/node_modules
sudo rm -rf /usr/local/share/powershell
sudo rm -rf /usr/share/rust
sudo rm -rf /usr/local/.ghcup
# Remove cached files
sudo rm -rf /var/lib/apt/lists/*
sudo rm -rf /var/cache/apt/archives/*
echo "Disk space after cleanup:"
df -h
- name: Checkout
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: ${{ matrix.python-version }}
- name: Install CUDA ${{ matrix.torch-cuda.cuda-version }}
uses: Jimver/cuda-toolkit@v0.2.21
id: cuda-toolkit
with:
cuda: ${{ matrix.torch-cuda.cuda-version }}
linux-local-args: '["--toolkit"]'
method: 'network'
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
run: |
sudo apt update
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Allow Git to Access Safe Directory
git config --global --add safe.directory /__w/FastVideo/FastVideo
# Set CUDA environment variables
export CUDA_HOME=/usr/local/cuda-${{ matrix.torch-cuda.cuda-version }}
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Verify installation
gcc --version
g++ --version
clang-11 --version
nvcc --version
- name: Install PyTorch ${{ matrix.torch-cuda.torch-version }}+cu${{ matrix.torch-cuda.cuda-version }}
run: |
pip install --upgrade pip
pip install typing-extensions==4.12.2
pip install --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
nvcc --version
python --version
python -c "import torch; print('PyTorch:', torch.__version__)"
python -c "import torch; print('CUDA:', torch.version.cuda)"
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
- name: Build wheel
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
pip install setuptools ninja packaging wheel triton scikit-build-core cmake build
cd fastvideo-kernel
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
# Release builds are produced on GPU-less runners, so force-enable TK and target Hopper.
export TORCH_CUDA_ARCH_LIST="9.0a"
export CMAKE_ARGS="${CMAKE_ARGS:-} -DFASTVIDEO_KERNEL_BUILD_TK=ON -DCMAKE_CUDA_ARCHITECTURES=90a"
# Build standard wheel (no local version suffix) for PyPI
python -m build --wheel --outdir dist
# Fix the wheel to be manylinux compliant
pip install auditwheel
# Point auditwheel at torch libs, but do not vendor them into the wheel.
TORCH_LIB_DIR=$(python - <<'PY'
import os
import torch
print(os.path.join(os.path.dirname(torch.__file__), "lib"))
PY
)
export LD_LIBRARY_PATH="${TORCH_LIB_DIR}:${LD_LIBRARY_PATH}"
# Target manylinux_2_35 (Ubuntu 22.04 native)
auditwheel repair dist/*.whl --plat manylinux_2_35_x86_64 -w fixed_dist \
--exclude libtorch_cuda.so \
--exclude libtorch_cpu.so \
--exclude libtorch.so \
--exclude libc10.so \
--exclude libc10_cuda.so \
--exclude libtorch_python.so
# Move fixed wheels back to dist for upload consistency
rm dist/*.whl
mv fixed_dist/*.whl dist/
- name: Upload wheel artifact
# Only upload if it's the "main" CUDA version we want on PyPI
# We upload all to artifacts for inspection/GH releases, but give them distinct artifact names
uses: actions/upload-artifact@v4
with:
name: fastvideo_kernel-py${{ matrix.python-version }}-${{ matrix.torch-cuda.torch-cuda-short }}-torch${{ matrix.torch-cuda.torch-version }}
path: fastvideo-kernel/dist/*.whl
retention-days: 90
publish_package:
name: Publish package
needs: [build_wheels, check-version-change]
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
runs-on: ubuntu-22.04
permissions:
id-token: write # Needed for OIDC Trusted Publishing
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: '3.10'
- name: Download PyPI wheels
uses: actions/download-artifact@v4
with:
path: fastvideo-kernel/dist/
pattern: 'fastvideo_kernel-py*'
merge-multiple: true
- name: Build source distribution
run: |
pip install build scikit-build-core cmake ninja
cd fastvideo-kernel
# We don't need full CUDA/Torch to just package the source (sdist)
python -m build --sdist --outdir dist
- name: Publish release distributions to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: fastvideo-kernel/dist/
+1 -1
View File
@@ -13,4 +13,4 @@
]
}
]
}
}
+243 -27
View File
@@ -12,13 +12,11 @@ on:
paths:
- "fastvideo/**/*.py"
- ".github/workflows/pr-test.yml"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
- "csrc/**"
workflow_dispatch:
inputs:
custom_image:
description: "Custom image from this repository (default: fastvideo-dev:latest)"
required: false
default: "fastvideo-dev:latest"
type: string
run_encoder_test:
description: "Run encoder-test"
required: false
@@ -39,10 +37,41 @@ on:
required: false
default: false
type: boolean
run_training_test:
description: "Run training-test"
required: false
default: false
type: boolean
run_training_test_VSA:
description: "Run training-test-VSA"
required: false
default: false
type: boolean
run_inference_test_STA:
description: "Run inference-test-STA"
required: false
default: false
type: boolean
run_precision_test_STA:
description: "Run precision-test-STA"
required: false
default: false
type: boolean
run_precision_test_VSA:
description: "Run precision-test-VSA"
required: false
default: false
type: boolean
run_unit_test:
description: "Run unit-test"
required: false
default: false
type: boolean
env:
PYTHONUNBUFFERED: "1"
concurrency:
group: pr-test-${{ github.ref }}
cancel-in-progress: true
@@ -59,26 +88,79 @@ jobs:
encoder-test: ${{ steps.filter.outputs.encoder-test }}
vae-test: ${{ steps.filter.outputs.vae-test }}
transformer-test: ${{ steps.filter.outputs.transformer-test }}
training-test: ${{ steps.filter.outputs.training-test }}
training-test-VSA: ${{ steps.filter.outputs.training-test-VSA }}
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
unit-test: ${{ steps.filter.outputs.unit-test }}
steps:
- uses: actions/checkout@v4
- uses: dorny/paths-filter@v3
id: filter
with:
filters: |
# Define reusable path patterns
common-paths: &common-paths
- 'pyproject.toml'
- 'docker/Dockerfile.python3.10'
- 'docker/Dockerfile.python3.11'
- 'docker/Dockerfile.python3.12'
sta-kernel-paths: &sta-kernel-paths
- 'csrc/attn/sliding_tile_attn/**'
- 'csrc/attn/sliding_tile_attn/tk/**'
- 'csrc/attn/sliding_tile_attn/setup.py'
- 'csrc/attn/sliding_tile_attn/config_sta.py'
- 'csrc/attn/sliding_tile_attn/st_attn.cpp'
vsa-kernel-paths: &vsa-kernel-paths
- 'csrc/attn/video_sparse_attn/**'
- 'csrc/attn/video_sparse_attn/tk/**'
- 'csrc/attn/video_sparse_attn/setup.py'
- 'csrc/attn/video_sparse_attn/config_vsa.py'
- 'csrc/attn/video_sparse_attn/vsa.cpp'
vsa-paths: &vsa-paths
- 'fastvideo/**'
- *common-paths
- *vsa-kernel-paths
# Actual tests
encoder-test:
- 'fastvideo/v1/models/encoders/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/encoders/**'
- 'fastvideo/models/encoders/**'
- 'fastvideo/models/loader/**'
- 'fastvideo/tests/encoders/**'
- *common-paths
vae-test:
- 'fastvideo/v1/models/vaes/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/vaes/**'
- 'fastvideo/models/vaes/**'
- 'fastvideo/models/loader/**'
- 'fastvideo/tests/vaes/**'
- *common-paths
transformer-test:
- 'fastvideo/v1/models/dits/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/transformers/**'
- 'fastvideo/v1/layers/**'
- 'fastvideo/v1/attention/**'
- 'fastvideo/models/dits/**'
- 'fastvideo/models/loader/**'
- 'fastvideo/tests/transformers/**'
- 'fastvideo/layers/**'
- 'fastvideo/attention/**'
- *common-paths
training-test:
- 'fastvideo/**'
- *common-paths
training-test-VSA:
- 'fastvideo/**'
- *common-paths
- *vsa-kernel-paths
inference-test-STA:
- 'fastvideo/**'
- *common-paths
- *sta-kernel-paths
precision-test-STA:
- *common-paths
- *sta-kernel-paths
precision-test-VSA:
- *common-paths
- *vsa-kernel-paths
unit-test:
- 'fastvideo/**'
- *common-paths
encoder-test:
needs: change-filter
@@ -91,8 +173,8 @@ jobs:
gpu_type: "NVIDIA A40"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/encoders -s"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
@@ -109,8 +191,8 @@ jobs:
gpu_type: "NVIDIA A40"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/vaes -s"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
@@ -127,8 +209,8 @@ jobs:
gpu_type: "NVIDIA L40S"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/transformers -s"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
@@ -137,8 +219,7 @@ jobs:
ssim-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
github.event_name != 'workflow_dispatch' || (github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
strategy:
fail-fast: false
matrix:
@@ -155,14 +236,149 @@ jobs:
volume_size: 200
disk_size: 200
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:${{ matrix.python-version.tag }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/ssim -vs"
timeout_minutes: 60
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
training-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.training-test == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "training-test"
gpu_type: "NVIDIA A40"
gpu_count: 4
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/training/Vanilla -srP"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
training-test-VSA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.training-test-VSA == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test_VSA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "training-test-VSA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 2
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/training/VSA -srP"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
inference-test-STA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.inference-test-STA == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_inference_test_STA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "inference-test-STA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 2
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/inference/STA -srP"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
precision-test-STA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-STA == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_STA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "precision-test-STA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_sta.py"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
precision-test-VSA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-VSA == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_VSA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "precision-test-VSA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_vsa.py"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
unit-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.unit-test == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_unit_test == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "unit-test"
gpu_type: "NVIDIA L40S"
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/dataset/ -vs && pytest ./fastvideo/workflow/ -vs && pytest ./fastvideo/entrypoints/ -vs"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
# nightly-test:
# if: >-
# (github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
# uses: ./.github/workflows/runpod-test.yml
# with:
# job_id: "nightly-test"
# gpu_type: "NVIDIA A40"
# gpu_count: 4
# volume_size: 100
# disk_size: 100
# image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
# test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
# timeout_minutes: 30
# secrets:
# RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
# RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
# WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
runpod-cleanup:
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
# Add other jobs to this list as you create them
needs: [encoder-test, vae-test, transformer-test, ssim-test, training-test, training-test-VSA, inference-test-STA, precision-test-STA, precision-test-VSA]
if: ${{ always() && ((github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) || github.event_name == 'workflow_dispatch') }}
runs-on: ubuntu-latest
steps:
@@ -179,7 +395,7 @@ jobs:
- name: Cleanup all RunPod instances
env:
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12"]'
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12", "training-test", "training-test-VSA", "inference-test-STA", "precision-test-STA", "precision-test-VSA"]'
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
run: python .github/scripts/runpod_cleanup.py
+28
View File
@@ -0,0 +1,28 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
- master
paths:
- "pyproject.toml"
permissions:
issues: write
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'hao-ai-lab' }}
steps:
- name: Check out code
uses: actions/checkout@v4
with:
submodules: true
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@v1
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+4 -1
View File
@@ -43,6 +43,8 @@ on:
required: true
RUNPOD_PRIVATE_KEY:
required: true
WANDB_API_KEY:
required: false
jobs:
run-test:
@@ -55,7 +57,7 @@ jobs:
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
python-version: "3.12"
- name: Set up SSH key
run: |
@@ -72,6 +74,7 @@ jobs:
JOB_ID: ${{ inputs.job_id }}
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
timeout-minutes: ${{ inputs.timeout_minutes }}
run: >-
python .github/scripts/runpod_api.py
+13 -9
View File
@@ -5,7 +5,7 @@ on:
branches:
- main
paths:
- "csrc/sliding_tile_attention/setup.py"
- "csrc/attn/sliding_tile_attn/setup.py"
workflow_dispatch:
jobs:
@@ -23,7 +23,7 @@ jobs:
- name: Check if version changed
id: check-version
run: |
cd csrc/sliding_tile_attention
cd csrc/attn/sliding_tile_attn
# Get current commit's version
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
echo "New version: $NEW_VERSION"
@@ -136,19 +136,21 @@ jobs:
- name: Build wheel
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
# However this still fails so I'm using a newer version of setuptools
pip install setuptools
pip install ninja packaging wheel
cd csrc/sliding_tile_attention # Move into the correct folder
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
cd csrc/attn/sliding_tile_attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup.py bdist_wheel --dist-dir=dist
- name: Rename wheel file
run: |
cd csrc/sliding_tile_attention
cd csrc/attn/sliding_tile_attn
CUDA_SHORT_VERSION=$(echo ${{ matrix.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-version }} | cut -d. -f1,2)
@@ -163,7 +165,7 @@ jobs:
uses: actions/upload-artifact@v4
with:
name: ${{ env.wheel_name }}
path: csrc/sliding_tile_attention/dist/*.whl
path: csrc/attn/sliding_tile_attn/dist/*.whl
retention-days: 90
publish_package:
@@ -229,17 +231,19 @@ jobs:
- name: Build source distribution
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
# However this still fails so I'm using a newer version of setuptools
pip install setuptools
pip install ninja packaging wheel
cd csrc/sliding_tile_attention # Move into the correct folder
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
cd csrc/attn/sliding_tile_attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup.py sdist --dist-dir=dist
- name: Publish release distributions to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: csrc/sliding_tile_attention/dist/
packages-dir: csrc/attn/sliding_tile_attn/dist/
+1 -1
View File
@@ -28,4 +28,4 @@ jobs:
- name: Run Pytest
run: |
pytest --ignore csrc/sliding_tile_attention/test
pytest --ignore csrc/attn/test
+257
View File
@@ -0,0 +1,257 @@
name: Publish Video Sparse Attention Kernel to PyPI on Version Change
on:
push:
branches:
- main
paths:
- "csrc/attn/video_sparse_attn/setup.py"
workflow_dispatch:
jobs:
check-version-change:
runs-on: ubuntu-latest
outputs:
version-changed: ${{ steps.check-version.outputs.changed }}
new-version: ${{ steps.check-version.outputs.new-version }}
steps:
- name: Checkout code
uses: actions/checkout@v4
with:
fetch-depth: 2
- name: Check if version changed
id: check-version
run: |
cd csrc/attn/video_sparse_attn
# Get current commit's version
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
echo "New version: $NEW_VERSION"
# Get previous version from git history
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
echo "Old version: $OLD_VERSION"
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
echo "changed=true" >> $GITHUB_OUTPUT
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
else
echo "Version did not change"
echo "changed=false" >> $GITHUB_OUTPUT
fi
build_wheels:
name: Build Wheel
needs: check-version-change
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
runs-on: ${{ matrix.os }}
strategy:
fail-fast: false
matrix:
# Using ubuntu-20.04 instead of 22.04 for more compatibility (glibc). Ideally we'd use the
# manylinux docker image, but I haven't figured out how to install CUDA on manylinux.
os: [ubuntu-22.04]
python-version: ['3.10', '3.11', '3.12', '3.13']
# For version reference https://pytorch.org/get-started/previous-versions/
torch-cuda:
- torch-version: '2.5.1'
cuda-version: '12.4.1'
torch-cuda-short: 'cu124'
- torch-version: '2.6.0'
cuda-version: '12.6.3'
torch-cuda-short: 'cu126'
- torch-version: '2.7.1'
cuda-version: '12.8.0'
torch-cuda-short: 'cu128'
steps:
- name: Free up disk space
run: |
echo "Initial disk space:"
df -h
# Remove large directories
sudo rm -rf /usr/share/dotnet
sudo rm -rf /usr/local/lib/android
sudo rm -rf /opt/ghc
sudo rm -rf /usr/local/share/boost
sudo rm -rf /usr/share/swift
sudo rm -rf /usr/local/lib/node_modules
sudo rm -rf /usr/local/share/powershell
sudo rm -rf /usr/share/rust
sudo rm -rf /usr/local/.ghcup
# Remove cached files
sudo rm -rf /var/lib/apt/lists/*
sudo rm -rf /var/cache/apt/archives/*
echo "Disk space after cleanup:"
df -h
- name: Checkout
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: ${{ matrix.python-version }}
- name: Install CUDA ${{ matrix.torch-cuda.cuda-version }}
uses: Jimver/cuda-toolkit@v0.2.21
id: cuda-toolkit
with:
cuda: ${{ matrix.torch-cuda.cuda-version }}
linux-local-args: '["--toolkit"]'
method: 'network'
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
run: |
sudo apt update
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Allow Git to Access Safe Directory
git config --global --add safe.directory /__w/FastVideo/FastVideo
# Set CUDA environment variables
export CUDA_HOME=/usr/local/cuda-${{ matrix.torch-cuda.cuda-version }}
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Verify installation
gcc --version
g++ --version
clang-11 --version
nvcc --version
- name: Install PyTorch ${{ matrix.torch-cuda.torch-version }}+cu${{ matrix.torch-cuda.cuda-version }}
run: |
pip install --upgrade pip
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
pip install typing-extensions==4.12.2
# We want to figure out the CUDA version to download pytorch
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
pip install --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
nvcc --version
python --version
python -c "import torch; print('PyTorch:', torch.__version__)"
python -c "import torch; print('CUDA:', torch.version.cuda)"
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
- name: Build wheel
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
# However this still fails so I'm using a newer version of setuptools
pip install setuptools
pip install ninja packaging wheel
cd csrc/attn/video_sparse_attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup.py bdist_wheel --dist-dir=dist
- name: Rename wheel file
run: |
cd csrc/attn/video_sparse_attn
CUDA_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.torch-version }} | cut -d. -f1,2)
# Get the correct version format
tmpname=cu${CUDA_SHORT_VERSION}torch${TORCH_SHORT_VERSION}
wheel_name=$(ls dist/*whl | xargs -n 1 basename | sed "s/-/+$tmpname-/2")
# Rename with version information
ls dist/*whl |xargs -I {} mv {} dist/${wheel_name}
echo "wheel_name=${wheel_name}" >> $GITHUB_ENV
- name: Upload wheel artifact
uses: actions/upload-artifact@v4
with:
name: ${{ env.wheel_name }}
path: csrc/attn/video_sparse_attn/dist/*.whl
retention-days: 90
publish_package:
name: Publish package
needs: [build_wheels, check-version-change]
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
runs-on: ubuntu-22.04
permissions:
id-token: write # Needed for OIDC Trusted Publishing
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: '3.10'
- name: Install CUDA 12.4.1
uses: Jimver/cuda-toolkit@v0.2.21
id: cuda-toolkit
with:
cuda: 12.4.1
linux-local-args: '["--toolkit"]'
method: 'network'
sub-packages: '["nvcc"]'
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
run: |
sudo apt update
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Allow Git to Access Safe Directory
git config --global --add safe.directory /__w/FastVideo/FastVideo
# Set CUDA environment variables
export CUDA_HOME=/usr/local/cuda-12.4.1
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Verify installation
gcc --version
g++ --version
clang-11 --version
nvcc --version
- name: Install PyTorch 2.5.1+cu12.4.1
run: |
pip install --upgrade pip
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
pip install typing-extensions==4.12.2
# We want to figure out the CUDA version to download pytorch
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
export TORCH_CUDA_VERSION=124
pip install --no-cache-dir torch==2.5.1 --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
nvcc --version
python --version
python -c "import torch; print('PyTorch:', torch.__version__)"
python -c "import torch; print('CUDA:', torch.version.cuda)"
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
- name: Build source distribution
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
# However this still fails so I'm using a newer version of setuptools
pip install setuptools
pip install ninja packaging wheel
cd csrc/attn/video_sparse_attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup.py sdist --dist-dir=dist
- name: Publish release distributions to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: csrc/attn/video_sparse_attn/dist/
+29 -6
View File
@@ -14,12 +14,15 @@ wandb/
*.pt
cache_dir/
wandb/
venv/
.venv/
runs/
samples/
*validation/
data/
outputs/
outputs_video
checkpoints/
sbatch.sh
*.out
env
@@ -27,6 +30,13 @@ env
**/build/
**.pyc
**.txt
*.log
weights/
official_weights/
converted_weights/
# SSIM test outputs
fastvideo/tests/ssim/generated_videos/
# Distribution / packaging
build/
@@ -36,10 +46,13 @@ dist/
eggs/
.eggs/
# Sphinx documentation
docs/_build/
docs/source/getting_started/examples/
docs/source/inference/examples/
# MkDocs documentation
site/
docs/getting_started/examples/
docs/inference/examples/
docs/training/examples/
docs/distillation/examples/
!requirements-mkdocs.txt
# VSCode
.vscode/
@@ -55,7 +68,17 @@ docs/source/inference/examples/
*.pkl
# Reference videos
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
!fastvideo/tests/ssim/reference_videos/**/*.mp4
# Static images
!docs/source/_static/images/**/*.png
!docs/assets/images/**/*.png
!comfyui/assets/**/*.png
!comfyui/assets/**/*.gif
!assets/images/**/*.png
!assets/images/**/*.jpg
!assets/images/**/*.jpeg
!assets/images/**/*.gif
!assets/videos/**/*.mp4
dmd_t2v_output/
preprocess_output_text/
+5 -2
View File
@@ -1,3 +1,6 @@
[submodule "csrc/sliding_tile_attention/tk"]
path = csrc/sliding_tile_attention/tk
[submodule "fastvideo-kernel/include/tk"]
path = fastvideo-kernel/include/tk
url = https://github.com/HazyResearch/ThunderKittens.git
[submodule "fastvideo-kernel/include/cutlass"]
path = fastvideo-kernel/include/cutlass
url = https://github.com/NVIDIA/cutlass.git
+10 -11
View File
@@ -3,18 +3,16 @@ default_stages:
- manual # Run in CI
exclude: |
(?x)(
fastvideo/v1/third_party/.*|
csrc/.*|
fastvideo/third_party/.*|
fastvideo-kernel/.*|
assets/.*|
tests/.*|
demo/.*|
predict\.py|
scripts/.*|
assets/prompts/.*|
fastvideo/data_preprocess/.*|
fastvideo/dataset/.*|
fastvideo/distill/.*|
fastvideo/distill\.py|
fastvideo/distill_adv\.py|
fastvideo/models/.*|
fastvideo/sample/.*|
fastvideo/train\.py|
@@ -22,6 +20,7 @@ exclude: |
examples/.*|
.github/workflows/fastvideo-publish.yml|
.github/workflows/sta-publish.yml|
.github/workflows/vsa-publish.yml|
.github/workflows/build-image-template.yml|
docs/source/inference/support_matrix.md
)
@@ -43,10 +42,10 @@ repos:
- id: codespell
additional_dependencies: ['tomli']
args: ['--toml', 'pyproject.toml']
- repo: https://github.com/PyCQA/isort
rev: 6.0.1
hooks:
- id: isort
# - repo: https://github.com/PyCQA/isort
# rev: 6.0.1
# hooks:
# - id: isort
- repo: https://github.com/jackdewinter/pymarkdown
rev: v0.9.30
hooks:
@@ -60,7 +59,7 @@ repos:
rev: v1.15.0
hooks:
- id: mypy
args: [--python-version, '3.10', --follow-imports, "skip", ]
args: [--python-version, '3.10', --follow-imports, "skip", "--disable-error-code", "union-attr", "--disable-error-code", "override" ]
additional_dependencies: [types-cachetools, types-setuptools, types-PyYAML, types-requests]
- repo: local
hooks:
@@ -69,7 +68,7 @@ repos:
entry: bash
args:
- -c
- 'git ls-files | grep -v "^fastvideo/v1/tests/ssim/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
- 'git ls-files | grep -v "^\"*fastvideo/tests/ssim/" | grep -v "^\"*fastvideo/tests/inference/lora/L40S_reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
language: system
always_run: true
pass_filenames: false
+42
View File
@@ -0,0 +1,42 @@
# Repository Guidelines
## Project Structure & Module Organization
- Core Python package: `fastvideo/` (models, pipelines, training, distributed runtime, CLI entrypoints).
- CUDA/custom kernels: `fastvideo-kernel/` (separate build/test flow).
- Tests:
- `fastvideo/tests/` for package-level tests (dataset, encoders, inference, training, SSIM, workflow).
- `tests/local_tests/` for additional local/component checks.
- Docs and guides: `docs/` (MkDocs source), with contributor docs in `docs/contributing/`.
- Runnable examples and scripts: `examples/` and `scripts/`.
- Static assets: `assets/` (including `assets/images/`, `assets/videos/`, and `assets/prompts/`) and `comfyui/assets/`.
## Build, Test, and Development Commands
- `uv pip install -e .[dev]`: editable install with lint/test extras.
- `pre-commit install --hook-type pre-commit --hook-type commit-msg`: enable local hooks.
- `pre-commit run --all-files`: run formatter/lint/type/spelling checks.
- `pytest tests/`: run top-level test suite.
- `pytest fastvideo/tests/ -v`: run package tests.
- `pytest fastvideo/tests/ssim/ -vs`: run SSIM regression tests (GPU-heavy).
- `cd fastvideo-kernel && ./build.sh`: build kernel extensions.
## Coding Style & Naming Conventions
- Python 3.10+; 4-space indentation; keep code and imports readable and explicit.
- Style tools are configured in `pyproject.toml` and `.pre-commit-config.yaml`:
- `yapf` (format), `ruff` (lint, auto-fix), `mypy` (typing), `codespell`.
- Target line length is 80.
- Naming: `snake_case` for functions/files, `PascalCase` for classes, `UPPER_SNAKE_CASE` for constants.
## Testing Guidelines
- Use `pytest` and place tests near relevant domains (e.g., `fastvideo/tests/encoders/`).
- Prefer descriptive names like `test_<feature>_<expected_behavior>.py`.
- For new pipelines/backends, include at least one regression-oriented test; add SSIM coverage when output quality must be preserved.
- Document GPU assumptions in tests that require specific hardware.
## Commit & Pull Request Guidelines
- Follow existing commit style: short subject with optional tag prefix, e.g. `[bugfix]: ...`, `[feat]: ...`, `[misc]: ...`, and include PR reference like `(#1234)` when applicable.
- Keep commits focused by concern (feature, refactor, fix).
- PRs should include:
- clear problem/solution summary,
- test evidence (`pytest`/SSIM outputs or rationale if skipped),
- linked issue/PR context,
- screenshots or sample outputs for UI/demo/docs changes.
+1
View File
@@ -0,0 +1 @@
@AGENTS.md
+81 -74
View File
@@ -1,39 +1,49 @@
<div align="center">
<img src=assets/logo.jpg width="30%"/>
<img src=assets/logos/logo.svg width="30%"/>
</div>
**FastVideo is a unified framework for accelerated video generation.**
It features a clean, consistent API that works across popular video models, making it easier for developers to author new models and incorporate system- or kernel-level optimizations.
With FastVideo's optimizations, you can achieve more than 3x inference improvement compared to other systems.
<p align="center">
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> <b>Slack</b> </a> |
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | <a href="https://github.com/hao-ai-lab/FastVideo/discussions/982" target="_blank"><b>Weekly Dev Meeting</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://github.com/hao-ai-lab/FastVideo/discussions/1097" target="_blank"> <b> WeChat </b> </a> |
</p>
<div align="center">
<img src=assets/perf.png width="90%"/>
</div>
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
## NEWS
- `2025/11/19`: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
- `2025/08/04`: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
### More News
- `2025/06/14`: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
- `2025/04/24`: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
- `2025/02/18`: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
## Key Features
FastVideo has the following features:
- End-to-end post-training support for bidirectional and autoregressive models:
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs
- Data preprocessing pipeline for video, image, and text data
- Distribution Matching Distillation (DMD2) stepwise distillation.
- Sparse attention with [Video Sparse Attention](https://arxiv.org/pdf/2505.13389)
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) to achieve >50x denoising speedup
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing.
- Causal distillation through Self-Forcing
- See this [page](https://hao-ai-lab.github.io/FastVideo/training/overview/) for full list of supported models and recipes.
- State-of-the-art performance optimizations for inference
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
- [TeaCache](https://arxiv.org/pdf/2411.19108)
- [Sage Attention](https://arxiv.org/abs/2410.02367)
- Cutting edge models
- Wan2.1 T2V, I2V
- HunyuanVideo
- FastHunyuan: consistency distilled video diffusion models for 8x inference speedup.
- StepVideo T2V
- Distillation support
- Recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
- Scalable training with FSDP, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
- Memory efficient finetuning with LoRA, precomputed latent, and precomputed text embeddings.
- Sequence Parallelism for distributed inference
- Multiple state-of-the-art attention backends
- User-friendly CLI and Python API
- See this [page](https://hao-ai-lab.github.io/FastVideo/inference/optimizations/) for full list of supported optimizations.
- Diverse hardware and OS support
- Support H100, A100, 4090
- Support Linux, Windows, MacOS
- See this [page](https://hao-ai-lab.github.io/FastVideo/inference/hardware_support/) for full list of supported hardware and OS.
## Getting Started
We recommend using an environment manager such as `Conda` to create a clean environment:
```bash
@@ -45,19 +55,35 @@ conda activate fastvideo
pip install fastvideo
```
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) for more detailed installation instructions.
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/) for more detailed installation instructions.
## Sparse Distillation
For our sparse distillation techniques, please see our [distillation docs](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/) and check out our [blog](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
See below for recipes and datasets:
| Model | Sparse Distillation | Dataset |
| ------------------------------------------------------------------------------------- | --------------------------------------------------------------------------------------------------------------- | -------------------------------------------------------------------------------------------------------- |
| [FastWan2.1-T2V-1.3B](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P) | [FastVideo Synthetic Wan2.1 480P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x448x832_600k) |
| [FastWan2.2-TI2V-5B](https://huggingface.co/FastVideo/FastWan2.2-TI2V-5B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free) | [FastVideo Synthetic Wan2.2 720P](https://huggingface.co/datasets/FastVideo/Wan2.2-Syn-121x704x1280_32k) |
## Inference
### Generating Your First Video
Here's a minimal example to generate a video using the default settings. Create a file called `example.py` with the following code:
Here's a minimal example to generate a video using the default settings. Make sure VSA kernels are [installed](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation/). Create a file called `example.py` with the following code:
```python
import os
from fastvideo import VideoGenerator
def main():
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
# Create a video generator with a pre-trained model
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
num_gpus=1, # Adjust based on your hardware
)
@@ -82,68 +108,49 @@ Run the script with:
python example.py
```
For a more detailed guide, please see our [inference quick start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html).
For a more detailed guide, please see our [inference quick start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/).
### Other docs:
## More Guides
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview.html)
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html)
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview/)
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/)
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/contributing/overview/)
## Distillation and Finetuning
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/training/distillation.html)
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetuning.html)
## Awesome work using FastVideo or our research projects
## 📑 Development Plan
<!-- - More distillation methods -->
<!-- - [ ] Add Distribution Matching Distillation -->
- More models support
<!-- - [ ] Add CogvideoX model -->
- [x] Add StepVideo to V1
- Optimization features
- [x] Teacache in V1
- [x] SageAttention in V1
- Code updates
- [x] V1 Configuration API
- [ ] Support Training in V1
<!-- - [ ] fp8 support -->
<!-- - [ ] faster load model and save model support -->
- [SGLang](https://github.com/sgl-project/sglang/tree/main/python/sglang/multimodal_gen): SGLang's diffusion inference functionality is based on a fork of FastVideo on Sept. 24, 2025.
- [DanceGRPO](https://github.com/XueZeyue/DanceGRPO): A unified framework to adapt Group Relative Policy Optimization (GRPO) to visual generation paradigms. Code based on FastVideo.
- [SRPO](https://github.com/Tencent-Hunyuan/SRPO): A method to directly align the full diffusion trajectory with fine-grained human preference. Code based on FastVideo.
- [DCM](https://github.com/Vchitect/DCM): Dual-expert consistency model for efficient and high-quality video generation. Code based on FastVideo.
- [Hunyuan Video 1.5](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5): A leading lightweight video generation model, where they proposed SSTA based on Sliding Tile Attention.
- [Kandinsky-5.0](https://github.com/kandinskylab/kandinsky-5): A family of diffusion models for video & image generation, where their NABLA attention includes a Sliding Tile Attention branch.
- [LongCat Video](https://github.com/meituan-longcat/LongCat-Video): A foundational video generation model with 13.6B parameters with block-sparse attention similar to Video Sparse Attention.
## 🤝 Contributing
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/developer_guide/overview.html)
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview/).
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/899).
## Acknowledgement
We learned and reused code from the following projects:
- [PCM](https://github.com/G-U-N/Phased-Consistency-Model)
- [diffusers](https://github.com/huggingface/diffusers)
- [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan)
- [xDiT](https://github.com/xdit-project/xDiT)
- [vLLM](https://github.com/vllm-project/vllm)
- [SGLang](https://github.com/sgl-project/sglang)
We thank MBZUAI and [Anyscale](https://www.anyscale.com/) for their support throughout this project.
We learned the design and reused code from the following projects: [Wan-Video](https://github.com/Wan-Video), [ThunderKittens](https://github.com/HazyResearch/ThunderKittens), [DMD2](https://github.com/tianweiy/DMD2), [diffusers](https://github.com/huggingface/diffusers), [xDiT](https://github.com/xdit-project/xDiT), [vLLM](https://github.com/vllm-project/vllm), [SGLang](https://github.com/sgl-project/sglang). We thank [MBZUAI](https://ifm.mbzuai.ac.ae/), [Anyscale](https://www.anyscale.com/), and [GMI Cloud](https://www.gmicloud.ai/) for their support throughout this project.
## Citation
If you use FastVideo for your research, please cite our paper:
If you find FastVideo useful, please consider citing our research work:
```bibtex
@misc{zhang2025fastvideogenerationsliding,
title={Fast Video Generation with Sliding Tile Attention},
author={Peiyuan Zhang and Yongqi Chen and Runlong Su and Hangliang Ding and Ion Stoica and Zhenghong Liu and Hao Zhang},
year={2025},
eprint={2502.04507},
archivePrefix={arXiv},
primaryClass={cs.CV},
url={https://arxiv.org/abs/2502.04507},
@article{zhang2025vsa,
title={Vsa: Faster video diffusion with trainable sparse attention},
author={Zhang, Peiyuan and Chen, Yongqi and Huang, Haofeng and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
journal={arXiv preprint arXiv:2505.13389},
year={2025}
}
@misc{ding2025efficientvditefficientvideodiffusion,
title={Efficient-vDiT: Efficient Video Diffusion Transformers With Attention Tile},
author={Hangliang Ding and Dacheng Li and Runlong Su and Peiyuan Zhang and Zhijie Deng and Ion Stoica and Hao Zhang},
year={2025},
eprint={2502.06155},
archivePrefix={arXiv},
primaryClass={cs.CV},
url={https://arxiv.org/abs/2502.06155},
@article{zhang2025fast,
title={Fast video generation with sliding tile attention},
author={Zhang, Peiyuan and Chen, Yongqi and Su, Runlong and Ding, Hangliang and Stoica, Ion and Liu, Zhengzhong and Zhang, Hao},
journal={arXiv preprint arXiv:2502.04507},
year={2025}
}
```
+15
View File
@@ -0,0 +1,15 @@
try:
from .comfyui.video_generator.nodes import (NODE_CLASS_MAPPINGS,
NODE_DISPLAY_NAME_MAPPINGS)
WEB_DIRECTORY = "./web"
__all__ = [
'NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY'
]
except ImportError:
# ComfyUI environment not available, skip comfyui imports
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
WEB_DIRECTORY = "./web"
__all__ = [
'NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY'
]
Binary file not shown.

After

Width:  |  Height:  |  Size: 194 KiB

+18
View File
@@ -0,0 +1,18 @@
<svg width="252" height="105" viewBox="0 0 252 105" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM100.768 13.1217L87.7028 29.4852H103.143L100.768 13.1217Z" fill="#356CFF"/>
<path d="M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM109.081 90.697L116.802 65.8487C116.802 65.8487 120.959 65.8487 132.242 65.8487C143.525 65.8487 137.586 78.5759 135.211 84.0304C133.307 88.4021 127.491 90.697 122.74 90.697C117.989 90.697 109.081 90.697 109.081 90.697Z" fill="#356CFF"/>
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944H159.747C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273L125.188 48.273L124 37.97L147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852H131.836C120.142 29.4852 125.897 1.00043 141.337 1.00043L173.188 1.00056Z" fill="#356CFF"/>
<path d="M179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056Z" fill="#356CFF"/>
<path d="M161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM237.948 77.9692C239.984 70.6965 240.917 65.242 228.446 65.242C215.975 65.242 211.818 71.9087 210.037 77.9692C208.255 84.0298 208.255 91.3025 219.538 91.3025C230.821 91.3025 235.911 85.2419 237.948 77.9692Z" fill="#356CFF"/>
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944M173.188 1.00056C173.188 1.00056 156.777 1.00043 141.337 1.00043M173.188 1.00056L141.337 1.00043M141.337 20.3944C146.088 20.3944 150.839 20.3944 159.747 20.3944M141.337 20.3944H159.747M159.747 20.3944C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273M148.463 48.273C139.556 48.273 125.188 48.273 125.188 48.273M148.463 48.273L125.188 48.273M125.188 48.273L124 37.97M124 37.97C124 37.97 141.931 37.97 147.87 37.97M124 37.97L147.87 37.97M147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852M151.433 29.4852C146.682 29.4852 138.962 29.4852 131.836 29.4852M151.433 29.4852H131.836M131.836 29.4852C120.142 29.4852 125.897 1.00043 141.337 1.00043M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057ZM96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM87.7028 29.4852L100.768 13.1217L103.143 29.4852H87.7028ZM89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457ZM108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM116.802 65.8487L109.081 90.697C109.081 90.697 117.989 90.697 122.74 90.697C127.491 90.697 133.307 88.4021 135.211 84.0304C137.586 78.5759 143.525 65.8487 132.242 65.8487C120.959 65.8487 116.802 65.8487 116.802 65.8487ZM179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056ZM161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457ZM230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM228.446 65.242C240.917 65.242 239.984 70.6965 237.948 77.9692C235.911 85.2419 230.821 91.3025 219.538 91.3025C208.255 91.3025 208.255 84.0298 210.037 77.9692C211.818 71.9087 215.975 65.242 228.446 65.242Z" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M15.2524 55.5451L21.191 100.999L24.7541 100.999L18.8156 55.5451L15.2524 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M8.12646 55.5451L14.065 100.999L15.2527 100.999L9.31417 55.5451L8.12646 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M1 55.5451L6.93853 100.999L7.53239 100.999L1.59385 55.5451L1 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
<path d="M15.2524 48.2724L30.0988 1H33.6619L18.8156 48.2724H15.2524Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M8.12646 48.2724L22.9728 1H24.1605L9.31417 48.2724H8.12646Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M1 48.2724L15.8463 1H16.4402L1.59385 48.2724H1Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
<path d="M85.3271 55.5457H67.5116L87 12.7363L44.3513 68.2729H58.6038L43.1636 101L85.3271 55.5457Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.18771" stroke-miterlimit="16"/>
</svg>

After

Width:  |  Height:  |  Size: 5.7 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 490 KiB

+6
View File
@@ -0,0 +1,6 @@
<svg width="160" height="93" viewBox="0 0 160 93" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M28.8511 91.66L57.6319 1.86368H64.5394L35.7585 91.66H28.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M15.0376 91.66L43.8185 1.86368H46.1209L17.3401 91.66H15.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M1.22217 91.66L30.003 1.86366H31.1543L2.3734 91.66H1.22217Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
<path d="M71.4465 1.86483L42.666 91.6599H69.144L78.3538 58.2746H123.251L129.007 39.855H84.1099L89.866 22.5868H152.032L157.788 1.86483H71.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
</svg>

After

Width:  |  Height:  |  Size: 691 B

Binary file not shown.

After

Width:  |  Height:  |  Size: 113 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.2 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 229 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 168 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 103 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 148 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 155 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 723 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 723 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 875 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 664 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 62 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 686 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 957 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 585 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 558 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 942 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 890 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 433 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 595 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 781 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 783 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 762 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 68 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 147 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 89 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 133 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 213 KiB

BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 149 KiB

+6
View File
@@ -0,0 +1,6 @@
<svg width="160" height="93" viewBox="0 0 160 93" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M28.8511 91.66L57.6319 1.86368H64.5394L35.7585 91.66H28.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M15.0376 91.66L43.8185 1.86368H46.1209L17.3401 91.66H15.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M1.22217 91.66L30.003 1.86366H31.1543L2.3734 91.66H1.22217Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
<path d="M71.4465 1.86483L42.666 91.6599H69.144L78.3538 58.2746H123.251L129.007 39.855H84.1099L89.866 22.5868H152.032L157.788 1.86483H71.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
</svg>

After

Width:  |  Height:  |  Size: 691 B

+18
View File
@@ -0,0 +1,18 @@
<svg width="252" height="105" viewBox="0 0 252 105" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM100.768 13.1217L87.7028 29.4852H103.143L100.768 13.1217Z" fill="#356CFF"/>
<path d="M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM109.081 90.697L116.802 65.8487C116.802 65.8487 120.959 65.8487 132.242 65.8487C143.525 65.8487 137.586 78.5759 135.211 84.0304C133.307 88.4021 127.491 90.697 122.74 90.697C117.989 90.697 109.081 90.697 109.081 90.697Z" fill="#356CFF"/>
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944H159.747C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273L125.188 48.273L124 37.97L147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852H131.836C120.142 29.4852 125.897 1.00043 141.337 1.00043L173.188 1.00056Z" fill="#356CFF"/>
<path d="M179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056Z" fill="#356CFF"/>
<path d="M161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457Z" fill="#356CFF"/>
<path fill-rule="evenodd" clip-rule="evenodd" d="M230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM237.948 77.9692C239.984 70.6965 240.917 65.242 228.446 65.242C215.975 65.242 211.818 71.9087 210.037 77.9692C208.255 84.0298 208.255 91.3025 219.538 91.3025C230.821 91.3025 235.911 85.2419 237.948 77.9692Z" fill="#356CFF"/>
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944M173.188 1.00056C173.188 1.00056 156.777 1.00043 141.337 1.00043M173.188 1.00056L141.337 1.00043M141.337 20.3944C146.088 20.3944 150.839 20.3944 159.747 20.3944M141.337 20.3944H159.747M159.747 20.3944C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273M148.463 48.273C139.556 48.273 125.188 48.273 125.188 48.273M148.463 48.273L125.188 48.273M125.188 48.273L124 37.97M124 37.97C124 37.97 141.931 37.97 147.87 37.97M124 37.97L147.87 37.97M147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852M151.433 29.4852C146.682 29.4852 138.962 29.4852 131.836 29.4852M151.433 29.4852H131.836M131.836 29.4852C120.142 29.4852 125.897 1.00043 141.337 1.00043M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057ZM96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM87.7028 29.4852L100.768 13.1217L103.143 29.4852H87.7028ZM89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457ZM108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM116.802 65.8487L109.081 90.697C109.081 90.697 117.989 90.697 122.74 90.697C127.491 90.697 133.307 88.4021 135.211 84.0304C137.586 78.5759 143.525 65.8487 132.242 65.8487C120.959 65.8487 116.802 65.8487 116.802 65.8487ZM179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056ZM161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457ZM230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM228.446 65.242C240.917 65.242 239.984 70.6965 237.948 77.9692C235.911 85.2419 230.821 91.3025 219.538 91.3025C208.255 91.3025 208.255 84.0298 210.037 77.9692C211.818 71.9087 215.975 65.242 228.446 65.242Z" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M15.2524 55.5451L21.191 100.999L24.7541 100.999L18.8156 55.5451L15.2524 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M8.12646 55.5451L14.065 100.999L15.2527 100.999L9.31417 55.5451L8.12646 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M1 55.5451L6.93853 100.999L7.53239 100.999L1.59385 55.5451L1 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
<path d="M15.2524 48.2724L30.0988 1H33.6619L18.8156 48.2724H15.2524Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M8.12646 48.2724L22.9728 1H24.1605L9.31417 48.2724H8.12646Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
<path d="M1 48.2724L15.8463 1H16.4402L1.59385 48.2724H1Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
<path d="M85.3271 55.5457H67.5116L87 12.7363L44.3513 68.2729H58.6038L43.1636 101L85.3271 55.5457Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.18771" stroke-miterlimit="16"/>
</svg>

After

Width:  |  Height:  |  Size: 5.7 KiB

Binary file not shown.
+110
View File
@@ -0,0 +1,110 @@
[
{
"prompt": "Young man skating with a skateboard on the ramps with graffiti of a park with trees, on a sunny day.",
"image_path": "assets/images/mixkit-boy-skating-with-a-skateboard-in-a-park-with-ramps-34389.png"
},
{
"prompt": "In the midst of the joyous New Year's Eve celebration, the cheerful group of friends, their spirits lifted by the festivities, decides to immortalize the moment with a vibrant snapshot",
"image_path": "assets/images/mixkit-a-cheerful-group-of-friends-celebrate-new-years-eve-and-51525.png"
},
{
"prompt": "A man and a woman playing in a field with grass, during a bright afternoon, while cars pass by in the distance.",
"image_path": "assets/images/mixkit-a-cute-couple-playing-on-the-grass-4688.png"
},
{
"prompt": "Aerial view of a rocky mountain in the forest at a sunny day drone flight footage",
"image_path": "assets/images/mixkit-aerial-view-of-a-rocky-mountain-in-the-forest-50589.png"
},
{
"prompt": "A little girl wearing a pink security helmet and denim overall discovers the art of cycling amidst the serene park, as the camera captures her graceful progress.",
"image_path": "assets/images/mixkit-a-little-girl-cruises-through-the-forest-path-on-her-50088.png"
},
{
"prompt": "Aerial shot of a beach shore with sea waves. Big rocks on the sand at an alone beach.",
"image_path": "assets/images/mixkit-aerial-shot-of-a-beach-with-sea-waves-1087.png"
},
{
"prompt": "Young woman cleaning her house decorated with plants and decorations, while dancing happily to music in her headphones.",
"image_path": "assets/images/mixkit-woman-cleaning-her-house-dancing-happy-43379.png"
},
{
"prompt": "Aerial tour in a meadow surrounded by hills on the horizon, while some birds fly low over a lake.",
"image_path": "assets/images/mixkit-birds-flying-low-over-a-lake-in-a-meadow-41417.png"
},
{
"prompt": "In the video, a lone rider guides a majestic horse across an expansive, open field as the sun sets in the background. The rider, dressed in a classic blue shirt and wide-brimmed hat, sits confidently in the saddle, silhouetted against the warm glow of the evening sky. The horse moves gracefully, its mane and tail flowing with each step, creating a sense of harmony between horse and rider. Surrounding the pair, towering trees form a natural border, their leaves gently rustling in the breeze. The shadows lengthen on the ground, accentuating the serene and timeless feel of the scene. The distant hills and wooden fences frame the horizon, adding depth to the tranquil landscape. A few horses graze peacefully in the background, blending into the pastoral setting. The overall ambiance evokes a sense of calmness and quietude, capturing a perfect moment in the golden light of dusk.",
"image_path": "assets/images/mixkit-a-rancher-riding-a-horse-at-sunset-1143.png"
},
{
"prompt": "In the video, a martial artist dressed in a traditional white uniform with a black belt demonstrates a series of precise movements against a stark black background. The individual gracefully transitions between stances, embodying a sense of focused discipline and control. Each motion is executed with a deliberate pace, showcasing the fluidity of martial arts techniques. The soft lighting creates subtle highlights on the uniform, adding depth to the figure as it moves. The practitioner begins with an open-hand pose, feet firmly grounded, gradually shifting to a powerful forward punch. The fluidity of the sequence displays a mastery of balance and poise. Every trajectory of the limbs is precise and deliberate, capturing the elegance and strength of martial arts. The serene, isolated setting enhances the intensity and concentration of the practitioner. This visual presentation is an elegant interplay of motion and stillness, displaying the art form's discipline and grace.",
"image_path": "assets/images/mixkit-a-young-man-practicing-his-karate-moves-49635.png"
},
{
"prompt": "In a serene and softly lit yoga studio, three individuals engage in a yoga session, each performing an upward-facing stretch. The central figure is a woman with shoulder-length brown hair, dressed in a light cropped top and green leggings, her posture reflecting grace and concentration. To her right, another participant, a woman in a purple outfit, mirrors the pose with equal poise. On her left, a person with a bun focuses intently, supported slightly by yoga blocks beneath their hands. The warm-colored wooden floor contrasts soothingly with the soft pastel mural on the back wall, featuring an abstract design and partial visage of a serene face. Natural light floods the space from a large window on the right, where lush greens peek through, adding an element of tranquility. In the corner of the room, a collection of meditation instruments, including a gong and a Buddha statue, subtly frame the peaceful setting. The mood is calm yet focused, as all three participants are deeply engaged in their practice. The scene combines elements of balance, harmony, and a shared journey towards mindfulness. This depiction captures the essence of a yoga session that blends personal growth with collective experience.",
"image_path": "assets/images/mixkit-small-group-of-people-doing-yoga-together-43730.png"
},
{
"prompt": "In the deep blue expanse of the ocean, two dolphins glide effortlessly, their sleek bodies reflecting the sunlight filtering through the water. The prominent shadows and caustics create a shimmering effect on their skin, capturing the beauty of their natural habitat. Each dolphin moves with a fluid grace, occasionally interacting with gentle nudges, showcasing their playful and social nature. The scene is vibrant and dynamic, with the clear blue background accentuating the dolphins' movements, making it an ideal subject for AI recreation.",
"image_path": "assets/images/mixkit-dolphins-underwater-4133.png"
},
{
"prompt": "A bustling ski slope comes alive with skiers descending a pristine, snow-covered hill, surrounded by towering, snow-draped evergreens. Several figures stand atop the slope, silhouetted against a clear blue sky, preparing to embark on their ski run. The chair lift on the right continuously drops off eager adventurers, adding to the excitement at the hilltop. Each skier, clad in colorful winter gear, carves distinct paths into the textured snow as they weave their way down. The interplay of sunlight and shadows accentuates the myriad tracks etched into the slope, creating a dynamic visual rhythm. The scene captures a vibrant winter wonderland, full of action and the thrill of a perfect ski day.",
"image_path": "assets/images/mixkit-skiers-on-a-snowy-slope-3327.png"
},
{
"prompt": "A determined climber is scaling a massive rock face, showcasing exceptional strength and skill. The person, clad in a teal shirt and dark pants, climbs with precision, their movements measured and deliberate. They are secured by climbing gear, which includes ropes and a harness, emphasizing their commitment to safety. The rugged texture of the sandy-colored rock provides an imposing backdrop, adding drama and scale to the climb. In the distance, other large rock formations and sparse vegetation can be seen under a bright, overcast sky, contributing to the natural and adventurous atmosphere. The scene captures a moment of focus and challenge, highlighting the climber's tenacity and the breathtaking environment.",
"image_path": "assets/images/mixkit-alpinist-climbing-a-huge-rock-in-a-desert-43306.png"
},
{
"prompt": "A silver SUV drives along a winding, snow-covered mountain road, with dense pine trees blanketed in snow lining both sides. The scene is serene, with the vehicle moving smoothly, possibly on a winter journey or vacation. As the SUV disappears around the bend, another, darker SUV follows, creating a sense of motion and perspective on the snow-dusted asphalt. The towering, snow-laden rock formation to the right contrasts with the dark green of the pines, highlighting the peacefulness of the wintry landscape.",
"image_path": "assets/images/mixkit-curve-on-a-snowy-forest-road-3317.png"
},
{
"prompt": "A solitary boat glides across the expansive, tranquil expanse of a serene lake. The vessel leaves a gentle wake behind, creating delicate ripples across the mirror-like surface. The water appears a rich shade of teal, seamlessly blending with the sky at the horizon. Silhouettes of distant trees are faintly visible, creating a picturesque backdrop that enhances the solitary journey of the boat. The sky is a calm gradient, shifting from soft oranges near the shore to the pale blues above. In the distance, a few slender poles emerge from the water, remnants of an old structure or natural formation. The mood of the scene is one of peace and solitude, with the boat journeying steadily through the quiet landscape. There is a sense of endless possibilities as the boat moves toward the unseen beyond the frame. The simplicity and stillness of the scene invite contemplation and reflection, encapsulating a perfect moment of quietude on the water.",
"image_path": "assets/images/mixkit-motorboat-on-a-large-lake-with-turquoise-blue-waters-4996.png"
},
{
"prompt": "A man wearing grey shorts jumps rope in a gym, weights and gym equipment in the background.",
"image_path": "assets/images/gray_short_man.jpg"
},
{
"prompt": "Flying over a peninsula covered in bushy trees, while discovering the sea around it, painted a beautiful turquoise blue, on a sunny day.",
"image_path": "assets/images/peninsula.jpg"
},
{
"prompt": "Skillful cyclist doing a wheelie on a bike while riding through a forest, on a dirt road, surrounded by many trees, in the morning.",
"image_path": "assets/images/cyclist.jpg"
},
{
"prompt": "Some friends dancing and having fun together in circles, at a party surrounded by colored lights at a party, in a fancy old place, in a view from below them.",
"image_path": "assets/images/friends.jpg"
},
{
"prompt": "A saxophonist wearing a blazer dances while playing a song in a park.",
"image_path": "assets/images/saxophonist.jpg"
},
{
"prompt": "Romantic couple embracing and looking at each other in the middle of a forest, during a break on a road trip through nature.",
"image_path": "assets/images/romance.jpg"
},
{
"prompt": "Man dressed in 80's style dances very happily in his kitchen while listening to music on his radio and drinking wine.",
"image_path": "assets/images/80s_dance.jpg"
},
{
"prompt": "Pair of jazz musicians performing a song with their saxophone and trombone on an abandoned train.",
"image_path": "assets/images/jazz.jpg"
},
{
"prompt": "A young woman with short hair wearing pink sunglasses chews gum and makes a bubble gum with the city in the background.",
"image_path": "assets/images/pink.jpg"
},
{
"prompt": "Natural aerial landscape with a relief covered with abundant trees and vegetation and a thick layer of mist.",
"image_path": "assets/images/natural.jpg"
},
{
"prompt": "Loving couple sitting on a log on the shore of a lake outside, sharing an affectionate hug.",
"image_path": "assets/images/couple.jpg"
}
]
+7
View File
@@ -0,0 +1,7 @@
# FastVideo/assets/videos
This folder is used to store **video assets for examples**, primarily **input videos** consumed by scripts under `FastVideo/examples/`.
- **Typical contents**: short input clips for demos (e.g., video2world / image2video examples).
- **Non-critical**: these assets are for convenience and are not required to use the FastVideo library.
- **Large files**: avoid committing large videos to git; prefer shared storage or download-on-demand.
Binary file not shown.
+106
View File
@@ -0,0 +1,106 @@
# FVD (Fréchet Video Distance) Benchmark
Evaluate generated video quality using FVD with the I3D feature extractor.
## Quick Start
**Run the benchmark:**
```bash
bash benchmarks/scripts/run.sh
```
That's it! The script auto-installs dependencies and runs the benchmark.
**To customize:** Edit `benchmarks/fvd/run_fvd.py` to change:
- Video paths (`real_dir`, `gen_dir`)
- Number of videos, frames, sampling strategy
- Device, batch size, caching, etc.
## Advanced Usage (CLI)
For more control without editing Python files, use the CLI.
**First-time setup** (one-time per pod/environment):
```bash
bash benchmarks/scripts/setup_fvd.sh
```
Then run any configuration you want:
```bash
# Custom configuration
python -m benchmarks.fvd.cli \
--real-path data/real/ \
--gen-path outputs/gen/ \
--num-videos 1024 \
--num-frames 32 \
--clip-strategy random \
--batch-size 32 \
--seed 42 \
--extractor clip
```
**Standard protocols:**
```bash
# Use predefined protocols
python -m benchmarks.fvd.cli \
--real-path data/real/ \
--gen-path outputs/gen/ \
--protocol fvd2048_16f # or fvd2048_128f, quick_test, etc.
```
This would use i3d model by default as the feature extractor
**Feature caching** (speed up repeated evaluations):
```bash
python -m benchmarks.fvd.cli \
--real-path data/real/ \
--gen-path outputs/gen/ \
--protocol fvd2048_16f \
--cache-real-features fvd-cache/extractor_name # Directory path (will save/load fvd-cache/extractor_name/extractor-name_real_features.pkl)
```
Run `python -m benchmarks.fvd.cli --help` for all options.
## Available Protocols
- `fvd2048_16f` - Standard (2048 videos, 16 frames)
- `fvd2048_128f` - Long videos (128 frames)
- `fvd2048_128f_subsample8` - Subsampled long videos
- `quick_test` - Fast testing (10 videos)
## Configuration Options
Key options in `FVDConfig`:
```python
num_videos=2048, # Videos to evaluate
num_frames_per_clip=16, # Frames per clip
clip_strategy='beginning', # beginning|random|uniform|middle|sliding
frame_stride=1, # Frame subsampling
batch_size=32, # GPU batch size
device='cuda', # cuda|cpu
cache_real_features=None, # Cache path for speed
seed=42, # Reproducibility
extractor='i3d', # i3d|clip|videomae
```
## Programmatic Usage
```python
from benchmarks.fvd import compute_fvd_with_config, FVDConfig
config = FVDConfig.fvd2048_16f() # or custom config
results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
print(f"FVD: {results['fvd']:.2f}")
```
## Notes
- Requires minimum 10 frames per clip
- Supports both video files (.mp4, .avi, etc.) and frame directories
- `--cache-real-features` expects a **directory path** (e.g., `cache/real`), it will automatically create/load `real_features.pkl` inside that directory
+38
View File
@@ -0,0 +1,38 @@
"""
FastVideo Frechet Video Distance (FVD) Benchmark Module.
>>> from fastvideo.benchmarks.fvd import compute_fvd_with_config, FVDConfig
>>> config = FVDConfig.fvd2048_16f() # Standard protocol
>>> results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
>>> print(f"FVD: {results['fvd']:.2f}")
"""
from .fvd import (
compute_fvd,
compute_fvd_with_config,
compute_frechet_distance,
compute_statistics,
FVDConfig,
)
from .feature_extractors import (BaseFeatureExtractor, I3DFeatureExtractor,
load_extractor)
from .video_utils import (
load_video_auto,
sample_clips_from_video,
load_video_clips_streaming,
ClipSamplingStrategy,
)
__all__ = [
'compute_fvd',
'compute_fvd_with_config',
'compute_frechet_distance',
'compute_statistics',
'FVDConfig',
'BaseFeatureExtractor',
'I3DFeatureExtractor',
'load_extractor',
'load_video_auto',
'sample_clips_from_video',
'load_video_clips_streaming',
'ClipSamplingStrategy',
]
+107
View File
@@ -0,0 +1,107 @@
import argparse
import sys
import traceback
from .fvd import compute_fvd_with_config, FVDConfig
def main() -> int:
parser = argparse.ArgumentParser(
description='Compute Fréchet Video Distance (FVD)')
# Required arguments
parser.add_argument('--real-path',
type=str,
required=True,
help='Path to real videos')
parser.add_argument('--gen-path',
type=str,
required=True,
help='Path to generated videos')
# Extractor selection
parser.add_argument('--extractor',
type=str,
default='i3d',
choices=['i3d', 'clip', 'videomae'],
help='Feature extractor model to use (default: i3d)')
# Standard args
parser.add_argument('--seed',
type=int,
default=None,
help='Random seed for reproducibility')
parser.add_argument('--protocol',
type=str,
default=None,
choices=['fvd2048_16f', 'fvd2048_128f', 'quick_test'],
help='Use standard protocol (overrides other settings)')
parser.add_argument('--num-videos',
type=int,
default=2048,
help='Number of videos to use')
parser.add_argument('--num-frames',
type=int,
default=16,
help='Number of frames per clip')
parser.add_argument('--clip-strategy',
type=str,
default='beginning',
help='Clip sampling strategy')
parser.add_argument('--batch-size',
type=int,
default=32,
help='Batch size for feature extraction')
parser.add_argument('--device',
type=str,
default='cuda',
help='Device to use (cuda or cpu)')
parser.add_argument('--cache-real-features',
type=str,
default=None,
help='Path to cache real video features')
parser.add_argument('--quiet',
action='store_true',
help='Suppress progress output')
args = parser.parse_args()
# Create config
if args.protocol:
protocol_map = {
'fvd2048_16f': FVDConfig.fvd2048_16f,
'fvd2048_128f': FVDConfig.fvd2048_128f,
'quick_test': FVDConfig.quick_test,
}
config = protocol_map[args.protocol]()
# Apply overrides
config.device = args.device
config.cache_real_features = args.cache_real_features
config.extractor_model = args.extractor # Apply extractor arg
else:
config = FVDConfig(
num_videos=args.num_videos,
num_frames_per_clip=args.num_frames,
extractor_model=args.extractor, # Apply extractor arg
clip_strategy=args.clip_strategy,
batch_size=args.batch_size,
device=args.device,
cache_real_features=args.cache_real_features,
seed=args.seed)
try:
_ = compute_fvd_with_config(
args.real_path, # noqa: F841
args.gen_path,
config,
verbose=not args.quiet)
return 0
except Exception as e:
print(f"Error: {e}", file=sys.stderr)
traceback.print_exc(file=sys.stderr)
return 1
if __name__ == '__main__':
sys.exit(main())
+264
View File
@@ -0,0 +1,264 @@
"""
Pluggable Feature Extractors for FVD Computation.
Supports I3D (standard), CLIP, and VideoMAE via a common interface.
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from abc import ABC, abstractmethod
from huggingface_hub import hf_hub_download
from tqdm import tqdm
try:
from transformers import CLIPModel, CLIPProcessor, VideoMAEModel
TRANSFORMERS_AVAILABLE = True
except ImportError:
TRANSFORMERS_AVAILABLE = False
class BaseFeatureExtractor(ABC, nn.Module):
"""Abstract base class for all video feature extractors."""
def __init__(self, device: str = 'cuda'):
super().__init__()
self.device = torch.device(
device if torch.cuda.is_available() else 'cpu')
@property
@abstractmethod
def feature_dim(self) -> int:
"""Dimension of the output feature vector."""
pass
@abstractmethod
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
"""
Args:
videos: [B, T, C, H, W] in [0, 255] range.
Returns:
Preprocessed tensor ready for the model.
"""
pass
@abstractmethod
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
"""
Extract features for a single batch.
Args:
videos: [B, T, C, H, W] (raw input)
Returns:
Features: [B, feature_dim]
"""
pass
@torch.no_grad()
def extract_features(self,
videos: torch.Tensor,
batch_size: int = 32,
verbose: bool = True) -> torch.Tensor:
"""
Extract features for a large tensor of videos by batching.
"""
N = len(videos)
all_features = []
iterator = range(0, N, batch_size)
if verbose:
iterator = tqdm(
iterator,
desc=f"Extracting features ({self.__class__.__name__})")
for i in iterator:
batch = videos[i:i + batch_size].to(self.device)
features = self.extract_features_batch(batch)
all_features.append(features.cpu())
return torch.cat(all_features, dim=0)
# 1. I3D Extractor (The Standard FVD Metric)
class I3DFeatureExtractor(BaseFeatureExtractor):
REPO_ID = 'flateon/FVD-I3D-torchscript'
MODEL_FILENAME = 'i3d_torchscript.pt'
def __init__(self, device: str = 'cuda', cache_dir: str | None = None):
super().__init__(device)
self.cache_dir = cache_dir
self.model = self._load_model()
self.model.eval()
self.model.to(self.device)
@property
def feature_dim(self) -> int:
return 400
def _load_model(self) -> torch.nn.Module:
try:
model_path = hf_hub_download(repo_id=self.REPO_ID,
filename=self.MODEL_FILENAME,
cache_dir=self.cache_dir)
return torch.jit.load(model_path, map_location=self.device)
except Exception as e:
raise RuntimeError(f"Failed to load I3D model: {e}") from e
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
"""Standard I3D preprocessing: Resize to 224, Norm to [-1, 1]."""
B, T, C, H, W = videos.shape
if T < 10:
raise ValueError(f"I3D requires at least 10 frames, got {T}")
# Normalize to [0, 1]
if videos.max() > 1.0:
videos = videos / 255.0
# Scale to [-1, 1]
videos = videos * 2.0 - 1.0
# Resize to 224x224
if H != 224 or W != 224:
videos = videos.reshape(B * T, C, H, W)
videos = F.interpolate(videos,
size=(224, 224),
mode='bilinear',
align_corners=False)
videos = videos.reshape(B, T, C, 224, 224)
# [B, T, C, H, W] -> [B, C, T, H, W]
return videos.permute(0, 2, 1, 3, 4).contiguous()
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
batch = self.preprocess(videos)
# TorchScript I3D returns raw logits when return_features=True
return self.model(batch,
rescale=False,
resize=False,
return_features=True)
# 2. CLIP Extractor (Semantic/Content Quality)
class CLIPFeatureExtractor(BaseFeatureExtractor):
def __init__(self,
device: str = 'cuda',
model_name: str = "openai/clip-vit-base-patch32"):
if not TRANSFORMERS_AVAILABLE:
raise ImportError(
"Please install transformers: pip install transformers")
super().__init__(device)
self.processor = CLIPProcessor.from_pretrained(model_name)
self.model = CLIPModel.from_pretrained(model_name).to(self.device)
self.model.eval()
self._feature_dim = self.model.config.projection_dim
@property
def feature_dim(self) -> int:
return self._feature_dim
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
# Ensure values are [0, 255]
if videos.max() <= 1.0:
videos = videos * 255.0
return videos.to(torch.uint8)
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
# Input: [B, T, C, H, W]
B, T, C, H, W = videos.shape
videos = self.preprocess(videos)
# Flatten B*T to treat frames as images
images = videos.view(B * T, C, H, W)
# HF Processor
inputs = self.processor(images=images,
return_tensors="pt",
padding=True)
inputs = {k: v.to(self.device) for k, v in inputs.items()}
# Extract features [B*T, Dim]
outputs = self.model.get_image_features(**inputs)
# Reshape [B, T, Dim] and Average Pooling over time
outputs = outputs.view(B, T, -1)
return outputs.mean(dim=1)
# 3. VideoMAE Extractor (Structure/Motion Quality)
class VideoMAEFeatureExtractor(BaseFeatureExtractor):
def __init__(self,
device: str = 'cuda',
model_name: str = "MCG-NJU/videomae-base"):
if not TRANSFORMERS_AVAILABLE:
raise ImportError(
"Please install transformers: pip install transformers")
super().__init__(device)
self.model = VideoMAEModel.from_pretrained(model_name).to(self.device)
self.model.eval()
self.register_buffer(
'mean',
torch.tensor([0.485, 0.456, 0.406],
device=self.device).view(1, 1, 3, 1, 1))
self.register_buffer(
'std',
torch.tensor([0.229, 0.224, 0.225],
device=self.device).view(1, 1, 3, 1, 1))
@property
def feature_dim(self) -> int:
return self.model.config.hidden_size
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
"""
Efficient GPU-based preprocessing.
Input: [B, T, C, H, W] in range [0, 255]
"""
B, T, C, H, W = videos.shape
# 1. Resize to 224x224
if H != 224 or W != 224:
videos = videos.view(B * T, C, H, W)
videos = F.interpolate(videos,
size=(224, 224),
mode='bilinear',
align_corners=False)
videos = videos.view(B, T, C, 224, 224)
# 2. Normalize to [0, 1]
if videos.dtype != torch.float32:
videos = videos.float()
if videos.max() > 1.0:
videos = videos / 255.0
# 3. Apply ImageNet Mean/Std
return (videos - self.mean) / self.std
def extract_features_batch(self, videos: torch.Tensor) -> torch.Tensor:
# Input: [B, T, C, H, W]
# Fast GPU Preprocessing
pixel_values = self.preprocess(videos)
# Forward pass
outputs = self.model(pixel_values)
# Global Average Pooling of last hidden state [B, T_patches, 768] -> [B, 768]
return outputs.last_hidden_state.mean(dim=1)
# Factory
def load_extractor(name: str, device: str = 'cuda') -> BaseFeatureExtractor:
name = name.lower()
if name == 'i3d':
return I3DFeatureExtractor(device)
elif name == 'clip':
return CLIPFeatureExtractor(device)
elif name == 'videomae':
return VideoMAEFeatureExtractor(device)
else:
raise ValueError(
f"Unknown extractor: {name}. Options: i3d, clip, videomae")
+405
View File
@@ -0,0 +1,405 @@
import numpy as np
import scipy.linalg
import torch
from pathlib import Path
from collections.abc import Iterator
import pickle
from dataclasses import dataclass, field
from .feature_extractors import BaseFeatureExtractor, load_extractor
from .video_utils import ClipSamplingStrategy, load_video_clips_streaming
def compute_statistics(features: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
"""Compute mean and covariance."""
mu = np.mean(features, axis=0)
sigma = np.cov(features, rowvar=False)
return mu, sigma
def compute_frechet_distance(mu1: np.ndarray,
sigma1: np.ndarray,
mu2: np.ndarray,
sigma2: np.ndarray,
eps: float = 1e-6) -> float:
"""
Compute Fréchet distance between two Gaussians.
"""
sigma1 = sigma1 + eps * np.eye(sigma1.shape[0])
sigma2 = sigma2 + eps * np.eye(sigma2.shape[0])
diff = mu1 - mu2
mean_distance = np.sum(diff**2)
trace_sum = np.trace(sigma1 + sigma2)
covmean = scipy.linalg.sqrtm(sigma1 @ sigma2)
if np.iscomplexobj(covmean):
if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3):
print(
f"Warning: Imaginary component: {np.max(np.abs(covmean.imag))}")
covmean = covmean.real
trace_product = np.trace(covmean)
fvd = mean_distance + trace_sum - 2 * trace_product
return float(fvd)
@dataclass
class FVDConfig:
# default configuration for FVD computation:
# Video selection
num_videos: int = 2048
# Feature Extractor Selection
extractor_model: str = 'i3d' # Options: 'i3d', 'clip', 'videomae'
# Clip sampling
num_frames_per_clip: int = 16
num_clips_per_video: int = 1
clip_strategy: str | ClipSamplingStrategy = 'beginning'
# Temporal subsampling
frame_stride: int = 1 # 1=no subsampling, 2=every 2nd, 8=every 8th
temporal_stride: int = 1 # For sliding window clips
# Data processing
video_extensions: list[str] = field(
default_factory=lambda: ['.mp4', '.avi', '.mov', '.mkv'])
support_frame_dirs: bool = True
# Computation
batch_size: int = 32
device: str = 'cuda'
use_streaming: bool = True
resize_before_extraction: bool = True
# Caching
cache_real_features: str | None = None
i3d_model_path: str | None = None
# Reproducibility
seed: int | None = None
@classmethod
def fvd2048_16f(cls) -> 'FVDConfig':
"""Standard FVD protocol: 2048 videos, 16 frames, beginning clip."""
return cls(num_videos=2048,
num_frames_per_clip=16,
clip_strategy='beginning',
use_streaming=True)
@classmethod
def fvd2048_128f(cls) -> 'FVDConfig':
"""Long video protocol: 2048 videos, 128 frames."""
return cls(num_videos=2048,
num_frames_per_clip=128,
clip_strategy='beginning',
use_streaming=True)
@classmethod
def quick_test(cls) -> 'FVDConfig':
"""Quick test config: 100 videos, 16 frames."""
return cls(num_videos=100,
num_frames_per_clip=16,
clip_strategy='beginning')
def to_dict(self) -> dict:
"""Export config to dict for logging"""
d = self.__dict__.copy()
d['clip_strategy'] = str(self.clip_strategy)
return d
def __str__(self) -> str:
"""Human-readable protocol name"""
desc = f"FVD_{self.extractor_model.upper()}_{self.num_videos}_{self.num_frames_per_clip}f"
if self.frame_stride > 1:
desc += f"_subsample{self.frame_stride}"
if self.num_clips_per_video > 1:
desc += f"_{self.num_clips_per_video}clips"
if self.clip_strategy != 'beginning':
desc += f"_{self.clip_strategy}"
return desc
def extract_features_streaming(video_generator: Iterator[torch.Tensor],
extractor: BaseFeatureExtractor,
batch_size: int = 32,
max_clips: int | None = None,
verbose: bool = True) -> np.ndarray:
"""
Extract features from a video clip generator using streaming.
"""
all_features = []
batch = []
if verbose:
print(f"Extracting features with batch_size={batch_size}...")
with torch.no_grad():
for clip_count, clip in enumerate(video_generator):
batch.append(clip)
# Process batch when full
if len(batch) == batch_size:
batch_tensor = torch.stack(batch).to(extractor.device)
features = extractor.extract_features_batch(batch_tensor)
all_features.append(features.detach().cpu().numpy())
batch = []
if verbose and clip_count % (batch_size * 10) == 0:
print(f"Processed {clip_count} clips...")
if max_clips is not None and clip_count >= max_clips:
break
# Process remaining clips
if len(batch) > 0:
batch_tensor = torch.stack(batch).to(extractor.device)
features = extractor.extract_features_batch(batch_tensor)
all_features.append(features.detach().cpu().numpy())
if len(all_features) == 0:
raise RuntimeError("No features extracted - check video loading")
features = np.concatenate(all_features, axis=0)
if verbose:
print(f"Extracted {len(features)} feature vectors")
return features
def load_or_compute_features(videos: str | Path | torch.Tensor,
extractor: BaseFeatureExtractor,
config: FVDConfig,
cache_path: str | None = None,
cache_name: str = "real_features") -> np.ndarray:
"""Load features from cache or compute (with streaming support)"""
if cache_path is not None:
script_dir = Path(__file__).parent
cache_dir = script_dir / cache_path
cache_file = cache_dir / f"{config.extractor_model}_{cache_name}.pkl"
if cache_file.exists():
print(f"Loading cached features from {cache_file}")
with open(cache_file, 'rb') as f:
features = pickle.load(f)
# Validate and limit based on config
max_features = config.num_videos * config.num_clips_per_video
if len(features) < max_features:
print(
f"WARNING: Cache has {len(features)} features but need {max_features}"
)
print("Cached features insufficient - will recompute...")
elif len(features) > max_features:
print(
f"Using {max_features} features from cache (truncated from {len(features)})"
)
features = features[:max_features]
return features
else:
print(f"Using all {len(features)} cached features")
return features
print("Computing features from scratch...")
if isinstance(videos, (str | Path)):
target_size = (224, 224) if config.resize_before_extraction else None
video_generator = load_video_clips_streaming(
videos,
num_frames=config.num_frames_per_clip,
max_videos=config.num_videos,
clip_strategy=config.clip_strategy,
frame_stride=config.frame_stride,
num_clips_per_video=config.num_clips_per_video,
video_extensions=config.video_extensions,
support_frame_dirs=config.support_frame_dirs,
target_size=target_size,
verbose=True)
max_clips = config.num_videos * config.num_clips_per_video
features = extract_features_streaming(video_generator,
extractor,
batch_size=config.batch_size,
max_clips=max_clips,
verbose=True)
else:
print(f"Extracting features from {len(videos)} video tensors...")
features = extractor.extract_features(videos,
batch_size=config.batch_size,
verbose=True)
features = features.numpy()
# Validate feature count
expected_count = config.num_videos * config.num_clips_per_video
if len(features) < expected_count:
raise ValueError(
f"ERROR: Only extracted {len(features)} features, but need {expected_count}!\n"
f"Found fewer videos than expected. Check your video directory.")
elif len(features) > expected_count:
print(f"Truncating {len(features)} features to {expected_count}")
features = features[:expected_count]
# Cache features if requested
if cache_path is not None:
script_dir = Path(__file__).parent
cache_dir = script_dir / cache_path
cache_dir.mkdir(parents=True, exist_ok=True)
cache_file = cache_dir / f"{config.extractor_model}_{cache_name}.pkl"
print(f"Caching features to {cache_file}")
with open(cache_file, 'wb') as f:
pickle.dump(features, f)
return features
def compute_fvd_with_config(real_videos: str | Path | torch.Tensor,
gen_videos: str | Path | torch.Tensor,
config: FVDConfig,
verbose: bool = True) -> dict:
"""
Compute FVD using a standardized configuration.
This is the recommended way to compute FVD for reproducibility.
Args:
real_videos: Path or tensors
gen_videos: Path or tensors
config: FVDConfig specifying protocol
verbose: Print progress
Returns:
results: Dictionary with:
- 'fvd': FVD score (float)
- 'protocol': Protocol name (str)
- 'model': Feature extractor model name (str)
- 'config': Configuration dict
Example:
>>> config = FVDConfig.fvd2048_16f()
>>> results = compute_fvd_with_config('data/real/', 'outputs/gen/', config)
>>> print(f"FVD: {results['fvd']:.2f}")
"""
# Seed for reproducibility
if config.seed is not None:
import random as _rnd
_rnd.seed(config.seed)
np.random.seed(config.seed)
torch.manual_seed(config.seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(config.seed)
if verbose:
print("=" * 70)
print(f"Computing FVD with protocol: {config}")
print(f"Model: {config.extractor_model.upper()}")
print("=" * 70)
print("\nConfiguration:")
for key, value in config.to_dict().items():
print(f" {key}: {value}")
print()
# Initialize Extractor using Factory
if verbose:
print(
f"\nInitializing {config.extractor_model.upper()} model on {config.device}..."
)
extractor = load_extractor(config.extractor_model, device=config.device)
# Extract features
if verbose:
print(f"\n{'='*70}")
print("Extracting REAL video features...")
print(f"{'='*70}")
real_features = load_or_compute_features(
videos=real_videos,
extractor=extractor,
config=config,
cache_path=config.cache_real_features,
cache_name="real_features")
if verbose:
print(f"\n{'='*70}")
print("Extracting GENERATED video features...")
print(f"{'='*70}")
gen_features = load_or_compute_features(videos=gen_videos,
extractor=extractor,
config=config,
cache_path=None,
cache_name="gen_features")
if verbose:
print(f"\nReal videos/clips: {len(real_features)}")
print(f"Generated videos/clips: {len(gen_features)}")
print(f"\n{'='*70}")
print("Computing statistics...")
print(f"{'='*70}")
mu_real, sigma_real = compute_statistics(real_features)
mu_gen, sigma_gen = compute_statistics(gen_features)
if verbose:
print(f"\n{'='*70}")
print("Computing Fréchet distance...")
print(f"{'='*70}")
fvd = compute_frechet_distance(mu_real, sigma_real, mu_gen, sigma_gen)
if verbose:
print(f"\n{'='*70}")
print(f"FVD Score ({config.extractor_model.upper()}): {fvd:.4f}")
print(f"Protocol: {config}")
print(f"{'='*70}\n")
results = {
'fvd': fvd,
'protocol': str(config),
'model': config.extractor_model,
'config': config.to_dict(),
}
return results
def compute_fvd(real_videos: str | Path | torch.Tensor,
gen_videos: str | Path | torch.Tensor,
num_frames: int = 16,
batch_size: int = 32,
device: str = 'cuda',
num_videos: int | None = 2048,
cache_real_features: str | None = None,
i3d_model_path: str | None = None,
seed: int | None = None,
verbose: bool = True) -> float:
"""
Backward compatibility wrapper for computing FVD (defaults to I3D).
"""
num_videos = num_videos if num_videos is not None else 2048
config = FVDConfig(
num_videos=num_videos,
num_frames_per_clip=num_frames,
extractor_model='i3d', # Default to I3D
batch_size=batch_size,
device=device,
cache_real_features=cache_real_features,
i3d_model_path=i3d_model_path,
seed=seed,
)
result = compute_fvd_with_config(real_videos, gen_videos, config, verbose)
return result['fvd']
+142
View File
@@ -0,0 +1,142 @@
"""I3D Feature Extractor for FVD Computation"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from pathlib import Path
from huggingface_hub import hf_hub_download
from tqdm import tqdm
from contextlib import suppress
class I3DFeatureExtractor(nn.Module):
"""
I3D feature extractor for FVD computation.
Extracts 400-dimensional features from videos using I3D model
trained on Kinetics-400.
"""
REPO_ID = 'flateon/FVD-I3D-torchscript'
MODEL_FILENAME = 'i3d_torchscript.pt'
def __init__(self,
device: str = 'cuda',
cache_dir: str | Path | None = None):
super().__init__()
self.device_str = device
if device == 'cuda' and not torch.cuda.is_available():
print(
"Warning: CUDA requested but not available – falling back to CPU"
)
self.device = torch.device('cpu')
else:
self.device = torch.device(device)
self.cache_dir: str | None
if cache_dir is not None:
self.cache_dir = str(Path(cache_dir).resolve())
else:
self.cache_dir = None # Use HF default cache
self.model = self._load_model()
self.model.eval()
with suppress(Exception):
self.model.to(self.device)
def _load_model(self) -> torch.nn.Module:
"""Download and load I3D TorchScript model from Hugging Face Hub."""
print(f"Loading I3D model from Hugging Face Hub ({self.REPO_ID})...")
try:
# Download model from Hugging Face Hub
model_path = hf_hub_download(repo_id=self.REPO_ID,
filename=self.MODEL_FILENAME,
cache_dir=self.cache_dir)
# Load directly to chosen device
model = torch.jit.load(model_path, map_location=self.device)
print("I3D model loaded successfully")
return model
except Exception as e:
raise RuntimeError(
f"Failed to load I3D model from Hugging Face Hub. Error: {e}\n"
f"Ensure you have internet connection and huggingface_hub installed:\n"
f"pip install huggingface_hub") from e
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
"""
Preprocess videos for I3D.
Args:
videos: [B, T, C, H, W], values in [0, 255]
Returns:
Preprocessed videos [B, C, T, 224, 224] (normalized and resized)
"""
B, T, C, H, W = videos.shape
if T < 10:
raise ValueError(f"I3D requires at least 10 frames, got {T}")
# Normalize to [0, 1] if needed
if videos.max() > 1.0:
videos = videos / 255.0
# Resize to 224x224 if needed
if H != 224 or W != 224:
videos = videos.reshape(B * T, C, H, W)
videos = F.interpolate(videos,
size=(224, 224),
mode='bilinear',
align_corners=False)
videos = videos.reshape(B, T, C, 224, 224)
# Convert to [B, C, T, H, W] format
videos = videos.permute(0, 2, 1, 3, 4).contiguous()
return videos
@torch.no_grad()
def extract_features(self,
videos: torch.Tensor,
batch_size: int = 32,
verbose: bool = True) -> torch.Tensor:
"""
Extract I3D features
Args:
videos: [N, T, C, H, W], values in [0, 255]
batch_size: Batch size for processing
verbose: Show progress bar
Returns:
Features [N, 400]
"""
N = len(videos)
all_features = []
iterator = range(0, N, batch_size)
if verbose:
iterator = tqdm(iterator, desc="Extracting I3D features")
for i in iterator:
batch = videos[i:i + batch_size].to(self.device)
batch = self.preprocess(batch) # Now returns [B, C, T, H, W]
# Use the HF model without rescale/resize (we handle it in preprocess)
features = self.model(batch,
rescale=False,
resize=False,
return_features=True)
all_features.append(features.cpu())
return torch.cat(all_features, dim=0)
def __call__(self,
videos: torch.Tensor,
batch_size: int = 32) -> torch.Tensor:
return self.extract_features(videos, batch_size=batch_size)
+54
View File
@@ -0,0 +1,54 @@
import sys
from pathlib import Path
root_dir = Path(__file__).parent.parent.parent
sys.path.insert(0, str(root_dir))
from benchmarks.fvd.fvd import FVDConfig, compute_fvd_with_config # noqa: E402
def main() -> None:
script_dir = Path(__file__).parent.resolve()
# Define directories
real_dir = "benchmarks/data/real_videos"
gen_dir = "benchmarks/data/generated_videos"
# Compare all 3 models
models_to_test = ['i3d', 'clip', 'videomae']
print(f"\n{'='*60}")
print("STARTING COMPARISON BENCHMARK")
print(f"{'='*60}")
for model_name in models_to_test:
print(f"\n>>> Running evaluation with {model_name.upper()}...")
try:
cfg = FVDConfig(
num_videos=650,
num_frames_per_clip=16,
extractor_model=model_name,
clip_strategy='beginning',
device='cuda',
seed=42,
# Use separate cache folders for each model to avoid conflicts
cache_real_features=str(script_dir / f'fvd-cache/{model_name}'),
)
results = compute_fvd_with_config(real_dir,
gen_dir,
cfg,
verbose=False)
print(f"FVD: {results['fvd']}\nModel: {results['model']}")
except Exception as e:
print(f"{model_name.upper()} Failed: {e}")
print(f"\n{'='*60}")
print("BENCHMARK COMPLETE")
print(f"{'='*60}")
if __name__ == '__main__':
main()
+97
View File
@@ -0,0 +1,97 @@
#!/usr/bin/env python3
import sys
from pathlib import Path
import shutil
import random
from fvd import compute_fvd_with_config, FVDConfig
script_path = Path(__file__).resolve()
fastvideo_root = script_path.parent.parent.parent
sys.path.insert(0, str(fastvideo_root))
def split_videos(video_dir: Path, n_per_subset: int = 128, seed: int = 42):
subset_a = video_dir.parent / 'bair_full_subset_A'
subset_b = video_dir.parent / 'bair_full_subset_B'
if subset_a.exists():
shutil.rmtree(subset_a)
if subset_b.exists():
shutil.rmtree(subset_b)
subset_a.mkdir(parents=True)
subset_b.mkdir(parents=True)
videos = sorted(video_dir.glob('*.mp4'))
random.seed(seed)
shuffled = list(videos)
random.shuffle(shuffled)
needed = n_per_subset * 2
if len(shuffled) > needed:
shuffled = shuffled[:needed]
mid = len(shuffled) // 2
print(f"\nSplitting {len(shuffled)} BAIR FULL videos:")
print(f" Subset A: {mid} videos")
print(f" Subset B: {len(shuffled) - mid} videos")
for v in shuffled[:mid]:
shutil.copy2(v, subset_a / v.name)
for v in shuffled[mid:]:
shutil.copy2(v, subset_b / v.name)
return subset_a, subset_b, mid
def validate_fvd(subset_a: Path, subset_b: Path, num_videos: int):
config = FVDConfig(num_videos=num_videos,
num_frames_per_clip=16,
clip_strategy='beginning',
batch_size=8,
device='cuda',
seed=42)
print("\n" + "=" * 70)
print("TEST 1: Identity Test")
print("=" * 70)
result1 = compute_fvd_with_config(real_videos=str(subset_a),
gen_videos=str(subset_a),
config=config,
verbose=False)
fvd_identity = result1['fvd']
print(f"\nIdentity FVD: {fvd_identity:.2f}")
print("\n" + "=" * 70)
print("TEST 2: Real vs Real")
print("=" * 70)
result2 = compute_fvd_with_config(real_videos=str(subset_a),
gen_videos=str(subset_b),
config=config,
verbose=False)
fvd_real = result2['fvd']
print(f"\nReal vs Real FVD: {fvd_real:.2f}")
print("\n" + "=" * 70)
print("RESULTS")
print("=" * 70)
print(f"Identity: {fvd_identity:.2f}")
print(f"Real vs Real: {fvd_real:.2f}")
def main() -> None:
bair_dir = Path('benchmarks/data/bair_full_videos')
subset_a, subset_b, count = split_videos(bair_dir,
n_per_subset=128,
seed=42)
validate_fvd(subset_a, subset_b, count)
if __name__ == '__main__':
main()
+490
View File
@@ -0,0 +1,490 @@
import torch
import cv2
import numpy as np
from pathlib import Path
from collections.abc import Iterator
from tqdm import tqdm
from enum import Enum
class ClipSamplingStrategy(Enum):
"""Clip sampling strategies for FVD evaluation."""
BEGINNING = 'beginning' # Take first N frames (most common)
RANDOM = 'random' # Random N consecutive frames
UNIFORM = 'uniform' # Uniformly spaced frames across video
MIDDLE = 'middle' # Middle N frames
SLIDING = 'sliding' # Multiple sliding windows
ALL = 'all' # All possible clips
def _load_video_cv2(video_path: str | Path,
num_frames: int | None = 16,
sample_strategy: str = 'uniform') -> torch.Tensor:
"""
Load video from video file using OpenCV.
Args:
video_path: Path to video file (MP4, AVI, MOV, MKV)
num_frames: Number of frames to extract
sample_strategy: 'uniform' or 'random'
Returns:
video: [T, C, H, W]
"""
video_path = str(video_path)
cap = cv2.VideoCapture(video_path)
if not cap.isOpened():
raise RuntimeError(f"Cannot open video: {video_path}")
frames = []
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
if num_frames is None:
# Read all available frames
while True:
ret, frame = cap.read()
if not ret:
break
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
frames.append(frame)
cap.release()
if len(frames) == 0:
raise RuntimeError(f"Video has 0 frames: {video_path}")
frames = np.stack(frames) # [T, H, W, C]
frames = torch.from_numpy(frames).permute(0, 3, 1,
2).float() # [T, C, H, W]
return frames
if total_frames == 0:
raise RuntimeError(f"Video has 0 frames: {video_path}")
# Determine frame indices for sampling
if total_frames < num_frames:
frame_indices = list(range(
total_frames)) + [total_frames - 1] * (num_frames - total_frames)
elif sample_strategy == 'uniform':
frame_indices = np.linspace(0, total_frames - 1, num_frames,
dtype=int).tolist()
elif sample_strategy == 'random':
frame_indices = sorted(
np.random.choice(total_frames, num_frames, replace=False))
else:
raise ValueError(f"Unknown sample_strategy: {sample_strategy}")
# Extract frames
for idx in frame_indices:
cap.set(cv2.CAP_PROP_POS_FRAMES, idx)
ret, frame = cap.read()
if not ret:
if len(frames) > 0:
frames.append(frames[-1].copy())
else:
h = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
w = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
frames.append(np.zeros((h, w, 3), dtype=np.uint8))
continue
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
frames.append(frame)
cap.release()
frames = np.stack(frames) # [T, H, W, C]
frames = torch.from_numpy(frames).permute(0, 3, 1,
2).float() # [T, C, H, W]
return frames
def _load_video_from_frames(
frame_dir: str | Path,
num_frames: int | None = 16,
sample_strategy: str = 'uniform',
frame_extensions: list[str] | None = None) -> torch.Tensor:
"""
Load video from directory of frame images.
Args:
frame_dir: Directory containing frames
num_frames: Number of frames to sample
sample_strategy: 'uniform' or 'random'
frame_extensions: Image file extensions to look for
Returns:
video: [T, C, H, W]
"""
if frame_extensions is None:
frame_extensions = ['.jpg', '.png', '.jpeg', '.bmp']
frame_dir = Path(frame_dir)
if not frame_dir.exists():
raise FileNotFoundError(f"Frame directory not found: {frame_dir}")
# Find all frames
frame_files: list[Path] = []
for ext in frame_extensions:
frame_files.extend(frame_dir.glob(f"*{ext}"))
if len(frame_files) == 0:
raise ValueError(
f"No frames found in {frame_dir} with extensions {frame_extensions}"
)
frame_files = sorted(frame_files, key=lambda x: x.name)
total_frames = len(frame_files)
# Determine frame indices
if num_frames is None:
frame_indices = list(range(total_frames))
else:
if total_frames < num_frames:
frame_indices = list(range(total_frames)) + [total_frames - 1] * (
num_frames - total_frames)
elif sample_strategy == 'uniform':
frame_indices = np.linspace(0,
total_frames - 1,
num_frames,
dtype=int).tolist()
elif sample_strategy == 'random':
frame_indices = sorted(
np.random.choice(total_frames, num_frames, replace=False))
else:
raise ValueError(f"Unknown sample_strategy: {sample_strategy}")
# Load frames
frames = []
for idx in frame_indices:
frame_path = frame_files[idx]
frame = cv2.imread(str(frame_path))
if frame is None:
raise RuntimeError(f"Failed to load frame: {frame_path}")
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
frames.append(frame)
# Stack and convert to tensor
frames = np.stack(frames) # [T, H, W, C]
frames = torch.from_numpy(frames).permute(0, 3, 1,
2).float() # [T, C, H, W]
return frames
def _detect_video_format(path: str | Path) -> str:
"""
Detect if path is a video file or frame directory.
Returns:
'video_file', 'frame_directory', or 'unknown'
"""
path = Path(path)
if path.is_file():
return 'video_file'
elif path.is_dir():
# Check if contains image files
image_extensions = ['.jpg', '.jpeg', '.png', '.bmp']
for ext in image_extensions:
if list(path.glob(f"*{ext}")):
return 'frame_directory'
return 'unknown'
else:
raise ValueError(f"Path does not exist: {path}")
def load_video_auto(video_path: str | Path,
num_frames: int | None = 16,
sample_strategy: str = 'uniform') -> torch.Tensor:
"""
Automatically detect format and load video.
Supports:
- Video files (MP4, AVI, MOV, MKV)
- Frame directories (JPG, PNG)
Args:
video_path: Path to video file or frame directory
num_frames: Number of frames to extract
sample_strategy: 'uniform' or 'random'
Returns:
video: [T, C, H, W]
"""
format_type = _detect_video_format(video_path)
if format_type == 'video_file':
return _load_video_cv2(video_path, num_frames, sample_strategy)
elif format_type == 'frame_directory':
return _load_video_from_frames(video_path, num_frames, sample_strategy)
else:
raise ValueError(f"Unknown video format at {video_path}")
def sample_clips_from_video(
video: torch.Tensor,
num_frames_per_clip: int = 16,
num_clips: int = 1,
strategy: str | ClipSamplingStrategy = ClipSamplingStrategy.BEGINNING,
frame_stride: int = 1,
temporal_stride: int = 1) -> list[torch.Tensor]:
"""
Sample clips from a video with various strategies.
Args:
video: [T, C, H, W] full video
num_frames_per_clip: Frames per clip
num_clips: Number of clips to extract
strategy: ClipSamplingStrategy or string ('beginning', 'random', etc.)
frame_stride: Skip frames (FPS control: 1=all, 2=every 2nd, 8=every 8th)
temporal_stride: Stride between clips for sliding window
Returns:
List of clips, each [num_frames_per_clip, C, H, W]
Examples:
>>> # Beginning clip (most common for FVD)
>>> clips = sample_clips_from_video(video, 16, strategy='beginning')
>>> # Multiple random clips
>>> clips = sample_clips_from_video(video, 16, num_clips=4, strategy='random')
>>> # Subsample FPS by 2x (every 2nd frame)
>>> clips = sample_clips_from_video(video, 16, frame_stride=2)
>>> # Sliding window with overlap
>>> clips = sample_clips_from_video(video, 16, strategy='sliding', temporal_stride=8)
"""
# Convert string to enum if needed
if isinstance(strategy, str):
strategy = ClipSamplingStrategy(strategy)
T, C, H, W = video.shape
# Apply frame stride (FPS subsampling)
if frame_stride > 1:
video = video[::frame_stride]
T = len(video)
effective_clip_length = num_frames_per_clip
# Handle videos shorter than clip length
if effective_clip_length > T:
pad_length = effective_clip_length - T
last_frame = video[-1:].repeat(pad_length, 1, 1, 1)
video = torch.cat([video, last_frame], dim=0)
T = len(video)
clips = []
if strategy == ClipSamplingStrategy.BEGINNING:
# Take first clip (most common for FVD evaluation)
clip = video[:effective_clip_length]
clips.append(clip)
elif strategy == ClipSamplingStrategy.MIDDLE:
# Take middle clip
start = (T - effective_clip_length) // 2
clip = video[start:start + effective_clip_length]
clips.append(clip)
elif strategy == ClipSamplingStrategy.RANDOM:
# Sample N random clips
for _ in range(num_clips):
if effective_clip_length == T:
start = 0
else:
start = np.random.randint(0, T - effective_clip_length + 1)
clip = video[start:start + effective_clip_length]
clips.append(clip)
elif strategy == ClipSamplingStrategy.UNIFORM:
# Uniformly spaced clips
if num_clips == 1:
# Single clip from middle
start = (T - effective_clip_length) // 2
clip = video[start:start + effective_clip_length]
clips.append(clip)
else:
# Multiple uniformly spaced clips
step = (T - effective_clip_length) / (num_clips -
1) if num_clips > 1 else 0
for i in range(num_clips):
start = int(i * step)
start = min(start, T - effective_clip_length)
clip = video[start:start + effective_clip_length]
clips.append(clip)
elif strategy == ClipSamplingStrategy.SLIDING:
# Sliding window with stride
for start in range(0, T - effective_clip_length + 1, temporal_stride):
clip = video[start:start + effective_clip_length]
clips.append(clip)
if len(clips) >= num_clips:
break
elif strategy == ClipSamplingStrategy.ALL:
# All possible clips (overlapping)
for start in range(T - effective_clip_length + 1):
clip = video[start:start + effective_clip_length]
clips.append(clip)
else:
raise ValueError(f"Unknown strategy: {strategy}")
return clips
def load_video_clips_streaming(directory: str | Path,
num_frames: int = 16,
max_videos: int | None = None,
clip_strategy: str
| ClipSamplingStrategy = 'beginning',
frame_stride: int = 1,
num_clips_per_video: int = 1,
video_extensions: list[str] | None = None,
support_frame_dirs: bool = True,
target_size: tuple[int, int] | None = (224, 224),
verbose: bool = True) -> Iterator[torch.Tensor]:
"""
This generator yields clips one-by-one instead of loading all videos into RAM.
Perfect for large datasets where memory is limited.
Args:
directory: Path to directory with videos
num_frames: Frames per clip
max_videos: Max videos to load
clip_strategy: 'beginning', 'random', 'uniform', etc.
frame_stride: Frame skip (1=all, 2=every 2nd, 8=every 8th)
num_clips_per_video: Number of clips per video
video_extensions: Video file extensions
support_frame_dirs: Also load frame directories
target_size: Resize clips to (H, W). If None, keep original size.
verbose: Show progress
Yields:
clip: [T, C, H, W] individual clips
Example:
>>> for clip in load_video_clips_streaming('data/videos/', num_frames=16):
>>> features = model.extract_features(clip.unsqueeze(0))
>>> # Process one clip at a time - low memory usage!
"""
if video_extensions is None:
video_extensions = ['.mp4', '.avi', '.mov', '.mkv']
directory = Path(directory)
if not directory.exists():
raise FileNotFoundError(f"Directory not found: {directory}")
# Find video paths
video_paths: list[Path] = []
# Find video files
for ext in video_extensions:
video_paths.extend(directory.glob(f"**/*{ext}"))
# Find frame directories if enabled
if support_frame_dirs:
for subdir in directory.iterdir():
if subdir.is_dir():
# Check if it contains frames
image_extensions = ['.jpg', '.jpeg', '.png', '.bmp']
for ext in image_extensions:
if list(subdir.glob(f"*{ext}")):
video_paths.append(subdir)
break
if len(video_paths) == 0:
raise ValueError(f"No videos found in {directory}")
video_paths = sorted(video_paths)
if max_videos is not None:
video_paths = video_paths[:max_videos]
if verbose:
print(f"Found {len(video_paths)} videos in {directory}")
if num_clips_per_video > 1:
print(f"Extracting {num_clips_per_video} clips per video...")
if frame_stride > 1:
print(f"Subsampling frames with stride {frame_stride}...")
if target_size:
print(f"Resizing clips to {target_size}...")
# Track statistics
failed_count = 0
total_clips = 0
iterator = tqdm(video_paths,
desc="Loading videos") if verbose else video_paths
for video_path in iterator:
try:
# Load full video
video = load_video_auto(video_path,
num_frames=None,
sample_strategy='uniform')
# Sample clips from video
clips = sample_clips_from_video(video,
num_frames_per_clip=num_frames,
num_clips=num_clips_per_video,
strategy=clip_strategy,
frame_stride=frame_stride)
if target_size is not None:
resized_clips = []
for clip in clips:
T, C, H, W = clip.shape
if target_size != (H, W):
# Resize to target size
clip = clip.contiguous(
) # Fix non-contiguous tensors first
clip_flat = clip.view(T * C, H,
W).unsqueeze(0) # [1, T*C, H, W]
clip_resized = torch.nn.functional.interpolate(
clip_flat,
size=target_size,
mode='bilinear',
align_corners=False)
clip = clip_resized.squeeze(0).view(
T, C, target_size[0],
target_size[1]) # Back to [T, C, H, W]
resized_clips.append(clip)
clips = resized_clips
# Yield clips one by one
for clip in clips:
yield clip
total_clips += 1
# Free memory
del video, clips
except Exception as e:
failed_count += 1
if verbose:
print(f"\nWarning: Failed to load {video_path}: {e}")
continue
# Validate
if total_clips == 0:
raise RuntimeError(f"Failed to load any videos from {directory}")
failure_rate = failed_count / len(video_paths)
if failure_rate > 0.1: # More than 10% failed
print(
f"\nWARNING: {failure_rate:.1%} of videos failed to load ({failed_count}/{len(video_paths)})"
)
if verbose:
print(
f"\nSuccessfully loaded {total_clips} clips from {len(video_paths) - failed_count} videos"
)
+7
View File
@@ -0,0 +1,7 @@
#!/bin/bash
# 1. Install missing dependency
pip install -q opencv-python-headless transformers huggingface_hub
# 2. Run FVD script
python benchmarks/fvd/run_fvd.py
+4
View File
@@ -0,0 +1,4 @@
#!/bin/bash
# 1. Install missing dependency
pip install -q opencv-python-headless
-24
View File
@@ -1,24 +0,0 @@
# Configuration for Cog ⚙️
# Reference: https://cog.run/yaml
build:
gpu: true
cuda: "12.1"
python_version: "3.10"
python_packages:
- "torch==2.4.0"
- "torchvision"
- "ninja==1.11.1.3"
- "transformers==4.46.1"
- "git+https://github.com/huggingface/diffusers.git@bf64b32652a63a1865a0528a73a13652b201698b"
- "accelerate==1.0.1"
- "safetensors==0.4.5"
- "peft==0.13.2"
- "packaging==24.2"
- "git+https://github.com/hao-ai-lab/FastVideo"
run:
- FLASH_ATTENTION_SKIP_CUDA_BUILD=TRUE pip install flash-attn --no-build-isolation
- curl -o /usr/local/bin/pget -L "https://github.com/replicate/pget/releases/latest/download/pget_$(uname -s)_$(uname -m)" && chmod +x /usr/local/bin/pget
predict: "predict.py:Predictor"
@@ -16,7 +16,7 @@ import sys
# Run it with `python collect_env.py` or `python -m torch.utils.collect_env`
from collections import namedtuple
from fastvideo.v1.envs import environment_variables
from fastvideo.envs import environment_variables
try:
import torch
@@ -62,6 +62,7 @@ SystemEnv = namedtuple(
DEFAULT_CONDA_PATTERNS = {
"torch",
"numpy",
"mypy"
"cudatoolkit",
"soumith",
"mkl",
+138
View File
@@ -0,0 +1,138 @@
# ComfyUI-FastVideo
A custom node suite for ComfyUI that provides accelerated video generation using [FastVideo](https://github.com/hao-ai-labs/FastVideo). See the [blog post](https://hao-ai-lab.github.io/blogs/fastvideo/) about FastVideo V1 to learn more.
## Multi-GPU Parallel Inference
One of the key features ComfyUI-FastVideo brings to ComfyUI is its ability to distribute the generation workload across multiple GPUs, resulting in significantly faster inference times.
![Wan2.1-I2V-14B-480P-Diffusers](./assets/wani2v.gif)
Example of Wan2.1-I2V-14B-480P-Diffusers model running on 4 GPUs.
## Features
- Generate high-quality videos from text prompts and images
- Configurable video parameters (prompt, resolution, frame count, FPS)
- Support for multiple GPUs with tensor and sequence parallelism
- Advanced configuration options for VAE, Text Encoder, and DIT components
- Interruption/cancellation support for long-running generations
## Installation
### Requirements
- [ComfyUI](https://github.com/comfyanonymous/ComfyUI)
- CUDA-capable GPU(s) with sufficient VRAM
### Install using ComfyUI Manager
Coming soon!
### Manual Installation
#### Copy the FastVideo `comfyui` directory into your ComfyUI custom_nodes directory:
```bash
cp -r /path/to/FastVideo/comfyui /path/to/ComfyUI/custom_nodes/FastVideo
```
#### Install dependencies:
Currently, the only dependency is `fastvideo`, which can be installed using pip.
```bash
pip install fastvideo
```
#### Install missing custom nodes:
`ComfyUI-VideoHelperSuite`:
```bash
cd /path/to/ComfyUI/custom_nodes
git clone https://github.com/Kosinkadink/ComfyUI-VideoHelperSuite.git
```
If you're seeing `ImportError: libGL.so.1: cannot open shared object file: No such file or directory`,
you may need to install ffmpeg
```bash
apt-get update && apt-get install ffmpeg
```
## Usage
After installation, the following nodes will be available in the ComfyUI interface under the "fastvideo" category:
- **Video Generator**: The main node for generating videos from prompts
- **Inference Args**: Configure video generation parameters
- **VAE Config**
- **Text Encoder Config**
- **DIT Config**
- **Load Image Path**: Load images for potential conditioning
You may have noticed many arguments on the nodes have 'auto' as the default value. This is because FastVideo will automatically detect the best values for these parameters based on the model and the hardware. However, you can also manually configure these parameters to get the best performance for your specific use case. We plan on releasing more optimized workflow files for different models and hardware configurations in the future.
You can see what some of the default configurations are by looking at the FastVideo repo:
- [Wan2.1-I2V-14B-480P-Diffusers](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/configs/wan_14B_i2v_480p_pipeline.json)
- [FastHunyuan-diffusers](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/configs/fasthunyuan_t2v.json)
### Node Configuration
#### Video Generator
- **prompt**: Text description of the video to generate
- **output_path**: Directory where generated videos will be saved
- **num_gpus**: Number of GPUs to use for generation
- **model_path**: Path to the FastVideo model
- **embedded_cfg_scale**: Classifier-free guidance scale
- **sp_size**: Sequence parallelism size (usually should match num_gpus)
- **tp_size**: Tensor parallelism size (usually should match num_gpus)
- **precision**: Model precision (fp16 or bf16)
`model_path takes either a model id from huggingface or a local path to a model. Models by default will be downloaded to ~/.cache/huggingface/hub/ and cached for subsequent runs.`
#### Inference Args
- **height/width**: Resolution of the output video
- **num_frames**: Number of frames to generate
- **num_inference_steps**: Number of diffusion steps per frame
- **guidance_scale**: Classifier-free guidance scale
- **flow_shift**: Frame flow shift parameter
- **seed**: Random seed for reproducible generation
- **fps**: Frames per second of the output video
- **image_path**: Optional path to input image for conditioning (for i2v models)
## Memory Management
Models will remain loaded in GPU memory between runs when you only change inference arguments (such as prompt, resolution, frame count, FPS, guidance scale, etc.) or the prompt text. This allows for faster subsequent generations since the model doesn't need to be reloaded.
However, if you need to change the following parameters, you will need to restart the ComfyUI server:
- **Number of GPUs** (`num_gpus`)
- **Model path** (`model_path`)
- **Tensor parallelism size** (`tp_size`)
- **Sequence parallelism size** (`sp_size`)
These parameters affect the model's distribution across GPUs and require a complete reinitialization of the model pipeline.
## Example workflows
### Text to Video
FastVideo-FastHunyuan-diffusers
![FastVideo-FastHunyuan-diffusers](./assets/fasthunyuan.png)
- [FastHunyuan-diffusers.json](./examples/FastHunyuan-diffusers.json)
### Image to Video
Wan2.1-I2V-14B-480P-Diffusers
![Wan2.1-I2V-14B-480P-Diffusers](./assets/wani2v.png)
- [Wan2.1-I2V-14B-480P-Diffusers.json](./examples/Wan2.1-I2V-14B-480P-Diffusers.json)
## License
This project is licensed under Apache 2.0.
+5
View File
@@ -0,0 +1,5 @@
from .video_generator.nodes import (NODE_CLASS_MAPPINGS,
NODE_DISPLAY_NAME_MAPPINGS)
WEB_DIRECTORY = "./web"
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY']
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.3 MiB

+6
View File
@@ -0,0 +1,6 @@
<svg width="160" height="93" viewBox="0 0 160 93" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M28.8511 91.66L57.6319 1.86368H64.5394L35.7585 91.66H28.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M15.0376 91.66L43.8185 1.86368H46.1209L17.3401 91.66H15.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
<path d="M1.22217 91.66L30.003 1.86366H31.1543L2.3734 91.66H1.22217Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
<path d="M71.4465 1.86483L42.666 91.6599H69.144L78.3538 58.2746H123.251L129.007 39.855H84.1099L89.866 22.5868H152.032L157.788 1.86483H71.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
</svg>

After

Width:  |  Height:  |  Size: 691 B

Binary file not shown.

After

Width:  |  Height:  |  Size: 8.7 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 769 KiB

+645
View File
@@ -0,0 +1,645 @@
{
"id": "23a1f065-bbba-4a8f-b144-944e1318fcbf",
"revision": 0,
"last_node_id": 8,
"last_link_id": 7,
"nodes": [
{
"id": 4,
"type": "VAEConfig",
"pos": [
374.2159423828125,
554.85888671875
],
"size": [
334.080078125,
322
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "vae_config",
"type": "VAE_CONFIG",
"links": [
2
]
}
],
"properties": {
"Node name for S&R": "VAEConfig"
},
"widgets_values": [
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999
],
"auto_widget_states": {
"load_encoder": {
"isAuto": true,
"value": -99999,
"cachedValue": true
},
"load_decoder": {
"isAuto": true,
"value": -99999,
"cachedValue": true
},
"tile_sample_min_height": {
"isAuto": true,
"value": -99999,
"cachedValue": 256
},
"tile_sample_min_width": {
"isAuto": true,
"value": -99999,
"cachedValue": 256
},
"tile_sample_min_num_frames": {
"isAuto": true,
"value": -99999,
"cachedValue": 16
},
"tile_sample_stride_height": {
"isAuto": true,
"value": -99999,
"cachedValue": 192
},
"tile_sample_stride_width": {
"isAuto": true,
"value": -99999,
"cachedValue": 192
},
"tile_sample_stride_num_frames": {
"isAuto": true,
"value": -99999,
"cachedValue": 12
},
"blend_num_frames": {
"isAuto": true,
"value": -99999,
"cachedValue": 0
},
"use_tiling": {
"isAuto": true,
"value": -99999,
"cachedValue": true
},
"use_temporal_tiling": {
"isAuto": true,
"value": -99999,
"cachedValue": true
},
"use_parallel_tiling": {
"isAuto": true,
"value": -99999,
"cachedValue": true
}
}
},
{
"id": 5,
"type": "TextEncoderConfig",
"pos": [
416.4937744140625,
953.6171875
],
"size": [
270,
106
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "text_encoder_config",
"type": "TEXT_ENCODER_CONFIG",
"links": [
7
]
}
],
"properties": {
"Node name for S&R": "TextEncoderConfig"
},
"widgets_values": [
-99999,
-99999,
-99999
],
"auto_widget_states": {
"prefix": {
"isAuto": true,
"value": -99999,
"cachedValue": ""
},
"quant_config": {
"isAuto": true,
"value": -99999,
"cachedValue": ""
},
"lora_config": {
"isAuto": true,
"value": -99999,
"cachedValue": ""
}
}
},
{
"id": 6,
"type": "DITConfig",
"pos": [
415.1928405761719,
1154.1573486328125
],
"size": [
270,
82
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "dit_config",
"type": "DIT_CONFIG",
"links": [
6
]
}
],
"properties": {
"Node name for S&R": "DITConfig"
},
"widgets_values": [
-99999,
-99999
],
"auto_widget_states": {
"prefix": {
"isAuto": true,
"value": -99999,
"cachedValue": ""
},
"quant_config": {
"isAuto": true,
"value": -99999,
"cachedValue": ""
}
}
},
{
"id": 1,
"type": "VideoGenerator",
"pos": [
818.804931640625,
348.9299621582031
],
"size": [
400,
436
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "inference_args",
"shape": 7,
"type": "INFERENCE_ARGS",
"link": 3
},
{
"name": "vae_config",
"shape": 7,
"type": "VAE_CONFIG",
"link": 2
},
{
"name": "text_encoder_config",
"shape": 7,
"type": "TEXT_ENCODER_CONFIG",
"link": 7
},
{
"name": "dit_config",
"shape": 7,
"type": "DIT_CONFIG",
"link": 6
}
],
"outputs": [
{
"name": "video_path",
"type": "STRING",
"links": [
4
]
}
],
"properties": {
"Node name for S&R": "VideoGenerator"
},
"widgets_values": [
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest. The playful yet serene atmosphere is complemented by soft natural light filtering through the petals. Mid-shot, warm and cheerful tones.",
"/workspace/ComfyUI/outputs_video/",
2,
"FastVideo/FastHunyuan-diffusers",
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999
],
"auto_widget_states": {
"embedded_cfg_scale": {
"isAuto": true,
"value": -99999,
"cachedValue": 6
},
"sp_size": {
"isAuto": true,
"value": -99999,
"cachedValue": 2
},
"tp_size": {
"isAuto": true,
"value": -99999,
"cachedValue": 2
},
"vae_precision": {
"isAuto": true,
"value": -99999,
"cachedValue": "fp16"
},
"vae_tiling": {
"isAuto": true,
"value": -99999,
"cachedValue": true
},
"vae_sp": {
"isAuto": true,
"value": -99999,
"cachedValue": true
},
"text_encoder_precision": {
"isAuto": true,
"value": -99999,
"cachedValue": "fp16"
},
"precision": {
"isAuto": true,
"value": -99999,
"cachedValue": "fp16"
},
"dit_cpu_offload": {
"isAuto": true,
"value": -99999,
"cachedValue": true
}
}
},
{
"id": 3,
"type": "VHS_LoadVideoPath",
"pos": [
1350.136962890625,
331.20361328125
],
"size": [
231.8896484375,
286
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "meta_batch",
"shape": 7,
"type": "VHS_BatchManager",
"link": null
},
{
"name": "vae",
"shape": 7,
"type": "VAE",
"link": null
},
{
"name": "video",
"type": "STRING",
"widget": {
"name": "video"
},
"link": 4
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
5
]
},
{
"name": "frame_count",
"type": "INT",
"links": null
},
{
"name": "audio",
"type": "AUDIO",
"links": null
},
{
"name": "video_info",
"type": "VHS_VIDEOINFO",
"links": null
}
],
"properties": {
"Node name for S&R": "VHS_LoadVideoPath"
},
"widgets_values": {
"video": "",
"force_rate": 0,
"custom_width": 0,
"custom_height": 0,
"frame_load_cap": 0,
"skip_first_frames": 0,
"select_every_nth": 1,
"format": "Wan",
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "",
"type": "path",
"format": "video/",
"force_rate": 0,
"custom_width": 0,
"custom_height": 0,
"frame_load_cap": 0,
"skip_first_frames": 0,
"select_every_nth": 1
}
}
}
},
{
"id": 2,
"type": "InferenceArgs",
"pos": [
411.46307373046875,
178.18182373046875
],
"size": [
278.73828125,
298
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "inference_args",
"type": "INFERENCE_ARGS",
"links": [
3
]
}
],
"properties": {
"Node name for S&R": "InferenceArgs"
},
"widgets_values": [
720,
1280,
45,
6,
-99999,
-99999,
1025,
"fixed",
24,
-99999,
-99999
],
"auto_widget_states": {
"height": {
"isAuto": false,
"value": 720,
"cachedValue": 720
},
"width": {
"isAuto": false,
"value": 1280,
"cachedValue": 1280
},
"num_frames": {
"isAuto": false,
"value": 45,
"cachedValue": 45
},
"num_inference_steps": {
"isAuto": false,
"value": 6,
"cachedValue": 6
},
"guidance_scale": {
"isAuto": true,
"value": -99999,
"cachedValue": 1
},
"flow_shift": {
"isAuto": true,
"value": -99999,
"cachedValue": 17
},
"seed": {
"isAuto": false,
"value": 1025,
"cachedValue": 1024
},
"fps": {
"isAuto": false,
"value": 24,
"cachedValue": 24
},
"image_path": {
"isAuto": true,
"value": -99999,
"cachedValue": "X://insert/path/here.mp4"
},
"enable_teacache": {
"isAuto": true,
"value": -99999,
"cachedValue": true
}
}
},
{
"id": 8,
"type": "VHS_VideoCombine",
"pos": [
1668.3499755859375,
328.22625732421875
],
"size": [
507.507080078125,
622.2227172851562
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 5
},
{
"name": "audio",
"shape": 7,
"type": "AUDIO",
"link": null
},
{
"name": "meta_batch",
"shape": 7,
"type": "VHS_BatchManager",
"link": null
},
{
"name": "vae",
"shape": 7,
"type": "VAE",
"link": null
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 24,
"loop_count": 0,
"filename_prefix": "",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 19,
"save_metadata": true,
"trim_to_audio": false,
"pingpong": false,
"save_output": false,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "._00003.mp4",
"subfolder": "",
"type": "temp",
"format": "video/h264-mp4",
"frame_rate": 24,
"workflow": "._00003.png",
"fullpath": "/workspace/ComfyUI/temp/._00003.mp4"
}
}
}
}
],
"links": [
[
2,
4,
0,
1,
1,
"VAE_CONFIG"
],
[
3,
2,
0,
1,
0,
"INFERENCE_ARGS"
],
[
4,
1,
0,
3,
2,
"STRING"
],
[
5,
3,
0,
8,
0,
"IMAGE"
],
[
6,
6,
0,
1,
3,
"DIT_CONFIG"
],
[
7,
5,
0,
1,
2,
"TEXT_ENCODER_CONFIG"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.9090909090909091,
"offset": [
112.86678372727341,
-71.45635903989245
]
},
"frontendVersion": "1.20.4",
"VHS_latentpreview": false,
"VHS_latentpreviewrate": 0,
"VHS_MetadataImage": true,
"VHS_KeepIntermediate": true
},
"version": 0.4
}
@@ -0,0 +1,697 @@
{
"id": "23a1f065-bbba-4a8f-b144-944e1318fcbf",
"revision": 0,
"last_node_id": 8,
"last_link_id": 7,
"nodes": [
{
"id": 7,
"type": "LoadImagePath",
"pos": [
33.15385437011719,
191.2037353515625
],
"size": [
270,
334
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "image_path",
"type": "STRING",
"links": [
1
]
},
{
"name": "IMAGE",
"type": "IMAGE",
"links": null
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"Node name for S&R": "LoadImagePath"
},
"widgets_values": [
"woman.jpg",
"image"
]
},
{
"id": 3,
"type": "VHS_LoadVideoPath",
"pos": [
1350.136962890625,
331.20361328125
],
"size": [
231.8896484375,
286
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "meta_batch",
"shape": 7,
"type": "VHS_BatchManager",
"link": null
},
{
"name": "vae",
"shape": 7,
"type": "VAE",
"link": null
},
{
"name": "video",
"type": "STRING",
"widget": {
"name": "video"
},
"link": 4
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
5
]
},
{
"name": "frame_count",
"type": "INT",
"links": null
},
{
"name": "audio",
"type": "AUDIO",
"links": null
},
{
"name": "video_info",
"type": "VHS_VIDEOINFO",
"links": null
}
],
"properties": {
"Node name for S&R": "VHS_LoadVideoPath"
},
"widgets_values": {
"video": "",
"force_rate": 0,
"custom_width": 0,
"custom_height": 0,
"frame_load_cap": 0,
"skip_first_frames": 0,
"select_every_nth": 1,
"format": "Wan",
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "",
"type": "path",
"format": "video/",
"force_rate": 0,
"custom_width": 0,
"custom_height": 0,
"frame_load_cap": 0,
"skip_first_frames": 0,
"select_every_nth": 1
}
}
}
},
{
"id": 8,
"type": "VHS_VideoCombine",
"pos": [
1668.3499755859375,
328.22625732421875
],
"size": [
214.7587890625,
334
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 5
},
{
"name": "audio",
"shape": 7,
"type": "AUDIO",
"link": null
},
{
"name": "meta_batch",
"shape": 7,
"type": "VHS_BatchManager",
"link": null
},
{
"name": "vae",
"shape": 7,
"type": "VAE",
"link": null
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 24,
"loop_count": 0,
"filename_prefix": "",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 19,
"save_metadata": true,
"trim_to_audio": false,
"pingpong": false,
"save_output": false,
"videopreview": {
"hidden": false,
"paused": false,
"params": {}
}
}
},
{
"id": 4,
"type": "VAEConfig",
"pos": [
374.2159423828125,
554.85888671875
],
"size": [
334.080078125,
322
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "vae_config",
"type": "VAE_CONFIG",
"links": [
2
]
}
],
"properties": {
"Node name for S&R": "VAEConfig"
},
"widgets_values": [
true,
true,
256,
256,
16,
192,
192,
12,
0,
true,
true,
true
],
"auto_widget_states": {
"load_encoder": {
"isAuto": true,
"value": true,
"cachedValue": true
},
"load_decoder": {
"isAuto": true,
"value": true,
"cachedValue": true
},
"tile_sample_min_height": {
"isAuto": true,
"value": 256,
"cachedValue": 256
},
"tile_sample_min_width": {
"isAuto": true,
"value": 256,
"cachedValue": 256
},
"tile_sample_min_num_frames": {
"isAuto": true,
"value": 16,
"cachedValue": 16
},
"tile_sample_stride_height": {
"isAuto": true,
"value": 192,
"cachedValue": 192
},
"tile_sample_stride_width": {
"isAuto": true,
"value": 192,
"cachedValue": 192
},
"tile_sample_stride_num_frames": {
"isAuto": true,
"value": 12,
"cachedValue": 12
},
"blend_num_frames": {
"isAuto": true,
"value": 0,
"cachedValue": 0
},
"use_tiling": {
"isAuto": true,
"value": true,
"cachedValue": true
},
"use_temporal_tiling": {
"isAuto": true,
"value": true,
"cachedValue": true
},
"use_parallel_tiling": {
"isAuto": true,
"value": true,
"cachedValue": true
}
}
},
{
"id": 2,
"type": "InferenceArgs",
"pos": [
411.46307373046875,
178.18182373046875
],
"size": [
278.73828125,
298
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "image_path",
"shape": 7,
"type": "STRING",
"widget": {
"name": "image_path"
},
"link": 1
}
],
"outputs": [
{
"name": "inference_args",
"type": "INFERENCE_ARGS",
"links": [
3
]
}
],
"properties": {
"Node name for S&R": "InferenceArgs"
},
"widgets_values": [
832,
480,
45,
20,
1,
17,
1024,
"fixed",
24,
"X://insert/path/here.mp4",
true
],
"auto_widget_states": {
"height": {
"isAuto": false,
"value": 832,
"cachedValue": 720
},
"width": {
"isAuto": false,
"value": 480,
"cachedValue": 1280
},
"num_frames": {
"isAuto": false,
"value": 45,
"cachedValue": 45
},
"num_inference_steps": {
"isAuto": false,
"value": 20,
"cachedValue": 6
},
"guidance_scale": {
"isAuto": true,
"value": 1,
"cachedValue": 1
},
"flow_shift": {
"isAuto": true,
"value": 17,
"cachedValue": 17
},
"seed": {
"isAuto": false,
"value": 1024,
"cachedValue": 1024
},
"fps": {
"isAuto": false,
"value": 24,
"cachedValue": 24
},
"image_path": {
"isAuto": true,
"value": "X://insert/path/here.mp4",
"cachedValue": "X://insert/path/here.mp4"
},
"enable_teacache": {
"isAuto": true,
"value": true,
"cachedValue": true
}
}
},
{
"id": 1,
"type": "VideoGenerator",
"pos": [
818.804931640625,
348.9299621582031
],
"size": [
400,
436
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "inference_args",
"shape": 7,
"type": "INFERENCE_ARGS",
"link": 3
},
{
"name": "vae_config",
"shape": 7,
"type": "VAE_CONFIG",
"link": 2
},
{
"name": "text_encoder_config",
"shape": 7,
"type": "TEXT_ENCODER_CONFIG",
"link": 7
},
{
"name": "dit_config",
"shape": 7,
"type": "DIT_CONFIG",
"link": 6
}
],
"outputs": [
{
"name": "video_path",
"type": "STRING",
"links": [
4
]
}
],
"properties": {
"Node name for S&R": "VideoGenerator"
},
"widgets_values": [
"A woman crying from laughter.",
"/workspace/ComfyUI/outputs_video/",
4,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
6,
2,
2,
"fp16",
true,
true,
"fp16",
"fp16",
true
],
"auto_widget_states": {
"embedded_cfg_scale": {
"isAuto": true,
"value": 6,
"cachedValue": 6
},
"sp_size": {
"isAuto": true,
"value": 2,
"cachedValue": 2
},
"tp_size": {
"isAuto": true,
"value": 2,
"cachedValue": 2
},
"vae_precision": {
"isAuto": true,
"value": "fp16",
"cachedValue": "fp16"
},
"vae_tiling": {
"isAuto": true,
"value": true,
"cachedValue": true
},
"vae_sp": {
"isAuto": true,
"value": true,
"cachedValue": true
},
"text_encoder_precision": {
"isAuto": true,
"value": "fp16",
"cachedValue": "fp16"
},
"precision": {
"isAuto": true,
"value": "fp16",
"cachedValue": "fp16"
},
"dit_cpu_offload": {
"isAuto": true,
"value": true,
"cachedValue": true
}
}
},
{
"id": 5,
"type": "TextEncoderConfig",
"pos": [
416.4937744140625,
953.6171875
],
"size": [
270,
106
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "text_encoder_config",
"type": "TEXT_ENCODER_CONFIG",
"links": [
7
]
}
],
"properties": {
"Node name for S&R": "TextEncoderConfig"
},
"widgets_values": [
"",
"",
""
],
"auto_widget_states": {
"prefix": {
"isAuto": true,
"value": "",
"cachedValue": ""
},
"quant_config": {
"isAuto": true,
"value": "",
"cachedValue": ""
},
"lora_config": {
"isAuto": true,
"value": "",
"cachedValue": ""
}
}
},
{
"id": 6,
"type": "DITConfig",
"pos": [
415.1928405761719,
1154.1573486328125
],
"size": [
270,
82
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "dit_config",
"type": "DIT_CONFIG",
"links": [
6
]
}
],
"properties": {
"Node name for S&R": "DITConfig"
},
"widgets_values": [
"",
""
],
"auto_widget_states": {
"prefix": {
"isAuto": true,
"value": "",
"cachedValue": ""
},
"quant_config": {
"isAuto": true,
"value": "",
"cachedValue": ""
}
}
}
],
"links": [
[
1,
7,
0,
2,
0,
"STRING"
],
[
2,
4,
0,
1,
1,
"VAE_CONFIG"
],
[
3,
2,
0,
1,
0,
"INFERENCE_ARGS"
],
[
4,
1,
0,
3,
2,
"STRING"
],
[
5,
3,
0,
8,
0,
"IMAGE"
],
[
6,
6,
0,
1,
3,
"DIT_CONFIG"
],
[
7,
5,
0,
1,
2,
"TEXT_ENCODER_CONFIG"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.8264462809917354,
"offset": [
646.7950212991898,
66.17259910028655
]
},
"frontendVersion": "1.20.4",
"VHS_latentpreview": false,
"VHS_latentpreviewrate": 0,
"VHS_MetadataImage": true,
"VHS_KeepIntermediate": true
},
"version": 0.4
}
+31
View File
@@ -0,0 +1,31 @@
class DITConfig:
@classmethod
def INPUT_TYPES(cls):
return {
"optional": {
"prefix": ("STRING", {
"default": ""
}),
"quant_config": ("STRING", {
"default": ""
}),
}
}
@classmethod
def VALIDATE_INPUTS(cls, **kwargs):
return True
RETURN_TYPES = ("DIT_CONFIG", )
RETURN_NAMES = ("dit_config", )
FUNCTION = "set_args"
CATEGORY = "fastvideo"
def set_args(self, prefix, quant_config):
raw_args = {"prefix": prefix, "quant_config": quant_config}
# Filter out keys where value is -99999
args = {k: v for k, v in raw_args.items() if str(int(v)) != str(-99999)}
return (args, )
+89
View File
@@ -0,0 +1,89 @@
class InferenceArgs:
@classmethod
def INPUT_TYPES(cls):
return {
"optional": {
"height": ("INT", {
"default": 720
}),
"width": ("INT", {
"default": 1280
}),
"num_frames": ("INT", {
"default": 45
}),
"num_inference_steps": ("INT", {
"default": 6
}),
"guidance_scale": ("FLOAT", {
"default": 1.0
}),
"flow_shift": ("INT", {
"default": 17
}),
"seed": ("INT", {
"default": 1024
}),
"fps": ("INT", {
"default": 24
}),
"image_path": ("STRING", {
"default": "X://insert/path/here.mp4"
}),
"enable_teacache": ([True, False], {
"default": False
}),
}
}
@classmethod
def VALIDATE_INPUTS(cls, **kwargs):
return True
RETURN_TYPES = ("INFERENCE_ARGS", )
RETURN_NAMES = ("inference_args", )
FUNCTION = "set_args"
CATEGORY = "fastvideo"
def set_args(
self,
height,
width,
num_frames,
num_inference_steps,
guidance_scale,
flow_shift,
seed,
fps,
image_path,
enable_teacache,
):
raw_args = {
"height": height,
"width": width,
"num_frames": num_frames,
"num_inference_steps": num_inference_steps,
"guidance_scale": guidance_scale,
"flow_shift": flow_shift,
"seed": seed,
"fps": fps,
"image_path": image_path,
"enable_teacache": enable_teacache,
}
# Filter out keys where value is -99999, handling different types properly
args = {}
for k, v in raw_args.items():
try:
if isinstance(v, str):
if v != "-99999":
args[k] = v
elif v != -99999:
# If it's not a string, compare directly
args[k] = v
except (ValueError, TypeError):
# Include any value that causes an error in comparison
args[k] = v
return (args, )
+103
View File
@@ -0,0 +1,103 @@
import hashlib
import os
import folder_paths
import numpy as np
import torch
from PIL import Image, ImageOps, ImageSequence
from .node_helpers import pillow
class LoadImagePath:
@classmethod
def INPUT_TYPES(s):
input_dir = folder_paths.get_input_directory()
files = [
f for f in os.listdir(input_dir)
if os.path.isfile(os.path.join(input_dir, f))
]
files = folder_paths.filter_files_content_types(files, ["image"])
return {
"required": {
"image": (sorted(files), {
"image_upload": True
})
},
}
CATEGORY = "fastvideo"
RETURN_TYPES = ("STRING", "IMAGE", "MASK")
RETURN_NAMES = ("image_path", "IMAGE", "MASK")
FUNCTION = "load_image"
def load_image(self, image):
image_path = folder_paths.get_annotated_filepath(image)
img = pillow(Image.open, image_path)
output_images: list[torch.Tensor] = []
output_masks: list[torch.Tensor] = []
w, h = None, None
excluded_formats = ['MPO']
for i in ImageSequence.Iterator(img):
processed_image = pillow(ImageOps.exif_transpose, i)
if processed_image is None:
continue
if processed_image.mode == 'I':
processed_image = processed_image.point(lambda i: i * (1 / 255))
image = processed_image.convert("RGB")
if len(output_images) == 0:
w = image.size[0]
h = image.size[1]
if image.size[0] != w or image.size[1] != h:
continue
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[
None,
]
if 'A' in processed_image.getbands():
mask = np.array(processed_image.getchannel('A')).astype(
np.float32) / 255.0
mask = 1. - torch.from_numpy(mask)
elif processed_image.mode == 'P' and 'transparency' in processed_image.info:
mask = np.array(
processed_image.convert('RGBA').getchannel('A')).astype(
np.float32) / 255.0
mask = 1. - torch.from_numpy(mask)
else:
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
output_images.append(image)
output_masks.append(mask.unsqueeze(0))
if len(output_images) > 1 and img.format not in excluded_formats:
output_image = torch.cat(output_images, dim=0)
output_mask = torch.cat(output_masks, dim=0)
else:
output_image = output_images[0]
output_mask = output_masks[0]
return (image_path, output_image, output_mask)
@classmethod
def IS_CHANGED(s, image):
image_path = folder_paths.get_annotated_filepath(image)
m = hashlib.sha256()
with open(image_path, 'rb') as f:
m.update(f.read())
return m.digest().hex()
@classmethod
def VALIDATE_INPUTS(s, image):
if not folder_paths.exists_annotated_filepath(image):
return "Invalid image file: {}".format(image)
return True
+68
View File
@@ -0,0 +1,68 @@
import hashlib
from collections.abc import Callable
from typing import Any, TypeVar
import torch
from comfy.cli_args import args
from PIL import ImageFile, UnidentifiedImageError
T = TypeVar('T')
def conditioning_set_values(conditioning: list[Any],
values: dict[str, Any] | None = None) -> list[Any]:
if values is None:
values = {}
c = []
for t in conditioning:
n = [t[0], t[1].copy()]
for k in values:
n[1][k] = values[k]
c.append(n)
return c
def pillow(fn: Callable[[Any], T], arg: Any) -> T:
prev_value = None
try:
x = fn(arg)
except (OSError, UnidentifiedImageError, ValueError
): #PIL issues #4472 and #2445, also fixes ComfyUI issue #3416
prev_value = ImageFile.LOAD_TRUNCATED_IMAGES
ImageFile.LOAD_TRUNCATED_IMAGES = True
x = fn(arg)
finally:
if prev_value is not None:
ImageFile.LOAD_TRUNCATED_IMAGES = prev_value
return x
def hasher() -> Callable[[], Any]:
hashfuncs = {
"md5": hashlib.md5,
"sha1": hashlib.sha1,
"sha256": hashlib.sha256,
"sha512": hashlib.sha512
}
return hashfuncs[args.default_hashing_function]
def string_to_torch_dtype(string: str) -> torch.dtype | None:
if string == "fp32":
return torch.float32
if string == "fp16":
return torch.float16
if string == "bf16":
return torch.bfloat16
return None
def image_alpha_fix(destination: torch.Tensor,
source: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
if destination.shape[-1] < source.shape[-1]:
source = source[..., :destination.shape[-1]]
elif destination.shape[-1] > source.shape[-1]:
destination = torch.nn.functional.pad(destination, (0, 1))
destination[..., -1] = 1.0
return destination, source
+24
View File
@@ -0,0 +1,24 @@
from .dit_config import DITConfig
from .inference_args import InferenceArgs
from .load_image import LoadImagePath
from .text_encoder_config import TextEncoderConfig
from .vae_config import VAEConfig
from .video_generator import VideoGenerator
NODE_CLASS_MAPPINGS = {
"VideoGenerator": VideoGenerator,
"InferenceArgs": InferenceArgs,
"VAEConfig": VAEConfig,
"TextEncoderConfig": TextEncoderConfig,
"DITConfig": DITConfig,
"LoadImagePath": LoadImagePath
}
NODE_DISPLAY_NAME_MAPPINGS = {
"VideoGenerator": "Video Generator",
"InferenceArgs": "Inference Args",
"VAEConfig": "VAE Config",
"TextEncoderConfig": "Text Encoder Config",
"DITConfig": "DIT Config",
"LoadImagePath": "Load Image Path"
}
@@ -0,0 +1,38 @@
class TextEncoderConfig:
@classmethod
def INPUT_TYPES(cls):
return {
"optional": {
"prefix": ("STRING", {
"default": ""
}),
"quant_config": ("STRING", {
"default": ""
}),
"lora_config": ("STRING", {
"default": ""
}),
}
}
@classmethod
def VALIDATE_INPUTS(cls, **kwargs):
return True
RETURN_TYPES = ("TEXT_ENCODER_CONFIG", )
RETURN_NAMES = ("text_encoder_config", )
FUNCTION = "set_args"
CATEGORY = "fastvideo"
def set_args(self, prefix, quant_config, lora_config):
raw_args = {
"prefix": prefix,
"quant_config": quant_config,
"lora_config": lora_config
}
# Filter out keys where value is -99999
args = {k: v for k, v in raw_args.items() if str(int(v)) != str(-99999)}
return (args, )
+88
View File
@@ -0,0 +1,88 @@
class VAEConfig:
@classmethod
def INPUT_TYPES(cls):
return {
"optional": {
"load_encoder": ([True, False], {
"default": True
}),
"load_decoder": ([True, False], {
"default": True
}),
"tile_sample_min_height": ("INT", {
"default": 256
}),
"tile_sample_min_width": ("INT", {
"default": 256
}),
"tile_sample_min_num_frames": ("INT", {
"default": 16
}),
"tile_sample_stride_height": ("INT", {
"default": 192
}),
"tile_sample_stride_width": ("INT", {
"default": 192
}),
"tile_sample_stride_num_frames": ("INT", {
"default": 12
}),
"blend_num_frames": ("INT", {
"default": 0
}),
"use_tiling": ([True, False], {
"default": True
}),
"use_temporal_tiling": ([True, False], {
"default": True
}),
"use_parallel_tiling": ([True, False], {
"default": True
}),
}
}
@classmethod
def VALIDATE_INPUTS(cls, **kwargs):
return True
RETURN_TYPES = ("VAE_CONFIG", )
RETURN_NAMES = ("vae_config", )
FUNCTION = "set_args"
CATEGORY = "fastvideo"
def set_args(
self,
load_encoder,
load_decoder,
tile_sample_min_height,
tile_sample_min_width,
tile_sample_min_num_frames,
tile_sample_stride_height,
tile_sample_stride_width,
tile_sample_stride_num_frames,
blend_num_frames,
use_tiling,
use_temporal_tiling,
use_parallel_tiling,
):
raw_args = {
"load_encoder": load_encoder,
"load_decoder": load_decoder,
"tile_sample_min_height": tile_sample_min_height,
"tile_sample_min_width": tile_sample_min_width,
"tile_sample_min_num_frames": tile_sample_min_num_frames,
"tile_sample_stride_height": tile_sample_stride_height,
"tile_sample_stride_width": tile_sample_stride_width,
"tile_sample_stride_num_frames": tile_sample_stride_num_frames,
"blend_num_frames": blend_num_frames,
"use_tiling": use_tiling,
"use_temporal_tiling": use_temporal_tiling,
"use_parallel_tiling": use_parallel_tiling,
}
# Filter out any value explicitly set to -99999
args = {k: v for k, v in raw_args.items() if str(int(v)) != str(-99999)}
return (args, )
+315
View File
@@ -0,0 +1,315 @@
from __future__ import annotations
import glob
import os
import signal
import sys
import threading
import time
from typing import Any
from comfy.model_management import processing_interrupted
from fastvideo import PipelineConfig
from fastvideo import VideoGenerator as FastVideoGenerator
sys.path.insert(
0,
os.path.dirname(
os.path.dirname(
os.path.dirname(os.path.dirname(os.path.abspath(__file__))))))
# Custom exception for interruption
class GenerationInterruptedException(Exception):
pass
# Custom exception for interruption that ComfyUI will recognize
class GenerationCancelledException(Exception):
def __init__(self,
message: str = "Generation was cancelled by user") -> None:
self.message = message
super().__init__(self.message)
def update_config_from_args(config: Any, args_dict: dict[str, Any]) -> None:
"""
Update configuration object from arguments dictionary.
Args:
config: The configuration object to update
args_dict: Dictionary containing arguments
"""
for key, value in args_dict.items():
if hasattr(config, key) and value is not None:
if key == "text_encoder_precisions" and isinstance(value, list):
setattr(config, key, tuple(value))
else:
setattr(config, key, value)
class VideoGenerator:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"prompt": ("STRING", {
"multiline":
True,
"default":
"A ripe orange tumbles gently from a tree and lands on the head of a lounging capybara, "
"who blinks slowly in response. The moment is quietly humorous and oddly serene, framed by "
"lush green foliage and dappled sunlight. Mid-shot, warm and whimsical tones."
}),
"output_path": ("STRING", {
"default": "/workspace/ComfyUI/outputs_video/"
}),
"num_gpus": ("INT", {
"default": 2,
"min": 1,
"max": 16
}),
"model_path": ("STRING", {
"default": "FastVideo/FastHunyuan-diffusers"
})
},
"optional": {
"inference_args": ("INFERENCE_ARGS", ),
"embedded_cfg_scale": ("FLOAT", {
"default": 6.0
}),
"sp_size": ("INT", {
"default": 2
}),
"tp_size": ("INT", {
"default": 2
}),
"vae_config": ("VAE_CONFIG", ),
"vae_precision": (["fp16", "bf16"], {
"default": "fp16"
}),
"vae_tiling": ([True, False], {
"default": True
}),
"vae_sp": ([True, False], {
"default": False
}),
"text_encoder_config": ("TEXT_ENCODER_CONFIG", ),
"text_encoder_precision": (["fp16", "bf16"], {
"default": "fp16"
}),
"dit_config": ("DIT_CONFIG", ),
"precision": (["fp16", "bf16"], {
"default": "fp16"
}),
"dit_cpu_offload": ([True, False], {
"default": False
}),
}
}
@classmethod
def VALIDATE_INPUTS(cls, **kwargs):
return True
RETURN_TYPES = ("STRING", )
RETURN_NAMES = ("video_path", )
FUNCTION = "launch_inference"
CATEGORY = "fastvideo"
generator: FastVideoGenerator | None = None
_interrupt_thread: threading.Thread | None = None
_generation_active: bool = False
_generation_interrupted: bool = False
_interrupt_event: threading.Event = threading.Event()
_generation_thread: threading.Thread | None = None
_generation_result: str | None = None
_generation_exception: Exception | None = None
def _monitor_for_interruption(self):
"""Background thread that monitors for interruption requests"""
time.sleep(2) # Give the generation thread time to send execute_forward
while self._generation_active and not self._interrupt_event.is_set():
if processing_interrupted():
print("Video generation interrupted by user")
self._generation_interrupted = True
# Try to send interrupt signal to worker processes
if self.generator is not None and hasattr(
self.generator, 'executor'):
try:
# The MultiprocExecutor has a workers attribute
if hasattr(self.generator.executor, 'workers'):
for worker in self.generator.executor.workers:
if worker.is_alive():
os.kill(worker.pid, signal.SIGINT)
print("Interrupt signal sent to worker processes")
except Exception as e:
print(f"Error sending interrupt signal: {e}")
# Set the interrupt event to notify other threads
self._interrupt_event.set()
break
time.sleep(0.5)
def _run_generation(self, prompt: str, output_path: str,
inference_args: dict[str, Any]) -> None:
"""Thread function to run the generation"""
try:
if self.generator is not None:
self.generator.generate_video(prompt=prompt,
output_path=output_path,
**inference_args)
self._generation_result = os.path.join(output_path,
f"{prompt[:100]}.mp4")
else:
raise RuntimeError("Generator is not initialized")
except Exception as e:
self._generation_exception = e
self._interrupt_event.set()
def load_output_video(self, output_dir):
video_extensions = ["*.mp4", "*.avi", "*.mov", "*.mkv"]
video_files = []
for ext in video_extensions:
video_files.extend(glob.glob(os.path.join(output_dir, ext)))
if not video_files:
print("No video files found in output directory: %s", output_dir)
return ""
video_files.sort()
return video_files[0]
def launch_inference(
self,
prompt,
output_path,
num_gpus,
model_path,
embedded_cfg_scale,
sp_size,
tp_size,
vae_precision,
vae_tiling,
vae_sp,
text_encoder_precision,
precision,
inference_args=None,
vae_config=None,
text_encoder_config=None,
dit_config=None,
dit_cpu_offload=None,
):
print('Running FastVideo inference')
# Reset interruption flag and event
self._generation_interrupted = False
self._interrupt_event.clear()
self._generation_result = None
self._generation_exception = None
# Load pipeline config from model path
pipeline_config = PipelineConfig.from_pretrained(model_path)
print('pipeline_config', pipeline_config)
# Update configs with provided config dictionaries
if dit_config is not None:
update_config_from_args(pipeline_config.dit_config, dit_config)
if vae_config is not None:
update_config_from_args(pipeline_config.vae_config, vae_config)
if text_encoder_config is not None:
update_config_from_args(pipeline_config.text_encoder_configs,
text_encoder_config)
# Update top-level pipeline config with remaining arguments
raw_pipeline_args = {}
if embedded_cfg_scale is not None:
raw_pipeline_args['embedded_cfg_scale'] = embedded_cfg_scale
if precision is not None:
raw_pipeline_args['precision'] = precision
if vae_precision is not None:
raw_pipeline_args['vae_precision'] = vae_precision
if vae_tiling is not None:
raw_pipeline_args['vae_tiling'] = vae_tiling
if vae_sp is not None:
raw_pipeline_args['vae_sp'] = vae_sp
if text_encoder_precision is not None:
raw_pipeline_args['text_encoder_precision'] = text_encoder_precision
# Filter out any value explicitly set to -99999 (auto values)
pipeline_args = {
k: v
for k, v in raw_pipeline_args.items() if str(int(v)) != str(-99999)
}
update_config_from_args(pipeline_config, pipeline_args)
raw_generation_args = {}
if num_gpus is not None:
raw_generation_args['num_gpus'] = num_gpus
if tp_size is not None:
raw_generation_args['tp_size'] = tp_size
if sp_size is not None:
raw_generation_args['sp_size'] = sp_size
if dit_cpu_offload is not None:
raw_generation_args['dit_cpu_offload'] = dit_cpu_offload
generation_args = {
k: v
for k, v in raw_generation_args.items()
if str(int(v)) != str(-99999)
}
if self.generator is None:
print('generation_args', generation_args)
print('pipeline_config', pipeline_config)
self.generator = FastVideoGenerator.from_pretrained(
model_path=model_path,
**generation_args,
pipeline_config=pipeline_config)
print('inference_args', inference_args)
# Start a thread to run the generation
self._generation_thread = threading.Thread(target=self._run_generation,
args=(prompt, output_path,
inference_args),
daemon=True)
self._generation_thread.start()
# Start a background thread to monitor for interruptions
self._generation_active = True
self._interrupt_thread = threading.Thread(
target=self._monitor_for_interruption, daemon=True)
self._interrupt_thread.start()
# Wait for either completion or interruption
while self._generation_thread.is_alive(
) and not self._interrupt_event.is_set():
self._generation_thread.join(timeout=0.5)
self._generation_active = False
if self._interrupt_thread:
self._interrupt_thread.join(timeout=1.0)
self._interrupt_thread = None
if self._generation_interrupted:
print("Video generation was cancelled by user")
raise GenerationCancelledException()
elif self._generation_exception:
# Re-raise the exception from the generation thread
raise self._generation_exception
elif self._generation_result:
return (self._generation_result, )
else:
# This shouldn't happen, but just in case
print("Generation completed but no result was produced")
raise Exception("Generation failed to produce a result")
+593
View File
@@ -0,0 +1,593 @@
import { app } from '../../../scripts/app.js'
function chainCallback(object, property, callback) {
if (object == undefined) {
console.error("Tried to add callback to non-existent object");
return;
}
if (property in object && object[property]) {
const callback_orig = object[property];
object[property] = function () {
const r = callback_orig.apply(this, arguments);
return callback.apply(this, arguments) ?? r;
};
} else {
object[property] = callback;
}
}
function drawAutoAnnotated(ctx, node, widget_width, y, H) {
const litegraph_base = LiteGraph;
const show_text = app.canvas.ds.scale >= 0.5;
const margin = 15;
const autoTextWidth = 30;
const autoTextRightMargin = 5;
ctx.textAlign = 'left';
ctx.strokeStyle = litegraph_base.WIDGET_OUTLINE_COLOR;
ctx.fillStyle = litegraph_base.WIDGET_BGCOLOR;
ctx.beginPath();
if (show_text && ctx.roundRect) {
ctx.roundRect(margin, y, widget_width - margin * 2, H, [H * 0.5]);
} else {
ctx.rect(margin, y, widget_width - margin * 2, H);
}
ctx.fill();
if (show_text) {
if (!this.disabled) ctx.stroke();
const isAuto = this.isAuto === true;
ctx.save();
if (isAuto) {
ctx.fillStyle = litegraph_base.WIDGET_TEXT_COLOR;
ctx.strokeStyle = litegraph_base.WIDGET_TEXT_COLOR;
} else {
ctx.fillStyle = litegraph_base.WIDGET_SECONDARY_TEXT_COLOR;
ctx.strokeStyle = litegraph_base.WIDGET_SECONDARY_TEXT_COLOR;
}
// Position for the cog
const cogX = widget_width - autoTextRightMargin - autoTextWidth - 6;
const cogY = y + H * 0.5;
const cogRadius = 6;
const toothLength = 2;
const numTeeth = 8;
const holeRadius = 2; // Radius of the center hole
// Draw the cog
ctx.beginPath();
ctx.arc(cogX, cogY, cogRadius - toothLength, 0, Math.PI * 2);
ctx.fill();
// Draw the center hole (by clearing it)
ctx.beginPath();
ctx.arc(cogX, cogY, holeRadius, 0, Math.PI * 2);
ctx.fillStyle = litegraph_base.WIDGET_BGCOLOR;
ctx.fill();
// Reset fill style for the teeth
if (isAuto) {
ctx.fillStyle = litegraph_base.WIDGET_TEXT_COLOR;
} else {
ctx.fillStyle = litegraph_base.WIDGET_SECONDARY_TEXT_COLOR;
}
// Draw teeth
ctx.beginPath();
for (let i = 0; i < numTeeth; i++) {
const angle = (i / numTeeth) * Math.PI * 2;
const innerX = cogX + (cogRadius - toothLength) * Math.cos(angle);
const innerY = cogY + (cogRadius - toothLength) * Math.sin(angle);
const outerX = cogX + cogRadius * Math.cos(angle);
const outerY = cogY + cogRadius * Math.sin(angle);
ctx.moveTo(innerX, innerY);
ctx.lineTo(outerX, outerY);
}
ctx.lineWidth = 2;
ctx.stroke();
ctx.restore();
// Draw label
ctx.fillStyle = litegraph_base.WIDGET_SECONDARY_TEXT_COLOR;
const label = this.label || this.name;
if (label != null) {
ctx.fillText(label, margin * 2 + 5, y + H * 0.7);
}
// Draw value
ctx.textAlign = 'right';
const text = isAuto ? "auto" : this.displayValue();
ctx.fillStyle = isAuto ? litegraph_base.WIDGET_SECONDARY_TEXT_COLOR : litegraph_base.WIDGET_TEXT_COLOR;
ctx.fillText(text, widget_width - autoTextRightMargin - autoTextWidth - 15, y + H * 0.7);
// Draw increment/decrement buttons if not in AUTO mode and not a string widget
if (!isAuto && !this.disabled && this.config[0] !== "FVAUTOSTRING") {
// Draw decrement button (left triangle)
ctx.fillStyle = litegraph_base.WIDGET_TEXT_COLOR;
ctx.beginPath();
ctx.moveTo(margin + 16, y + 5);
ctx.lineTo(margin + 6, y + H * 0.5);
ctx.lineTo(margin + 16, y + H - 5);
ctx.fill();
// Draw increment button (right triangle)
ctx.beginPath();
ctx.moveTo(widget_width - margin - 16, y + 5);
ctx.lineTo(widget_width - margin - 6, y + H * 0.5);
ctx.lineTo(widget_width - margin - 16, y + H - 5);
ctx.fill();
}
}
}
function mouseAutoAnnotated(event, [x, y], node) {
const widget_width = node.size[0];
const margin = 15;
const H = 20; // Widget height
const autoTextWidth = 30;
const autoTextRightMargin = 5;
const cogRadius = 6;
if (this.isAuto) {
if (event.type === "pointerup" || event.type === "mouseup") {
const cogX = widget_width - autoTextRightMargin - autoTextWidth - 6;
const cogLeftEdge = cogX - cogRadius;
const cogRightEdge = cogX + cogRadius;
if (x > cogLeftEdge && x < cogRightEdge) {
this.isAuto = false;
this.value = this.cachedValue !== undefined ? this.cachedValue : (this.options.default || 0);
if (this.callback) {
this.callback(this.value);
}
node.graph.setDirtyCanvas(true, false);
}
}
// Block ALL events in auto mode except cog clicks
event.preventDefault?.();
event.stopPropagation?.();
event.stopImmediatePropagation?.();
return true; // Always return true to indicate event was handled
}
// Determine if clicking on increment/decrement buttons
const delta = this.config[0] === "FVAUTOSTRING" ? 0 :
(x < 40 ? -1 : x > widget_width - 48 ? 1 : 0);
if (event.type === "pointerdown" || event.type === "mousedown") {
// ComfyUI appears to intercept pointerdown events, so this code path is never reached
console.log("pointerdown received (unexpected)");
return false;
} else if (event.type === "pointerup" || event.type === "mouseup") {
// Stop event propagation to prevent double handling
event.preventDefault?.();
event.stopPropagation?.();
event.stopImmediatePropagation?.();
const cogX = widget_width - autoTextRightMargin - autoTextWidth - 6;
const cogLeftEdge = widget_width - autoTextRightMargin - autoTextWidth - 6 - cogRadius;
const cogRightEdge = widget_width - autoTextRightMargin - autoTextWidth - 6 + cogRadius;
if (x > cogLeftEdge && x < cogRightEdge) {
this.isAuto = !this.isAuto;
if (this.isAuto) {
this.cachedValue = this.value;
this.value = -99999;
} else {
this.value = this.cachedValue !== undefined ? this.cachedValue : (this.options.default || 0);
}
if (this.callback) {
this.callback(this.value);
}
node.graph.setDirtyCanvas(true, false);
return true;
}
// If in auto mode and NOT clicking the cog, block all other interactions
if (this.isAuto) {
return true;
}
// Handle increment/decrement buttons if not in auto mode
if (delta !== 0 && !this.isAuto) {
if (this.config[0] === "FVAUTOCOMBO") {
const options = this.options.values || [];
if (options.length === 0) return true;
let currentIndex = -1;
for (let i = 0; i < options.length; i++) {
const optValue = typeof options[i] === 'object' ? options[i].value : options[i];
if (optValue == this.value || String(optValue) === String(this.value)) {
currentIndex = i;
break;
}
}
if (currentIndex === -1) {
currentIndex = 0;
}
let newIndex = currentIndex + delta;
if (newIndex < 0) {
newIndex = options.length - 1;
} else if (newIndex >= options.length) {
newIndex = 0;
}
const newOption = options[newIndex];
this.value = typeof newOption === 'object' ? newOption.value : newOption;
if (this.callback) {
this.callback(this.value);
}
node.graph.setDirtyCanvas(true, false);
return true;
} else {
let v = parseFloat(this.value);
const increment = delta * 0.1 * (this.options.step || 1);
v += increment;
// Apply min/max constraints
if (this.options.min != null) {
v = Math.max(this.options.min, v);
}
if (this.options.max != null) {
v = Math.min(this.options.max, v);
}
// Round to precision or to integer
if (this.config[0] === "FVAUTOINT") {
v = Math.round(v);
} else if (this.options.precision !== undefined) {
const precision = Math.pow(10, this.options.precision);
v = Math.round(v * precision) / precision;
}
this.value = v;
if (this.callback) {
this.callback(this.value);
}
node.graph.setDirtyCanvas(true, false);
return true;
}
}
if (delta === 0 && !this.isAuto) {
if (this.config[0] === "FVAUTOCOMBO") {
const options = this.options.values || [];
// Create menu items
const menuItems = options.map(opt => {
const value = typeof opt === 'object' ? opt.value : opt;
const label = typeof opt === 'object' ? opt.label : opt.toString();
return {
content: label,
callback: () => {
this.value = value;
if (this.callback) {
this.callback(this.value);
}
node.graph.setDirtyCanvas(true, false);
}
};
});
new LiteGraph.ContextMenu(menuItems, {
event: event,
title: null,
callback: null,
extra: node
});
return true;
} else if (this.config[0] === "FVAUTOSTRING") {
const d_callback = (v) => {
this.value = v;
if (this.callback) {
this.callback(this.value);
}
node.graph.setDirtyCanvas(true, false);
};
const dialog = app.canvas.prompt(
'Value',
this.value,
d_callback,
event
);
return true;
} else {
// For numeric widgets, show input dialog
const d_callback = (v) => {
this.value = this.parseValue?.(v) ?? Number(v);
// Apply min/max constraints
if (this.options.min != null) {
this.value = Math.max(this.options.min, this.value);
}
if (this.options.max != null) {
this.value = Math.min(this.options.max, this.value);
}
// Round to precision or to integer
if (this.config[0] === "FVAUTOINT") {
this.value = Math.round(this.value);
} else if (this.options.precision !== undefined) {
const precision = Math.pow(10, this.options.precision);
this.value = Math.round(this.value * precision) / precision;
}
if (this.callback) {
this.callback(this.value);
}
node.graph.setDirtyCanvas(true, false);
};
const dialog = app.canvas.prompt(
'Value',
this.value,
d_callback,
event
);
return true;
}
}
return true;
}
return false;
}
function makeAutoAnnotated(widget, inputData) {
const original = {
callback: widget.callback,
type: widget.type,
value: widget.value
};
// Add AUTO properties to the widget
Object.assign(widget, {
type: "BOOLEAN",
draw: drawAutoAnnotated,
mouse: mouseAutoAnnotated,
onMouse: null, // Explicitly disable original onMouse handler
isAuto: true,
cachedValue: widget.value,
config: inputData,
options: Object.assign({}, inputData[1], widget.options),
original: original, // Store original properties for reference
// Disable other potential mouse handlers with no-op functions
onClick: function () {
return false;
},
onPointerUp: function () {
return false;
},
onPointerDown: function () {
return false;
},
onMouseUp: function () {
return false;
},
onMouseDown: function () {
return false;
},
computeSize(width) {
return [width, 20];
},
displayValue: function () {
if (this.config[0] === "FVAUTOINT") {
return Math.round(this.value).toString();
}
if (this.config[0] === "FVAUTOCOMBO") {
return this.value;
}
if (this.config[0] === "FVAUTOSTRING") {
return this.value;
}
// For FLOAT values, check if it's actually an integer
if (Number.isInteger(this.value)) {
return this.value.toString();
}
return this.value.toFixed(this.options.precision || 2);
},
parseValue: function (v) {
if (this.config[0] === "FVAUTOSTRING") {
return v;
}
if (typeof v === "string") {
return parseFloat(v);
}
return v;
},
serializeValue: function () {
// Return special value for AUTO mode
return this.isAuto ? -99999 : this.value;
},
deserializeValue: function (data) {
if (data === -99999) {
this.isAuto = true;
this.value = -99999;
} else {
this.isAuto = false;
this.value = data;
this.cachedValue = data;
}
}
});
// Override callback to handle AUTO mode
widget.callback = function (v) {
if (this.isAuto) {
return; // Don't call the original callback in AUTO mode
}
const result = original.callback?.call(this, v);
return result;
};
// Override any potential click handlers
const originalOnClick = widget.onClick;
if (originalOnClick) {
widget.onClick = function (...args) {
if (this.isAuto) {
return false;
}
return originalOnClick.call(this, ...args);
};
}
return widget;
}
app.registerExtension({
name: "FastVideo.AutoWidgets",
async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeData?.name == "VideoGenerator" || nodeData?.name === "InferenceArgs" || nodeData?.name === "VAEConfig" ||
nodeData?.name === "TextEncoderConfig" || nodeData?.name === "DITConfig") {
// Add serialization support
chainCallback(nodeType.prototype, "onSerialize", function (info) {
if (!this.widgets) {
return;
}
// Ensure widgets_values exists
if (!info.widgets_values) {
info.widgets_values = {};
}
// Store AUTO widget states in a separate property
if (!info.auto_widget_states) {
info.auto_widget_states = {};
}
// Handle AUTO widgets specially
for (const w of this.widgets) {
if (w.type === "BOOLEAN" && w.isAuto !== undefined) {
// Store the serialized value (for Python node)
info.widgets_values[w.name] = w.serializeValue();
// Store the full state (for UI restoration)
info.auto_widget_states[w.name] = {
isAuto: w.isAuto,
value: w.value,
cachedValue: w.cachedValue
};
}
}
});
// Add deserialization support
chainCallback(nodeType.prototype, "onConfigure", function (info) {
if (!this.widgets) {
return;
}
// First, restore from widgets_values (for backward compatibility)
if (info.widgets_values && Array.isArray(info.widgets_values)) {
for (let i = 0; i < this.widgets.length && i < info.widgets_values.length; i++) {
const w = this.widgets[i];
const value = info.widgets_values[i];
if (w.type === "BOOLEAN" && w.isAuto !== undefined) {
w.deserializeValue(value);
}
}
}
// Then, restore full state if available
if (info.auto_widget_states) {
for (const w of this.widgets) {
if (w.type === "BOOLEAN" && w.isAuto !== undefined && w.name in info.auto_widget_states) {
const state = info.auto_widget_states[w.name];
w.isAuto = state.isAuto;
w.cachedValue = state.cachedValue;
w.value = state.isAuto ? -99999 : state.value;
w.callback?.(w.value);
}
}
}
// Force a redraw
this.graph?.setDirtyCanvas(true, true);
});
// Override addInput to handle AUTO widgets
chainCallback(nodeType.prototype, "onNodeCreated", function () {
// Convert any existing widgets to AUTO widgets if needed
let new_widgets = [];
const intWidgetNames = ["sp_size", "tp_size", "height", "width", "num_frames", "num_inference_steps", "flow_shift", "seed", "fps", "scale_factor",
"tile_sample_min_height", "tile_sample_min_width", "tile_sample_min_num_frames", "tile_sample_stride_height", "tile_sample_stride_width",
"tile_sample_stride_num_frames", "blend_num_frames"
]
const floatWidgetNames = ["embedded_cfg_scale", "guidance_scale"]
const comboWidgetNames = ["vae_tiling", "vae_precision", "vae_sp", "text_encoder_precision", "precision",
"load_encoder", "load_decoder", "use_tiling", "use_temporal_tiling", "use_parallel_tiling", "dit_cpu_offload", "enable_teacache"
]
const stringWidgetNames = ["prefix", "quant_config", "lora_config", "image_path"]
if (this.widgets) {
for (let w of this.widgets) {
if (intWidgetNames.includes(w.name)) {
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOINT", { "default": 0 }]));
} else if (floatWidgetNames.includes(w.name)) {
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOFLOAT", { "default": 0 }]));
} else if (comboWidgetNames.includes(w.name)) {
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOCOMBO", { "default": 0 }]));
} else if (stringWidgetNames.includes(w.name)) {
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOSTRING", { "default": "" }]));
} else {
new_widgets.push(w);
}
}
this.widgets = new_widgets;
const autoWidgets = this.widgets.filter(w => w.type === "BOOLEAN" && w.isAuto !== undefined);
}
this.graph?.setDirtyCanvas(true, true);
});
}
},
async init() {
// Force a redraw of all nodes when the extension initializes
if (app.graph) {
setTimeout(() => {
app.graph.setDirtyCanvas(true, true);
}, 1000);
}
}
});
console.log("FastVideo.core.js loaded");
-2
View File
@@ -1,2 +0,0 @@
recursive-include tk *
include config.py
-68
View File
@@ -1,68 +0,0 @@
# Sliding Tile Atteniton Kernel
## Installation
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
First, install C++20 for ThunderKittens:
```bash
sudo apt update
sudo apt install gcc-11 g++-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
sudo apt update
sudo apt install clang-11
```
Install STA:
```bash
export CUDA_HOME=/usr/local/cuda-12.4
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
git submodule update --init --recursive
python setup.py install
```
## Usage
```python
from st_attn import sliding_tile_attention
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
# a tile is a cube of size (6, 8, 8)
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
# text_length: int ranging from 0 to 256
# If your attention contains text token (Hunyuan)
out = sliding_tile_attention(q, k, v, window_size, text_length)
# If your attention does not contain text token (StepVideo)
out = sliding_tile_attention(q, k, v, window_size, 0, False)
```
## Test
```bash
python test/test_sta.py
```
## How Does STA Work?
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
## Why is STA Fast?
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
STA removes mixed blocks.
<div align="center">
<img src=../../assets/sliding_tile_attn_map.png width="80%"/>
</div>
## Acknowledgement
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
-15
View File
@@ -1,15 +0,0 @@
### ADD TO THIS TO REGISTER NEW KERNELS
sources = {
'attn': {
'source_files': {
'h100': 'st_attn/st_attn_h100.cu' # define these source files for each GPU target desired.
}
}
}
### WHICH KERNELS DO WE WANT TO BUILD?
# (oftentimes during development work you don't need to redefine them all.)
kernels = ['attn']
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
target = 'h100'
-76
View File
@@ -1,76 +0,0 @@
import os
import subprocess
from config import kernels, sources, target
from setuptools import find_packages, setup
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
target = target.lower()
# Package metadata
PACKAGE_NAME = "st_attn"
VERSION = "0.0.4"
AUTHOR = "Hao AI Lab"
DESCRIPTION = "Sliding Tile Atteniton Kernel Used in FastVideo"
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/sliding_tile_attention"
# Set environment variables
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
python_include = subprocess.check_output(['python', '-c',
"import sysconfig; print(sysconfig.get_path('include'))"]).decode().strip()
torch_include = subprocess.check_output([
'python', '-c',
"import torch; from torch.utils.cpp_extension import include_paths; print(' '.join(['-I' + p for p in include_paths()]))"
]).decode().strip()
print('st_attn root:', tk_root)
print('Python include:', python_include)
print('Torch include directories:', torch_include)
# CUDA flags
cuda_flags = [
'-DNDEBUG', '-Xcompiler=-Wno-psabi', '-Xcompiler=-fno-strict-aliasing', '--expt-extended-lambda',
'--expt-relaxed-constexpr', '-forward-unknown-to-host-compiler', '--use_fast_math', '-std=c++20', '-O3',
'-Xnvlink=--verbose', '-Xptxas=--verbose', '-Xptxas=--warn-on-spills', f'-I{tk_root}/include',
f'-I{tk_root}/prototype', f'-I{python_include}', '-DTORCH_COMPILE'
] + torch_include.split()
cpp_flags = ['-std=c++20', '-O3']
if target == 'h100':
cuda_flags.append('-DKITTENS_HOPPER')
cuda_flags.append('-arch=sm_90a')
else:
raise ValueError(f'Target {target} not supported')
source_files = ['st_attn.cpp']
for k in kernels:
if target not in sources[k]['source_files']:
raise KeyError(f'Target {target} not found in source files for kernel {k}')
if isinstance(sources[k]['source_files'][target], list):
source_files.extend(sources[k]['source_files'][target])
else:
source_files.append(sources[k]['source_files'][target])
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
setup(name=PACKAGE_NAME,
version=VERSION,
author=AUTHOR,
description=DESCRIPTION,
url=URL,
packages=find_packages(),
ext_modules=[
CUDAExtension('st_attn_cuda',
sources=source_files,
extra_compile_args={
'cxx': cpp_flags,
'nvcc': cuda_flags
},
libraries=['cuda'])
],
cmdclass={'build_ext': BuildExtension},
classifiers=[
"Programming Language :: Python :: 3",
"Environment :: GPU :: NVIDIA CUDA :: 12",
"License :: OSI Approved :: Apache Software License",
],
python_requires='>=3.10',
install_requires=["torch>=2.5.0"])
-24
View File
@@ -1,24 +0,0 @@
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <vector>
#include <cuda_fp16.h>
#include <cuda_bf16.h>
#include <cuda_runtime.h>
#ifdef TK_COMPILE_ATTN
extern torch::Tensor sta_forward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag
);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "Sliding Block Attention Kernels"; // optional module docstring
#ifdef TK_COMPILE_ATTN
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention, assuming tile size is (6,8,8)");
#endif
}
@@ -1,46 +0,0 @@
import math
import torch
from st_attn_cuda import sta_fwd
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, img_latent_shape='30*48*80'):
seq_length = q_all.shape[2]
img_latent_shape_mapping = {
'30x48x80':1,
'36x48x48':2,
'18x48x80':3,
}
if has_text:
assert q_all.shape[
2] >= 115200, "STA currently only supports video with latent size (30, 48, 80), which is 117 frames x 768 x 1280 pixels"
assert q_all.shape[1] == len(window_size), "Number of heads must match the number of window sizes"
target_size = math.ceil(seq_length / 384) * 384
pad_size = target_size - seq_length
if pad_size > 0:
q_all = torch.cat([q_all, q_all[:, :, -pad_size:]], dim=2)
k_all = torch.cat([k_all, k_all[:, :, -pad_size:]], dim=2)
v_all = torch.cat([v_all, v_all[:, :, -pad_size:]], dim=2)
else:
if img_latent_shape == '36x48x48': # Stepvideo 204x768x68
assert q_all.shape[2] == 82944
elif img_latent_shape == '18x48x80': # Wan 69x768x1280
assert q_all.shape[2] == 69120
else:
raise ValueError(f"Unsupported {img_latent_shape}, current shape is {q_all.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
kernel_aspect_ratio_flag = img_latent_shape_mapping[img_latent_shape]
hidden_states = torch.empty_like(q_all)
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
for head_index, (t_kernel, h_kernel, w_kernel) in enumerate(window_size):
for batch in range(q_all.shape[0]):
q_head, k_head, v_head, o_head = (q_all[batch:batch + 1, head_index:head_index + 1],
k_all[batch:batch + 1,
head_index:head_index + 1], v_all[batch:batch + 1,
head_index:head_index + 1],
hidden_states[batch:batch + 1, head_index:head_index + 1])
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text, kernel_aspect_ratio_flag)
if has_text:
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True, kernel_aspect_ratio_flag)
return hidden_states[:, :, :seq_length]
@@ -1,831 +0,0 @@
// # Define TORCH_COMPILE macro
#include "kittens.cuh"
#include <cooperative_groups.h>
#include <iostream>
#include <stdio.h>
#define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
#define ABS(x) ((x) < 0 ? -(x) : (x))
constexpr int CONSUMER_WARPGROUPS = (3);
constexpr int PRODUCER_WARPGROUPS = (1);
constexpr int NUM_WARPGROUPS = (CONSUMER_WARPGROUPS+PRODUCER_WARPGROUPS);
constexpr int NUM_WORKERS = (NUM_WARPGROUPS*kittens::WARPGROUP_WARPS);
using namespace kittens;
namespace cg = cooperative_groups;
template<int D> struct fwd_attend_ker_tile_dims {};
template<> struct fwd_attend_ker_tile_dims<64> {
constexpr static int tile_width = (64);
constexpr static int qo_height = (4*16);
constexpr static int kv_height = (8*16);
constexpr static int stages = (4);
};
template<> struct fwd_attend_ker_tile_dims<128> {
constexpr static int tile_width = (128);
constexpr static int qo_height = (4*16);
constexpr static int kv_height = (8*16);
constexpr static int stages = (2);
};
template<int D> struct fwd_globals {
using q_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using k_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using v_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>>;
using o_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
using q_gl = gl<bf16, -1, -1, -1, -1, q_tile>;
using k_gl = gl<bf16, -1, -1, -1, -1, k_tile>;
using v_gl = gl<bf16, -1, -1, -1, -1, v_tile>;
using l_gl = gl<float, -1, -1, -1, -1, l_col_vec>;
using o_gl = gl<bf16, -1, -1, -1, -1, o_tile>;
q_gl q;
k_gl k;
v_gl v;
l_gl l;
o_gl o;
const int N;
const int text_L;
const int hr;
};
template<int D, bool is_causal, bool text_q, bool text_kv, int DT, int DH, int DW, int CT, int CH, int CW>
__global__ __launch_bounds__((NUM_WORKERS)*kittens::WARP_THREADS, 1)
void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
extern __shared__ int __shm[];
tma_swizzle_allocator al((int*)&__shm[0]);
int warpid = kittens::warpid(), warpgroupid = warpid/kittens::WARPGROUP_WARPS;
using K = fwd_attend_ker_tile_dims<D>;
using q_tile = st_bf<K::qo_height, K::tile_width>;
using k_tile = st_bf<K::kv_height, K::tile_width>;
using v_tile = st_bf<K::kv_height, K::tile_width>;
using l_col_vec = col_vec<st_fl<K::qo_height, K::tile_width>>;
using o_tile = st_bf<K::qo_height, K::tile_width>;
q_tile (&q_smem)[CONSUMER_WARPGROUPS] = al.allocate<q_tile, CONSUMER_WARPGROUPS>();
k_tile (&k_smem)[K::stages] = al.allocate<k_tile, K::stages >();
v_tile (&v_smem)[K::stages] = al.allocate<v_tile, K::stages >();
l_col_vec (&l_smem)[CONSUMER_WARPGROUPS] = al.allocate<l_col_vec, CONSUMER_WARPGROUPS>();
auto (*o_smem) = reinterpret_cast<o_tile(*)>(q_smem);
int img_kv_blocks;
int kv_blocks = g.N / (K::kv_height);
if constexpr (text_kv) {
img_kv_blocks = kv_blocks - 3;
} else {
img_kv_blocks = kv_blocks;
}
int kv_head_idx = blockIdx.y / g.hr;
int seq_idx;
if constexpr (text_q) {
seq_idx = CT * CH * CW * 6.0 + blockIdx.x * CONSUMER_WARPGROUPS;
} else {
seq_idx = blockIdx.x * CONSUMER_WARPGROUPS;
}
__shared__ kittens::semaphore qsmem_semaphore, k_smem_arrived[K::stages], v_smem_arrived[K::stages], compute_done[K::stages];
if (threadIdx.x == 0) {
init_semaphore(qsmem_semaphore, 0, 1);
for(int j = 0; j < K::stages; j++) {
init_semaphore(k_smem_arrived[j], 0, 1);
init_semaphore(v_smem_arrived[j], 0, 1);
init_semaphore(compute_done[j], CONSUMER_WARPGROUPS, 0);
}
tma::expect_bytes(qsmem_semaphore, sizeof(q_smem));
for (int wg = 0; wg < CONSUMER_WARPGROUPS; wg++) {
coord<q_tile> q_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + wg, 0};
tma::load_async(q_smem[wg], g.q, q_tile_idx, qsmem_semaphore);
}
if constexpr (text_q){
for (int j = 0; j < K::stages - 1; j++) {
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
tma::expect_bytes(k_smem_arrived[j], sizeof(k_tile));
tma::load_async(k_smem[j], g.k, kv_tile_idx, k_smem_arrived[j]);
tma::expect_bytes(v_smem_arrived[j], sizeof(v_tile));
tma::load_async(v_smem[j], g.v, kv_tile_idx, v_smem_arrived[j]);
}
} else {
int qt = seq_idx / 6 / (CH * CW);
int qh = (seq_idx / 6) % (CH * CW) / CW;
int qw = (seq_idx / 6) % CW;
qt = CLAMP(qt, DT, CT-DT-1);
qh = CLAMP(qh, DH, CH-DH-1);
qw = CLAMP(qw, DW, CW-DW-1);
int count = 0;
int j = 0;
while (count < K::stages - 1) {
int kt = j / 3 / (CH * CW);
int kh = (j / 3) % (CH * CW) / CW;
int kw = (j / 3) % CW;
bool mask = (ABS(qt - kt) <= DT) && (ABS(qh - kh) <= DH) && (ABS(qw - kw) <= DW);
if (mask){
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
tma::expect_bytes(k_smem_arrived[count], sizeof(k_tile));
tma::load_async(k_smem[count], g.k, kv_tile_idx, k_smem_arrived[count]);
tma::expect_bytes(v_smem_arrived[count], sizeof(v_tile));
tma::load_async(v_smem[count], g.v, kv_tile_idx, v_smem_arrived[count]);
count += 1;
}
j += 1;
}
}
}
__syncthreads();
int pipe_idx = K::stages - 1;
if(warpgroupid == NUM_WARPGROUPS-1) {
warpgroup::decrease_registers<32>();
int kv_iters;
if constexpr (is_causal) {
kv_iters = (seq_idx * (K::qo_height/kittens::TILE_ROW_DIM<bf16>)) - 1 + (CONSUMER_WARPGROUPS * (K::qo_height/kittens::TILE_ROW_DIM<bf16>));
kv_iters = ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) == 0) ? (0) : ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) - 1);
}
else { kv_iters = kv_blocks-2;}
if(warpid == NUM_WORKERS-4) {
if constexpr (text_q){
for (auto kv_idx = pipe_idx - 1; kv_idx <= kv_iters; kv_idx++) {
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, kv_idx + 1, 0};
tma::expect_bytes(k_smem_arrived[(kv_idx+1)%K::stages], sizeof(k_tile));
tma::load_async(k_smem[(kv_idx+1)%K::stages], g.k, kv_tile_idx, k_smem_arrived[(kv_idx+1)%K::stages]);
tma::expect_bytes(v_smem_arrived[(kv_idx+1)%K::stages], sizeof(v_tile));
tma::load_async(v_smem[(kv_idx+1)%K::stages], g.v, kv_tile_idx, v_smem_arrived[(kv_idx+1)%K::stages]);
kittens::wait(compute_done[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
}
} else {
int qt = seq_idx / 6 / (CH * CW);
int qh = (seq_idx / 6) % (CH * CW) / CW;
int qw = (seq_idx / 6) % CW;
qt = CLAMP(qt, DT, CT-DT-1);
qh = CLAMP(qh, DH, CH-DH-1);
qw = CLAMP(qw, DW, CW-DW-1);
int k_t_min = CLAMP(qt-DT, 0, CT-1);
int k_t_max = CLAMP(qt+DT, 0, CT-1);
int k_h_min = CLAMP(qh-DH, 0, CH-1);
int k_h_max = CLAMP(qh+DH, 0, CH-1);
int k_w_min = CLAMP(qw-DW, 0, CW-1);
int k_w_max = CLAMP(qw+DW, 0, CW-1);
int count = 0;
for (int kt = k_t_min; kt <= k_t_max; kt++) {
for (int kh = k_h_min; kh <= k_h_max; kh++) {
for (int kw = k_w_min; kw <= k_w_max; kw++) {
for (int j = 0; j <= 2; j++){
if (count >= K::stages - 1) {
int index = ((kt * (CH * CW)) + (kh * CW) + kw) * 3 + j;
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
count += 1;
} else {
count += 1;
}
}
}
}
}
// for text
for (int index = img_kv_blocks; index < kv_blocks; index++) {
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
count += 1;
}
}
}
}
else {
warpgroup::increase_registers<160>();
rt_fl<16, K::kv_height> att_block;
rt_bf<16, K::kv_height> att_block_mma;
rt_fl<16, K::tile_width> o_reg;
col_vec<rt_fl<16, K::kv_height>> max_vec, norm_vec, max_vec_last_scaled, max_vec_scaled;
neg_infty(max_vec);
zero(norm_vec);
zero(o_reg);
int kv_iters;
if constexpr (is_causal) {
kv_iters = (seq_idx * 4) - 1 + (CONSUMER_WARPGROUPS * 4);
kv_iters = (kv_iters/8);
}
else if constexpr (text_q){
// the last three kv blocks are for text, we process them separately
kv_iters = img_kv_blocks - 1;
} else {
kv_iters = CLAMP(DT*2+1, 1, CT) * CLAMP(DH*2+1, 1, CH) * CLAMP(DW*2+1, 1, CW) * 3 - 1 ;
}
kittens::wait(qsmem_semaphore, 0);
for (auto kv_idx = 0; kv_idx <= kv_iters; kv_idx++) {
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
copy(max_vec_last_scaled, max_vec);
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
warpgroup::mma_async_wait();
row_max(max_vec, att_block, max_vec);
if constexpr (D == 64) {
mul(att_block, att_block, 1.44269504089f*0.125f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
}
else {
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
}
sub_row(att_block, att_block, max_vec_scaled);
exp2(att_block, att_block);
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
exp2(max_vec_last_scaled, max_vec_last_scaled);
mul(norm_vec, norm_vec, max_vec_last_scaled);
row_sum(norm_vec, att_block, norm_vec);
add(att_block, att_block, 0.f);
copy(att_block_mma, att_block);
mul_row(o_reg, o_reg, max_vec_last_scaled);
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
warpgroup::mma_async_wait();
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
}
// the last three kv blocks are for text, we process them separately
if constexpr(text_kv) {
for (auto kv_idx = kv_iters + 1; kv_idx <= kv_iters + 3; kv_idx++) {
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
copy(max_vec_last_scaled, max_vec);
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
warpgroup::mma_async_wait();
// apply non-pad mask
int offset = g.text_L - (kv_idx - (kv_iters + 1)) * K::kv_height;
// printf("k_idx_start: %d, k_idx_end: %d, text_end: %d, offset: %d\n", k_idx_start, k_idx_end, text_end, offset);
right_fill(att_block, att_block, offset, base_types::constants<float>::neg_infty());
row_max(max_vec, att_block, max_vec);
if constexpr (D == 64) {
mul(att_block, att_block, 1.44269504089f*0.125f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
}
else {
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
}
sub_row(att_block, att_block, max_vec_scaled);
exp2(att_block, att_block);
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
exp2(max_vec_last_scaled, max_vec_last_scaled);
mul(norm_vec, norm_vec, max_vec_last_scaled);
row_sum(norm_vec, att_block, norm_vec);
add(att_block, att_block, 0.f);
copy(att_block_mma, att_block);
mul_row(o_reg, o_reg, max_vec_last_scaled);
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
warpgroup::mma_async_wait();
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
}
}
div_row(o_reg, o_reg, norm_vec);
warpgroup::store(o_smem[warpgroupid], o_reg);
warpgroup::sync(warpgroupid+4);
if (warpid % 4 == 0) {
coord<o_tile> o_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + warpgroupid, 0};
tma::store_async(g.o, o_smem[warpgroupid], o_tile_idx);
}
mul(max_vec_scaled, max_vec_scaled, 0.69314718056f);
log(norm_vec, norm_vec);
add(norm_vec, norm_vec, max_vec_scaled);
if constexpr (D == 64) { mul(norm_vec, norm_vec, -8.0f); }
else { mul(norm_vec, norm_vec, -11.313708499f); }
warpgroup::store(l_smem[warpgroupid], norm_vec);
warpgroup::sync(warpgroupid+4);
if (warpid % 4 == 0) {
coord<l_col_vec> tile_idx = {blockIdx.z, blockIdx.y, 0, (seq_idx) + warpgroupid};
tma::store_async(g.l, l_smem[warpgroupid], tile_idx);
}
tma::store_async_wait();
}
}
#include "pyutils/torch_helpers.cuh"
#include <ATen/cuda/CUDAContext.h>
#include <iostream>
torch::Tensor
sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_h_size, int kernel_w_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag)
{
CHECK_INPUT(q);
CHECK_INPUT(k);
CHECK_INPUT(v);
auto batch = q.size(0);
auto seq_len = q.size(2);
auto head_dim = q.size(3);
auto qo_heads = q.size(1);
auto kv_heads = k.size(1);
// check to see that these dimensions match for all inputs
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(v.size(0) == batch, "V batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(k.size(3) == head_dim, "K head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(v.size(3) == head_dim, "V head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(qo_heads >= kv_heads, "QO heads must be greater than or equal to KV heads");
TORCH_CHECK(qo_heads % kv_heads == 0, "QO heads must be divisible by KV heads");
TORCH_CHECK(q.size(1) == qo_heads, "QO head dimension - idx 1 - must match for all inputs");
TORCH_CHECK(k.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
TORCH_CHECK(v.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
auto hr = qo_heads / kv_heads;
c10::BFloat16* q_ptr = q.data_ptr<c10::BFloat16>();
c10::BFloat16* k_ptr = k.data_ptr<c10::BFloat16>();
c10::BFloat16* v_ptr = v.data_ptr<c10::BFloat16>();
bf16* d_q = reinterpret_cast<bf16*>(q_ptr);
bf16* d_k = reinterpret_cast<bf16*>(k_ptr);
bf16* d_v = reinterpret_cast<bf16*>(v_ptr);
torch::Tensor l_vec = torch::empty({static_cast<const uint>(batch),
static_cast<const uint>(qo_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(1)},
torch::TensorOptions().dtype(torch::kFloat).device(q.device()).memory_format(at::MemoryFormat::Contiguous));
bf16* o_ptr = reinterpret_cast<bf16*>(o.data_ptr<c10::BFloat16>());
bf16* d_o = reinterpret_cast<bf16*>(o_ptr);
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
float* d_l = reinterpret_cast<float*>(l_ptr);
cudaDeviceSynchronize();
auto stream = at::cuda::getCurrentCUDAStream().stream();
if (head_dim == 128) {
using q_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using k_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using v_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>>;
using o_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
using q_global = gl<bf16, -1, -1, -1, -1, q_tile>;
using k_global = gl<bf16, -1, -1, -1, -1, k_tile>;
using v_global = gl<bf16, -1, -1, -1, -1, v_tile>;
using l_global = gl<float, -1, -1, -1, -1, l_col_vec>;
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
using globals = fwd_globals<128>;
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(text_length), static_cast<int>(hr)};
auto mem_size = kittens::MAX_SHARED_MEMORY;
auto threads = NUM_WORKERS * kittens::WARP_THREADS;
if (has_text) {
// TORCH_CHECK(seq_len % (CONSUMER_WARPGROUPS*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 192");
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4)-2, qo_heads, batch);
dim3 grid_text(2, qo_heads, batch);
if (!process_text) {
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 1, 1, 1, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 1, 1, 1, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 1, 1, 2, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true,1, 1, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 1, 1, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 1, 1, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
}else if (kernel_t_size ==3 && kernel_h_size == 5 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 1, 2, 2, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 1, 2, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==5 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 3, 0, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 3, 0, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==5 && kernel_h_size == 3 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 1, 2, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 1, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 2, 2, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 2, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 2, 3, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 2, 3, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 3, 5, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 3, 5, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 0, 0, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 2, 0, 0, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 0, 3, 5, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true, 0, 3, 5, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, true, 2, 0, 5, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, true,2, 0, 5, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else {
// print error
std::cout << "Invalid kernel size" << std::endl;
//print kernel size
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
}
} else {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10><<<grid_text, (32*NUM_WORKERS), mem_size, stream>>>(g);
}
} else {
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
if (kernel_aspect_ratio_flag == 2){
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
}else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 3){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size ==6 && kernel_h_size == 3 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else {
// print error
std::cout << "Invalid kernel size" << std::endl;
//print kernel size
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
}
}
else if (kernel_aspect_ratio_flag == 3) {
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 1, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 1, 1, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 2, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 2, 2, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 2, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 0, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 0, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 2, 3, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 2, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 9){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 2, 4, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 2, 4, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 1, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 1, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 0, 0, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 0, 0, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 2, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 2, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 3, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 2, 3, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 2, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 9){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 2, 4, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 2, 4, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 0, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 0, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 3 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 1, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,0, 1, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 2, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,0, 3, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else {
// print error
std::cout << "Invalid kernel size" << std::endl;
//print kernel size
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
}
}
else {
std::cout << "Unsupported kernel_aspect_ratio_flag: " << kernel_aspect_ratio_flag << std::endl;
}
}
CHECK_CUDA_ERROR(cudaGetLastError());
cudaStreamSynchronize(stream);
}
return o;
cudaDeviceSynchronize();
}
-151
View File
@@ -1,151 +0,0 @@
import os
from collections import defaultdict
import matplotlib.pyplot as plt
import numpy as np
import torch
from st_attn import sliding_tile_attention
def flops(batch, seqlen, nheads, headdim, causal, mode="fwd"):
assert mode in ["fwd", "bwd", "fwd_bwd"]
f = 4 * batch * seqlen**2 * nheads * headdim // (2 if causal else 1)
return f if mode == "fwd" else (2.5 * f if mode == "bwd" else 3.5 * f)
def efficiency(flop, time):
flop = flop / 1e12
time = time / 1e6
return flop / time
def benchmark_attention(configurations):
results = {'fwd': defaultdict(list), 'bwd': defaultdict(list)}
for B, H, N, D, causal in configurations:
print("=" * 60)
print(f"Timing forward and backward pass for B={B}, H={H}, N={N}, D={D}, causal={causal}")
q = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
k = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
v = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
grad_output = torch.randn_like(q, requires_grad=False).contiguous()
qg = torch.zeros_like(q, requires_grad=False, dtype=torch.float).contiguous()
kg = torch.zeros_like(k, requires_grad=False, dtype=torch.float).contiguous()
vg = torch.zeros_like(v, requires_grad=False, dtype=torch.float).contiguous()
# Prepare for timing forward pass
start_events_fwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
end_events_fwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
torch.cuda.empty_cache()
torch.cuda.synchronize()
# Warmup for forward pass
for _ in range(10):
o = sliding_tile_attention(q, k, v, [[6, 6, 6]] * 24, 0, False)
# Time the forward pass
for i in range(10):
start_events_fwd[i].record()
o = sliding_tile_attention(q, k, v, [[6, 6, 6]] * 24, 0, False)
end_events_fwd[i].record()
torch.cuda.synchronize()
times_fwd = [s.elapsed_time(e) for s, e in zip(start_events_fwd, end_events_fwd)]
time_us_fwd = np.mean(times_fwd) * 1000
tflops_fwd = efficiency(flops(B, N, H, D, causal, 'fwd'), time_us_fwd)
results['fwd'][(D, causal)].append((N, tflops_fwd))
print(f"Average time for forward pass in us: {time_us_fwd:.2f}")
print(f"Average efficiency for forward pass in TFLOPS: {tflops_fwd}")
print("-" * 60)
# torch.cuda.empty_cache()
# torch.cuda.synchronize()
# # Prepare for timing backward pass
# start_events_bwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
# end_events_bwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
# # Warmup for backward pass
# for _ in range(10):
# qg, kg, vg = tk.mha_backward(q, k, v, o, l_vec, grad_output, causal)
# # Time the backward pass
# for i in range(10):
# start_events_bwd[i].record()
# qg, kg, vg = tk.mha_backward(q, k, v, o, l_vec, grad_output, causal)
# end_events_bwd[i].record()
# torch.cuda.synchronize()
# times_bwd = [s.elapsed_time(e) for s, e in zip(start_events_bwd, end_events_bwd)]
# time_us_bwd = np.mean(times_bwd) * 1000
# tflops_bwd = efficiency(flops(B, N, H, D, causal, 'bwd'), time_us_bwd)
# results['bwd'][(D, causal)].append((N, tflops_bwd))
# print(f"Average time for backward pass in us: {time_us_bwd:.2f}")
# print(f"Average efficiency for backward pass in TFLOPS: {tflops_bwd}")
print("=" * 60)
torch.cuda.empty_cache()
torch.cuda.synchronize()
return results
def plot_results(results):
os.makedirs('benchmark_results', exist_ok=True)
for mode in ['fwd', 'bwd']:
for (D, causal), values in results[mode].items():
seq_lens = [x[0] for x in values]
tflops = [x[1] for x in values]
plt.figure(figsize=(10, 6))
bars = plt.bar(range(len(seq_lens)), tflops, tick_label=seq_lens)
plt.xlabel('Sequence Length')
plt.ylabel('TFLOPS')
plt.title(f'{mode.upper()} Pass - Head Dim: {D}, Causal: {causal}')
plt.grid(True)
# Adding the numerical y value on top of each bar
for bar in bars:
yval = bar.get_height()
plt.text(bar.get_x() + bar.get_width() / 2, yval, round(yval, 2), ha='center', va='bottom')
filename = f'benchmark_results/{mode}_D{D}_causal{causal}.png'
plt.savefig(filename)
plt.close()
# Example list of configurations to test
configurations = [
(2, 24, 82944, 128, False),
# (16, 16, 768*16, 128, False),
# (16, 16, 768*2, 128, False),
# (16, 16, 768*4, 128, False),
# (16, 16, 768*8, 128, False),
# (16, 16, 768*16, 128, False),
# (16, 16, 768, 128, True),
# (16, 16, 768*2, 128, True),
# (16, 16, 768*4, 128, True),
# (16, 16, 768*8, 128, True),
# (16, 16, 768*16, 128, True),
# (16, 32, 768, 64, False),
# (16, 32, 768*2, 64, False),
# (16, 32, 768*4, 64, False),
# (16, 32, 768*8, 64, False),
# (16, 32, 768*16, 64, False),
# (16, 32, 768, 64, True),
# (16, 32, 768*2, 64, True),
# (16, 32, 768*4, 64, True),
# (16, 32, 768*8, 64, True),
# (16, 32, 768*16, 64, True),
]
results = benchmark_attention(configurations)
# plot_results(results)

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