Compare commits

...
Author SHA1 Message Date
Will Lin 3572c6821e move tests 2026-01-21 14:09:58 -08:00
Will Lin deb901f6fc cleanup 2026-01-21 14:09:40 -08:00
Will Lin beae2943cf cleanup 2026-01-21 13:28:43 -08:00
Will Lin fadb71bb64 fix 2026-01-20 18:45:26 -08:00
Shao Duan 65c6fbaa46 Use PyAV for audio muxing instead of ffmpeg CLI
Replace ffmpeg subprocess call with PyAV library for muxing audio
into video files. PyAV bundles FFmpeg libraries, so users no longer
need ffmpeg CLI installed separately.
2026-01-20 21:32:47 +00:00
Shao Duan e32a8a9504 added tiling to ltx vae, simplified example for ltx2 2026-01-20 08:32:06 +00:00
Shao Duan 6907d87871 cleanup 2026-01-20 08:32:06 +00:00
Shao Duan 9775fed31a Implement native LTX2 Video and Audio VAEs 2026-01-20 08:32:06 +00:00
Shao Duan 5c6f635d73 Use LTX2 config defaults and HF model ID 2026-01-20 08:32:06 +00:00
Shao Duan c19708fd58 Fix test_ltx2_audio imports and use consistent attention backend 2026-01-20 08:32:06 +00:00
Shao Duan b2e4fb0743 Use HuggingFace model path for LTX2 examples and CI tests 2026-01-20 08:32:06 +00:00
Shao Duan a3ad4852b0 Fix yapf formatting 2026-01-20 08:32:06 +00:00
Shao Duan add2be21b5 Fix missing LTX2 exports and add layerwise offload compatibility
- Restore LTX2VideoConfig and add to dits __init__
- Add LTX2VAEConfig export to vaes __init__
- Add LTX2Transformer3DModel and CausalVideoAutoencoder to model registry
- Add compatibility check for layerwise offload (skip for models without nn.ModuleList)
2026-01-20 08:32:06 +00:00
Shao Duan 41203d92b8 Fix pre-commit issues and code cleanup
- Fix ruff SIM102 errors in ltx2_denoising.py (combine nested if statements)
- Remove exit() call from create_hf_repo.py
- Remove global torch.backends.cuda settings from gemma.py property
- Apply yapf formatting fixes
2026-01-20 08:32:06 +00:00
Shao Duan 5fa8415c0b Remove development markdown files 2026-01-20 08:32:06 +00:00
Shao Duan 3a182925f3 Fix LTX2 distilled to skip CFG with guidance_scale=1.0 2026-01-20 08:32:06 +00:00
Shao Duan c1e4787775 Added audio to ltx2 and fixed issue with sigma values and encoder alignment (#1004) 2026-01-20 08:32:06 +00:00
Will Lin 2c6bf47b9f update 2026-01-20 08:32:06 +00:00
Will Lin 548cc08817 working 2026-01-20 08:32:06 +00:00
Will Lin 7521b06693 update 2026-01-20 08:32:06 +00:00
Will Lin fb9ad77086 encoder 2026-01-20 08:32:06 +00: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
920 changed files with 87999 additions and 37189 deletions
+90 -47
View File
@@ -13,40 +13,40 @@ steps:
- label: "Trigger Tests"
plugins:
- monorepo-diff#v1.4.0:
diff: "git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD"
diff: 'git fetch origin "$BUILDKITE_PULL_REQUEST_BASE_BRANCH" && git diff --name-only origin/"$BUILDKITE_PULL_REQUEST_BASE_BRANCH"...HEAD'
watch:
- path:
- "fastvideo/v1/models/encoders/**"
- "fastvideo/v1/models/loader/**"
- "fastvideo/v1/tests/encoders/**"
- "fastvideo/models/encoders/**"
- "fastvideo/models/loader/**"
- "fastvideo/tests/encoders/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 20m .buildkite/scripts/pr_test.sh"
label: "Encoder Tests"
env:
- TEST_TYPE=encoder
agents:
queue: "default"
- path:
- "fastvideo/v1/models/vaes/**"
- "fastvideo/v1/models/loader/**"
- "fastvideo/v1/tests/vaes/**"
- "fastvideo/models/vaes/**"
- "fastvideo/models/loader/**"
- "fastvideo/tests/vaes/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 20m .buildkite/scripts/pr_test.sh"
label: "VAE Tests"
env:
- TEST_TYPE=vae
agents:
queue: "default"
- path:
- "fastvideo/v1/models/dits/**"
- "fastvideo/v1/models/loader/**"
- "fastvideo/v1/tests/transformers/**"
- "fastvideo/v1/layers/**"
- "fastvideo/v1/attention/**"
- "fastvideo/models/dits/**"
- "fastvideo/models/loader/**"
- "fastvideo/tests/transformers/**"
- "fastvideo/layers/**"
- "fastvideo/attention/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
@@ -57,31 +57,33 @@ steps:
agents:
queue: "default"
- path:
- "fastvideo/v1/**/*.py"
- "fastvideo/**/*.py"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
label: "SSIM Tests"
env:
- TEST_TYPE=ssim
agents:
queue: "default"
- path:
- "fastvideo/v1/tests/lora/**"
- "fastvideo/v1/models/loader/**"
- "fastvideo/v1/tests/transformers/**"
- "fastvideo/v1/pipelines/**"
- "fastvideo/v1/layers/lora/**"
- "fastvideo/tests/lora/**"
- "fastvideo/models/loader/**"
- "fastvideo/tests/transformers/**"
- "fastvideo/pipelines/**"
- "fastvideo/layers/lora/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 20m .buildkite/scripts/pr_test.sh"
label: "LoRA Inference Tests"
env:
- TEST_TYPE=inference_lora
agents:
queue: "default"
- path:
- "fastvideo/v1/**"
- "fastvideo/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
@@ -92,12 +94,42 @@ steps:
agents:
queue: "default"
- path:
- "fastvideo/v1/**"
- "csrc/attn/vsa/**"
- "csrc/attn/tk/**"
- "csrc/attn/setup_vsa.py"
- "csrc/attn/config_vsa.py"
- "csrc/attn/vsa.cpp"
- "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:
@@ -108,11 +140,8 @@ steps:
agents:
queue: "default"
- path:
- "fastvideo/v1/**"
- "csrc/attn/st_attn/**"
- "csrc/attn/setup_sta.py"
- "csrc/attn/config_sta.py"
- "csrc/attn/st_attn.cpp"
- "fastvideo/**"
- "fastvideo-kernel/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
@@ -123,31 +152,45 @@ steps:
agents:
queue: "default"
- path:
- "csrc/attn/st_attn/**"
- "csrc/attn/setup_sta.py"
- "csrc/attn/config_sta.py"
- "csrc/attn/st_attn.cpp"
- "fastvideo-kernel/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests STA"
label: "Kernel Tests"
env:
- TEST_TYPE=precision_sta
- 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:
- "csrc/attn/vsa/**"
- "csrc/attn/tk/**"
- "csrc/attn/setup_vsa.py"
- "csrc/attn/config_vsa.py"
- "csrc/attn/vsa.cpp"
- "fastvideo/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VSA"
label: "Unit Tests"
env:
- TEST_TYPE=precision_vsa
- 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"
+35 -14
View File
@@ -31,9 +31,9 @@ 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)
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"
@@ -50,7 +50,7 @@ else
exit 1
fi
MODAL_TEST_FILE="fastvideo/v1/tests/modal/pr_test.py"
MODAL_TEST_FILE="fastvideo/tests/modal/pr_test.py"
if [ -z "${TEST_TYPE:-}" ]; then
log "Error: TEST_TYPE environment variable is not set"
@@ -63,24 +63,28 @@ MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUI
case "$TEST_TYPE" in
"encoder")
log "Running encoder tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_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 python3 -m modal run $MODAL_TEST_FILE::run_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 python3 -m modal run $MODAL_TEST_FILE::run_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 python3 -m modal run $MODAL_TEST_FILE::run_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"
@@ -89,18 +93,35 @@ case "$TEST_TYPE" in
log "Running inference STA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_STA"
;;
"precision_sta")
log "Running precision STA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_STA"
;;
"precision_vsa")
log "Running precision VSA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_VSA"
"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
+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
+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/
+69 -43
View File
@@ -62,8 +62,8 @@ on:
required: false
default: false
type: boolean
run_nightly_test:
description: "Run nightly-test"
run_unit_test:
description: "Run unit-test"
required: false
default: false
type: boolean
@@ -93,6 +93,7 @@ jobs:
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
@@ -102,50 +103,53 @@ jobs:
# 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/st_attn/**'
- 'csrc/attn/setup_sta.py'
- 'csrc/attn/config_sta.py'
- 'csrc/attn/st_attn.cpp'
- '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/vsa/**'
- 'csrc/attn/tk/**'
- 'csrc/attn/setup_vsa.py'
- 'csrc/attn/config_vsa.py'
- 'csrc/attn/vsa.cpp'
- '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/v1/**'
- 'fastvideo/**'
- *common-paths
- *vsa-kernel-paths
# Actual tests
encoder-test:
- 'fastvideo/v1/models/encoders/**'
- 'fastvideo/v1/models/loader/**'
- 'fastvideo/v1/tests/encoders/**'
- 'fastvideo/models/encoders/**'
- 'fastvideo/models/loader/**'
- 'fastvideo/tests/encoders/**'
- *common-paths
vae-test:
- 'fastvideo/v1/models/vaes/**'
- 'fastvideo/v1/models/loader/**'
- 'fastvideo/v1/tests/vaes/**'
- 'fastvideo/models/vaes/**'
- 'fastvideo/models/loader/**'
- 'fastvideo/tests/vaes/**'
- *common-paths
transformer-test:
- 'fastvideo/v1/models/dits/**'
- 'fastvideo/v1/models/loader/**'
- '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/v1/**'
- 'fastvideo/**'
- *common-paths
training-test-VSA:
- 'fastvideo/v1/**'
- 'fastvideo/**'
- *common-paths
- *vsa-kernel-paths
inference-test-STA:
- 'fastvideo/v1/**'
- 'fastvideo/**'
- *common-paths
- *sta-kernel-paths
precision-test-STA:
@@ -154,6 +158,9 @@ jobs:
precision-test-VSA:
- *common-paths
- *vsa-kernel-paths
unit-test:
- 'fastvideo/**'
- *common-paths
encoder-test:
needs: change-filter
@@ -167,7 +174,7 @@ jobs:
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/encoders -s"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
@@ -185,7 +192,7 @@ jobs:
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/vaes -s"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
@@ -203,7 +210,7 @@ jobs:
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/transformers -s"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
@@ -229,12 +236,12 @@ jobs:
volume_size: 200
disk_size: 200
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:${{ matrix.python-version.tag }}"
test_command: "uv 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: >-
@@ -248,7 +255,7 @@ jobs:
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/v1/tests/training/Vanilla -srP"
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 }}
@@ -268,7 +275,7 @@ jobs:
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/v1/tests/training/VSA -srP"
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 }}
@@ -288,7 +295,7 @@ jobs:
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/v1/tests/inference/STA -srP"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/inference/STA -srP"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
@@ -326,29 +333,48 @@ jobs:
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_block_sparse.py"
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 }}
nightly-test:
unit-test:
needs: change-filter
if: >-
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
(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: "nightly-test"
gpu_type: "NVIDIA A40"
gpu_count: 4
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: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/nightly/test_e2e_overfit_single_sample.py -vs"
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 }}
WANDB_API_KEY: ${{ secrets.WANDB_API_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:
# Add other jobs to this list as you create them
@@ -372,4 +398,4 @@ jobs:
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
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 }}
+11 -11
View File
@@ -5,7 +5,7 @@ on:
branches:
- main
paths:
- "csrc/attn/setup_sta.py"
- "csrc/attn/sliding_tile_attn/setup.py"
workflow_dispatch:
jobs:
@@ -23,13 +23,13 @@ jobs:
- name: Check if version changed
id: check-version
run: |
cd csrc/attn
cd csrc/attn/sliding_tile_attn
# Get current commit's version
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_sta.py)
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_sta.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
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
@@ -144,13 +144,13 @@ jobs:
pip install setuptools
pip install ninja packaging wheel
cd csrc/attn # Move into the correct folder
cd csrc/attn/sliding_tile_attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup_sta.py bdist_wheel --dist-dir=dist
python setup.py bdist_wheel --dist-dir=dist
- name: Rename wheel file
run: |
cd csrc/attn
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)
@@ -165,7 +165,7 @@ jobs:
uses: actions/upload-artifact@v4
with:
name: ${{ env.wheel_name }}
path: csrc/attn/dist/*.whl
path: csrc/attn/sliding_tile_attn/dist/*.whl
retention-days: 90
publish_package:
@@ -239,11 +239,11 @@ jobs:
pip install setuptools
pip install ninja packaging wheel
cd csrc/attn # Move into the correct folder
cd csrc/attn/sliding_tile_attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup_sta.py sdist --dist-dir=dist
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/dist/
packages-dir: csrc/attn/sliding_tile_attn/dist/
+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/
+17 -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,8 @@ env
**/build/
**.pyc
**.txt
*.log
weights/
# Distribution / packaging
build/
@@ -36,10 +41,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,9 +63,12 @@ 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
dmd_t2v_output/
preprocess_output_text/
+5 -2
View File
@@ -1,3 +1,6 @@
[submodule "csrc/attn/tk"]
path = csrc/attn/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/.*|
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 -v "^fastvideo/v1/tests/inference/lora/L40S_reference_videos/" | 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
+85 -75
View File
@@ -1,42 +1,47 @@
<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-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ" 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://ibb.co/sv3MMKyv" 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/).
<details>
<summary>More</summary>
- ```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/).
</details>
## 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 achineve >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:
@@ -50,19 +55,33 @@ 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.1-T2V-14B-Preview](https://huggingface.co/FastVideo/FastWan2.1-T2V-14B-Diffusers) | Coming soon! | [FastVideo Synthetic Wan2.1 720P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x768x1280_250k) |
| [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
)
@@ -87,77 +106,68 @@ 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:
- [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/)
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/)
## 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/finetune.html)
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd/)
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
## 📑 Development Plan
## Awesome work using FastVideo or our research projects
<!-- - 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. [![Star](https://img.shields.io/github/stars/sgl-project/sglang.svg?style=social&label=Star)](https://github.com/sgl-project/sglang)
- [DanceGRPO](https://github.com/XueZeyue/DanceGRPO): A unified framework to adapt Group Relative Policy Optimization (GRPO) to visual generation paradigms. Code based on FastVideo. [![Star](https://img.shields.io/github/stars/XueZeyue/DanceGRPO.svg?style=social&label=Star)](https://github.com/XueZeyue/DanceGRPO)
- [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. [![Star](https://img.shields.io/github/stars/Tencent-Hunyuan/SRPO.svg?style=social&label=Star)](https://github.com/Tencent-Hunyuan/SRPO)
- [DCM](https://github.com/Vchitect/DCM): Dual-expert consistency model for efficient and high-quality video generation. Code based on FastVideo. [![Star](https://img.shields.io/github/stars/Vchitect/DCM.svg?style=social&label=Star)](https://github.com/Vchitect/DCM)
- [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. [![Star](https://img.shields.io/github/stars/Tencent-Hunyuan/HunyuanVideo-1.5.svg?style=social&label=Star)](https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5)
- [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. [![Star](https://img.shields.io/github/stars/kandinskylab/kandinsky-5.svg?style=social&label=Star)](https://github.com/kandinskylab/kandinsky-5)
- [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. [![Star](https://img.shields.io/github/stars/meituan-longcat/LongCat-Video.svg?style=social&label=Star)](https://github.com/meituan-longcat/LongCat-Video)
## 🤝 Contributing
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/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)
- [Wan-Video](https://github.com/Wan-Video)
- [ThunderKittens](https://github.com/HazyResearch/ThunderKittens)
- [Triton](https://github.com/triton-lang/triton)
- [DMD2](https://github.com/tianweiy/DMD2)
- [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 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 considering citing our work:
```bibtex
@misc{zhang2025vsafastervideodiffusion,
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
author={Peiyuan Zhang and Haofeng Huang and Yongqi Chen and Will Lin and Zhengzhong Liu and Ion Stoica and Eric Xing and Hao Zhang},
year={2025},
eprint={2505.13389},
archivePrefix={arXiv},
primaryClass={cs.CV},
url={https://arxiv.org/abs/2505.13389},
@software{fastvideo2024,
title = {FastVideo: A Unified Framework for Accelerated Video Generation},
author = {The FastVideo Team},
url = {https://github.com/hao-ai-lab/FastVideo},
month = apr,
year = {2024},
}
@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

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.
+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
@@ -81,6 +81,7 @@ DEFAULT_CONDA_PATTERNS = {
DEFAULT_PIP_PATTERNS = {
"torch",
"numpy",
"mypy",
"flake8",
"triton",
"optree",
+2 -2
View File
@@ -74,8 +74,8 @@ After installation, the following nodes will be available in the ComfyUI interfa
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/v1/configs/wan_14B_i2v_480p_pipeline.json)
- [FastHunyuan-diffusers](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/v1/configs/fasthunyuan_t2v.json)
- [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
+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

+1 -1
View File
@@ -310,7 +310,7 @@
"value": -99999,
"cachedValue": "fp16"
},
"use_cpu_offload": {
"dit_cpu_offload": {
"isAuto": true,
"value": -99999,
"cachedValue": true
@@ -517,7 +517,7 @@
"value": "fp16",
"cachedValue": "fp16"
},
"use_cpu_offload": {
"dit_cpu_offload": {
"isAuto": true,
"value": true,
"cachedValue": true
+4 -4
View File
@@ -105,7 +105,7 @@ class VideoGenerator:
"precision": (["fp16", "bf16"], {
"default": "fp16"
}),
"use_cpu_offload": ([True, False], {
"dit_cpu_offload": ([True, False], {
"default": False
}),
}
@@ -204,7 +204,7 @@ class VideoGenerator:
vae_config=None,
text_encoder_config=None,
dit_config=None,
use_cpu_offload=None,
dit_cpu_offload=None,
):
print('Running FastVideo inference')
@@ -259,8 +259,8 @@ class VideoGenerator:
raw_generation_args['tp_size'] = tp_size
if sp_size is not None:
raw_generation_args['sp_size'] = sp_size
if use_cpu_offload is not None:
raw_generation_args['use_cpu_offload'] = use_cpu_offload
if dit_cpu_offload is not None:
raw_generation_args['dit_cpu_offload'] = dit_cpu_offload
generation_args = {
k: v
+1 -1
View File
@@ -552,7 +552,7 @@ app.registerExtension({
]
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", "use_cpu_offload", "enable_teacache"
"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"]
-2
View File
@@ -1,2 +0,0 @@
recursive-include tk *
include config.py
-83
View File
@@ -1,83 +0,0 @@
# Sliding Tile Atteniton Kernel
## Installation
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only support H100/H200, because ThunderKittens uses TMA but doesn't support Blackwell yet.
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
```
## Environment Setup
First, set up your CUDA environment:
```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
```
## Install Sliding Tile Attention (STA)
```bash
python setup_sta.py install
```
## Install Video Sparse Attention (VSA)
```bash
python setup_vsa.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 tests/test_sta.py # test STA
python tests/test_block_sparse.py # test VSA
```
## Benchmark
```bash
python benchmarks/bench_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.
-147
View File
@@ -1,147 +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
from triton.testing import do_bench
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 compute_TFLOPS(flops, ms):
flops = flops / 1e12
ms = ms / 1e3
return flops / ms
def benchmark_attention(configurations):
results = {'fwd': defaultdict(list), 'bwd': defaultdict(list)}
for B, H, N, D, causal, dit_seq_shape, window_size 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()
# # Warmup for forward pass
# for _ in range(10):
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
# # Time the forward pass
# for i in range(10):
# start_events_fwd[i].record()
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
# end_events_fwd[i].record()
ms = do_bench(lambda: sliding_tile_attention(q, k, v, [window_size] * 24, 0, False, dit_seq_shape))
# 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 = compute_TFLOPS(flops(B, N, H, D, causal, 'fwd'), ms)
results['fwd'][(D, causal)].append((N, tflops_fwd))
print(f"Average time for forward pass (ms): {ms:.2f}")
print(f"Average 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 = compute_TFLOPS(flops(B, N, H, D, causal, 'bwd'), ms)
# results['bwd'][(D, causal)].append((N, tflops_bwd))
# print(f"Average time for backward pass(ms): {ms:.2f}")
# print(f"Average TFLOPS: {tflops_bwd}")
# print("=" * 60)
torch.cuda.empty_cache()
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, 69120, 128, False, '18x48x80', [3, 6, 10]),
(2, 24, 69120, 128, True, '18x48x80', [3, 6, 10]),
(2, 24, 82944, 128, False, '36x48x48', [3, 3, 6]), # Stepvideo
(2, 24, 82944, 128, True, '36x48x48', [3, 3, 6]),
# (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)
-225
View File
@@ -1,225 +0,0 @@
import torch
import argparse
from flash_attn.utils.benchmark import benchmark_forward
from vsa import block_sparse_attention_fwd, block_sparse_attention_backward
from vsa import BLOCK_M, BLOCK_N
import numpy as np
import random
def set_seed(seed: int = 42):
# Python random module
random.seed(seed)
# NumPy
np.random.seed(seed)
# PyTorch
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed) # if using multi-GPU
def parse_arguments():
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
parser.add_argument('--batch_size', type=int, default=1, help='Batch size')
parser.add_argument('--num_heads', type=int, default=12, help='Number of heads')
parser.add_argument('--head_dim', type=int, default=64, help='Head dimension')
parser.add_argument('--topk', type=int, default=None, help='Number of kv blocks each q block attends to')
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[49152], help='Sequence lengths to benchmark')
return parser.parse_args()
def create_input_tensors(batch, head, seq_len, headdim):
"""Create random input tensors for attention."""
q = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
k = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
v = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
return q, k, v
def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device="cuda"):
"""
Generate a block sparse pattern where each q block attends to exactly k kv blocks.
Args:
bs: batch size
h: number of heads
num_q_blocks: number of query blocks
num_kv_blocks: number of key-value blocks
k: number of kv blocks each q block attends to
device: device to create tensors on
Returns:
q2k_block_sparse_index: [bs, h, num_q_blocks, k]
Contains the indices of kv blocks that each q block attends to.
q2k_block_sparse_num: [bs, h, num_q_blocks]
Contains the number of kv blocks that each q block attends to (all equal to k).
k2q_block_sparse_index: [bs, h, num_kv_blocks, num_q_blocks]
Contains the indices of q blocks that attend to each kv block.
k2q_block_sparse_num: [bs, h, num_kv_blocks]
Contains the number of q blocks that attend to each kv block.
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks]
Binary mask where 1 indicates attention connection.
"""
# Ensure k is not larger than num_kv_blocks
k = min(k, num_kv_blocks)
# Create random scores for sampling
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device=device)
# Get top-k indices for each q block
_, q2k_block_sparse_index = torch.topk(scores, k, dim=-1)
q2k_block_sparse_index = q2k_block_sparse_index.to(torch.int32)
# sort q2k_block_sparse_index
q2k_block_sparse_index, _ = torch.sort(q2k_block_sparse_index, dim=-1)
# All q blocks attend to exactly k kv blocks
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks), k, dtype=torch.int32, device=device)
# Create the corresponding mask
block_sparse_mask = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
# Fill in the mask based on the indices
for b in range(bs):
for head in range(h):
for q_idx in range(num_q_blocks):
kv_indices = q2k_block_sparse_index[b, head, q_idx]
block_sparse_mask[b, head, q_idx, kv_indices] = True
# Create the reverse mapping (k2q)
# First, initialize lists to collect q indices for each kv block
k2q_indices_list = [[[] for _ in range(num_kv_blocks)] for _ in range(bs * h)]
# Populate the lists based on q2k mapping
for b in range(bs):
for head in range(h):
flat_idx = b * h + head
for q_idx in range(num_q_blocks):
kv_indices = q2k_block_sparse_index[b, head, q_idx].tolist()
for kv_idx in kv_indices:
k2q_indices_list[flat_idx][kv_idx].append(q_idx)
# Find the maximum number of q blocks that attend to any kv block
max_q_per_kv = 0
for flat_idx in range(bs * h):
for kv_idx in range(num_kv_blocks):
max_q_per_kv = max(max_q_per_kv, len(k2q_indices_list[flat_idx][kv_idx]))
# Create tensors for k2q mapping
k2q_block_sparse_index = torch.full((bs, h, num_kv_blocks, max_q_per_kv), -1,
dtype=torch.int32, device=device)
k2q_block_sparse_num = torch.zeros((bs, h, num_kv_blocks),
dtype=torch.int32, device=device)
# Fill the tensors
for b in range(bs):
for head in range(h):
flat_idx = b * h + head
for kv_idx in range(num_kv_blocks):
q_indices = k2q_indices_list[flat_idx][kv_idx]
num_q = len(q_indices)
k2q_block_sparse_num[b, head, kv_idx] = num_q
if num_q > 0:
k2q_block_sparse_index[b, head, kv_idx, :num_q] = torch.tensor(
q_indices, dtype=torch.int32, device=device)
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
def benchmark_block_sparse_attention(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops):
"""Benchmark block sparse attention forward and backward passes."""
print("\n=== BLOCK SPARSE ATTENTION BENCHMARK ===")
# Forward pass
# Warm-up run
o, l_vec = block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
torch.cuda.synchronize()
# Benchmark forward
_, fwd_time = benchmark_forward(
block_sparse_attention_fwd,
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num,
repeats=20,
verbose=False,
desc='Block Sparse Forward'
)
sparse_tflops = flops / fwd_time.mean * 1e-12
print(f"Block Sparse Forward - TFLOPS: {sparse_tflops:.2f}")
# Backward pass
grad_output = torch.randn_like(o)
# Warm-up runs
for _ in range(5):
block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
torch.cuda.synchronize()
# Benchmark backward
_, bwd_time = benchmark_forward(
block_sparse_attention_backward,
q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num,
repeats=20,
verbose=False,
desc='Block Sparse Backward'
)
bwd_flops = 2.5 * flops # Approximation
sparse_bwd_tflops = bwd_flops / bwd_time.mean * 1e-12
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd_tflops:.2f}")
return sparse_tflops, sparse_bwd_tflops
def main():
args = parse_arguments()
set_seed(42)
# Extract parameters
batch = args.batch_size
head = args.num_heads
headdim = args.head_dim
print(f"Block Sparse Attention Benchmark")
print(f"batch: {batch}, head: {head}, headdim: {headdim}")
# Test with different sequence lengths
for seq_len in args.seq_lengths:
# Skip very long sequences if they might cause OOM
if seq_len > 16384 and batch > 1:
continue
print("="*100)
print(f"\nSequence length: {seq_len}")
# Calculate theoretical FLOPs for attention
flops = 4 * batch * head * headdim * seq_len * seq_len
# Create input tensors
q, k, v = create_input_tensors(batch, head, seq_len, headdim)
# Setup block sparse parameters
num_q_blocks = seq_len // BLOCK_M
num_kv_blocks = seq_len // BLOCK_N
# Determine k value (number of kv blocks per q block)
topk = args.topk
if topk is None:
topk = num_kv_blocks // 10 # Default to ~90% sparsity if k is not specified
topk = max(1, topk)
print(f"Using topk={topk} kv blocks per q block (out of {num_kv_blocks} total kv blocks)")
# Generate block sparse pattern
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, _ = generate_block_sparse_pattern(
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
# Benchmark block sparse attention
sparse_fwd, sparse_bwd = benchmark_block_sparse_attention(
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops
)
# Print results
print("\n=== PERFORMANCE RESULTS ===")
print(f"Block Sparse Forward - TFLOPS: {sparse_fwd:.2f}")
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd:.2f}")
if __name__ == "__main__":
main()
-15
View File
@@ -1,15 +0,0 @@
### ADD TO THIS TO REGISTER NEW KERNELS
sources = {
'st_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 = ['st_attn']
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
target = 'h100'
-15
View File
@@ -1,15 +0,0 @@
### ADD TO THIS TO REGISTER NEW KERNELS
sources = {
'block_sparse': {
'source_files': {
'h100': 'vsa/block_sparse_h100.cu'
}
}
}
### WHICH KERNELS DO WE WANT TO BUILD?
# (oftentimes during development work you don't need to redefine them all.)
kernels = ['block_sparse']
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
target = 'h100'
-76
View File
@@ -1,76 +0,0 @@
import os
import subprocess
from csrc.attn.config_sta 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"])
-76
View File
@@ -1,76 +0,0 @@
import os
import subprocess
from csrc.attn.config_vsa 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 = "vsa"
VERSION = "0.0.1"
AUTHOR = "Hao AI Lab"
DESCRIPTION = "Video Sparse Attention Kernel Used in FastVideo"
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn"
# 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('vsa 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 = ['vsa.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('vsa_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"])
-23
View File
@@ -1,23 +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_ST_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_ST_ATTN
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention, assuming tile size is (6,8,8)");
#endif
}
-49
View File
@@ -1,49 +0,0 @@
import math
import torch
from torch.utils.checkpoint import detach_variable
try:
from st_attn_cuda import sta_fwd
except ImportError:
sta_fwd = None
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
seq_length = q_all.shape[2]
dit_seq_shape_mapping = {
'30x48x80':1,
'36x48x48':2,
'18x48x80':3,
}
if has_text:
assert q_all.shape[
2] >= 115200 and q_all.shape[2] <= 115456, f"Unsupported {dit_seq_shape}, current shape is {q_all.shape}, only support '30x48x80' for HunyuanVideo"
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 dit_seq_shape == '36x48x48': # Stepvideo 204x768x68
assert q_all.shape[2] == 82944
elif dit_seq_shape == '18x48x80': # Wan 69x768x1280
assert q_all.shape[2] == 69120
else:
raise ValueError(f"Unsupported {dit_seq_shape}, current shape is {q_all.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
kernel_aspect_ratio_flag = dit_seq_shape_mapping[dit_seq_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]
-289
View File
@@ -1,289 +0,0 @@
import torch
import argparse
from flash_attn.utils.benchmark import benchmark_forward
from flash_attn import flash_attn_func
from vsa import block_sparse_attention_fwd, block_sparse_attention_backward, BlockSparseAttentionFunction
from vsa import BLOCK_M, BLOCK_N
import numpy as np
import random
import gc
def set_seed(seed: int = 42):
# Python random module
random.seed(seed)
# NumPy
np.random.seed(seed)
# PyTorch
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed) # if using multi-GPU
@torch.no_grad
def precision_metric(quant_o, fa2_o):
x, xx = quant_o.float(), fa2_o.float()
sim = torch.nn.functional.cosine_similarity(x.reshape(1, -1), xx.reshape(1, -1)).item()
l1 = ((x - xx).abs().sum() / xx.abs().sum() ).item()
rmse = torch.sqrt(torch.mean((x -xx) ** 2)).item()
return sim, l1, rmse
def create_input_tensors(batch, head, seq_len, headdim):
"""Create random input tensors for attention."""
q = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
k = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
v = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
return q, k, v
def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device="cuda"):
"""
Generate a block sparse pattern where each q block attends to exactly k kv blocks.
Args:
bs: batch size
h: number of heads
num_q_blocks: number of query blocks
num_kv_blocks: number of key-value blocks
k: number of kv blocks each q block attends to
device: device to create tensors on
Returns:
q2k_block_sparse_index: [bs, h, num_q_blocks, k]
Contains the indices of kv blocks that each q block attends to.
q2k_block_sparse_num: [bs, h, num_q_blocks]
Contains the number of kv blocks that each q block attends to (all equal to k).
k2q_block_sparse_index: [bs, h, num_kv_blocks, num_q_blocks]
Contains the indices of q blocks that attend to each kv block.
k2q_block_sparse_num: [bs, h, num_kv_blocks]
Contains the number of q blocks that attend to each kv block.
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks]
Binary mask where 1 indicates attention connection.
"""
# Ensure k is not larger than num_kv_blocks
k = min(k, num_kv_blocks)
# Create random scores for sampling
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device=device)
# Get top-k indices for each q block
_, q2k_block_sparse_index = torch.topk(scores, k, dim=-1)
q2k_block_sparse_index = q2k_block_sparse_index.to(torch.int32)
# sort q2k_block_sparse_index
q2k_block_sparse_index, _ = torch.sort(q2k_block_sparse_index, dim=-1)
# All q blocks attend to exactly k kv blocks
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks), k, dtype=torch.int32, device=device)
# Create the corresponding mask
block_sparse_mask = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
# Fill in the mask based on the indices
for b in range(bs):
for head in range(h):
for q_idx in range(num_q_blocks):
kv_indices = q2k_block_sparse_index[b, head, q_idx]
block_sparse_mask[b, head, q_idx, kv_indices] = True
# Create the reverse mapping (k2q)
# First, initialize lists to collect q indices for each kv block
k2q_indices_list = [[[] for _ in range(num_kv_blocks)] for _ in range(bs * h)]
# Populate the lists based on q2k mapping
for b in range(bs):
for head in range(h):
flat_idx = b * h + head
for q_idx in range(num_q_blocks):
kv_indices = q2k_block_sparse_index[b, head, q_idx].tolist()
for kv_idx in kv_indices:
k2q_indices_list[flat_idx][kv_idx].append(q_idx)
# Find the maximum number of q blocks that attend to any kv block
max_q_per_kv = 0
for flat_idx in range(bs * h):
for kv_idx in range(num_kv_blocks):
max_q_per_kv = max(max_q_per_kv, len(k2q_indices_list[flat_idx][kv_idx]))
# Create tensors for k2q mapping
k2q_block_sparse_index = torch.full((bs, h, num_kv_blocks, max_q_per_kv), -1,
dtype=torch.int32, device=device)
k2q_block_sparse_num = torch.zeros((bs, h, num_kv_blocks),
dtype=torch.int32, device=device)
# Fill the tensors
for b in range(bs):
for head in range(h):
flat_idx = b * h + head
for kv_idx in range(num_kv_blocks):
q_indices = k2q_indices_list[flat_idx][kv_idx]
num_q = len(q_indices)
k2q_block_sparse_num[b, head, kv_idx] = num_q
if num_q > 0:
k2q_block_sparse_index[b, head, kv_idx, :num_q] = torch.tensor(
q_indices, dtype=torch.int32, device=device)
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
def main(args):
set_seed(42)
# Extract parameters
batch = args.batch_size
head = args.num_heads
headdim = args.head_dim
num_iterations = args.num_iterations
print(f"Block Sparse Attention Benchmark")
print(f"batch: {batch}, head: {head}, headdim: {headdim}, iterations: {num_iterations}")
# Test with different sequence lengths
for seq_len in args.seq_lengths:
# Skip very long sequences if they might cause OOM
# if seq_len > 16384 and batch > 1:
# continue
print("="*100)
print(f"\nSequence length: {seq_len}")
# Collect metrics across iterations
forward_metrics = {'sim': [], 'l1': [], 'rmse': []}
grad_q_metrics = {'sim': [], 'l1': [], 'rmse': []}
grad_k_metrics = {'sim': [], 'l1': [], 'rmse': []}
grad_v_metrics = {'sim': [], 'l1': [], 'rmse': []}
for iter_idx in range(num_iterations):
if num_iterations > 1:
print(f"\nIteration {iter_idx+1}/{num_iterations}")
# Create input tensors
q, k, v = create_input_tensors(batch, head, seq_len, headdim)
# Setup block sparse parameters
num_q_blocks = seq_len // BLOCK_M
num_kv_blocks = seq_len // BLOCK_N
# Determine k value (number of kv blocks per q block)
topk = args.topk
if topk is None:
topk = num_kv_blocks // 10 # Default to ~90% sparsity if k is not specified
topk = max(1, topk)
if iter_idx == 0: # Only print this once
print(f"Using topk={topk} kv blocks per q block (out of {num_kv_blocks} total kv blocks)")
# Generate block sparse pattern
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask = generate_block_sparse_pattern(
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
# expand block_sparse_mask to full mask
block_mask_expanded = block_sparse_mask.unsqueeze(-1).unsqueeze(-2) # [b, h, num_q_blocks, num_kv_blocks, 1, 1]
block_mask_expanded = block_mask_expanded.expand(-1, -1, -1, -1, BLOCK_M, BLOCK_N) # [b, h, num_q_blocks, num_kv_blocks, BLOCK_M, BLOCK_N]
full_mask = block_mask_expanded.permute(0, 1, 2, 4, 3, 5).reshape(batch, head, seq_len, seq_len)
q.requires_grad = True
k.requires_grad = True
v.requires_grad = True
# testing forward
o = BlockSparseAttentionFunction.apply(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
del q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask, block_mask_expanded
grad_o = torch.randn_like(o)
o.backward(grad_o)
# clear memory
q_sdpa = q.detach().clone()
k_sdpa = k.detach().clone()
v_sdpa = v.detach().clone()
q_sdpa.requires_grad = True
k_sdpa.requires_grad = True
v_sdpa.requires_grad = True
q.data = torch.empty(0, device=q.device)
k.data = torch.empty(0, device=k.device)
v.data = torch.empty(0, device=v.device)
torch.cuda.empty_cache()
o_sdpa = torch.nn.functional.scaled_dot_product_attention(q_sdpa, k_sdpa, v_sdpa, attn_mask=full_mask)
sim, l1, rmse = precision_metric(o, o_sdpa)
assert sim > 0.9999, f"SSIM too low: {sim}"
assert l1 < 8e-5, f"l1 too large: {l1}"
assert rmse < 2e-5, f"RMSE too large: {rmse}"
forward_metrics['sim'].append(sim)
forward_metrics['l1'].append(l1)
forward_metrics['rmse'].append(rmse)
print(f"block_sparse_attention_fwd vs torch.nn.functional.scaled_dot_product_attention:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
# test backward
o_sdpa.backward(grad_o)
sim, l1, rmse = precision_metric(q.grad, q_sdpa.grad)
# Error bounds collected on H100
assert sim > 0.9999, f"SSIM too low: {sim}"
assert l1 < 4e-3, f"l1 too large: {l1}"
assert rmse < 3e-4, f"RMSE too large: {rmse}"
grad_q_metrics['sim'].append(sim)
grad_q_metrics['l1'].append(l1)
grad_q_metrics['rmse'].append(rmse)
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_q:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
sim, l1, rmse = precision_metric(k.grad, k_sdpa.grad)
assert sim > 0.9999, f"SSIM too low: {sim}"
assert l1 < 4e-3, f"l1 too large: {l1}"
assert rmse < 2e-4, f"RMSE too large: {rmse}"
grad_k_metrics['sim'].append(sim)
grad_k_metrics['l1'].append(l1)
grad_k_metrics['rmse'].append(rmse)
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_k:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
sim, l1, rmse = precision_metric(v.grad, v_sdpa.grad)
assert sim > 0.9999, f"SSIM too low: {sim}"
assert l1 < 1e-4, f"l1 too large: {l1}"
assert rmse < 2e-5, f"RMSE too large: {rmse}"
grad_v_metrics['sim'].append(sim)
grad_v_metrics['l1'].append(l1)
grad_v_metrics['rmse'].append(rmse)
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_v:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
del o, o_sdpa, grad_o, q_sdpa, k_sdpa, v_sdpa
gc.collect()
torch.cuda.empty_cache()
# Print summary statistics if multiple iterations were run
if num_iterations > 1:
print("\n" + "="*50)
print(f"Summary Statistics (over {num_iterations} iterations):")
print("\nForward metrics:")
print(f"Similarity: mean={np.mean(forward_metrics['sim']):.6f}, std={np.std(forward_metrics['sim']):.6f}, min={np.min(forward_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(forward_metrics['l1']):.6f}, std={np.std(forward_metrics['l1']):.6f}, max={np.max(forward_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(forward_metrics['rmse']):.6f}, std={np.std(forward_metrics['rmse']):.6f}, max={np.max(forward_metrics['rmse']):.6f}")
print("\nGradient Q metrics:")
print(f"Similarity: mean={np.mean(grad_q_metrics['sim']):.6f}, std={np.std(grad_q_metrics['sim']):.6f}, min={np.min(grad_q_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(grad_q_metrics['l1']):.6f}, std={np.std(grad_q_metrics['l1']):.6f}, max={np.max(grad_q_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(grad_q_metrics['rmse']):.6f}, std={np.std(grad_q_metrics['rmse']):.6f}, max={np.max(grad_q_metrics['rmse']):.6f}")
print("\nGradient K metrics:")
print(f"Similarity: mean={np.mean(grad_k_metrics['sim']):.6f}, std={np.std(grad_k_metrics['sim']):.6f}, min={np.min(grad_k_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(grad_k_metrics['l1']):.6f}, std={np.std(grad_k_metrics['l1']):.6f}, max={np.max(grad_k_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(grad_k_metrics['rmse']):.6f}, std={np.std(grad_k_metrics['rmse']):.6f}, max={np.max(grad_k_metrics['rmse']):.6f}")
print("\nGradient V metrics:")
print(f"Similarity: mean={np.mean(grad_v_metrics['sim']):.6f}, std={np.std(grad_v_metrics['sim']):.6f}, min={np.min(grad_v_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(grad_v_metrics['l1']):.6f}, std={np.std(grad_v_metrics['l1']):.6f}, max={np.max(grad_v_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(grad_v_metrics['rmse']):.6f}, std={np.std(grad_v_metrics['rmse']):.6f}, max={np.max(grad_v_metrics['rmse']):.6f}")
if __name__ == "__main__":
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
parser.add_argument('--batch_size', type=int, default=4, help='Batch size')
parser.add_argument('--num_heads', type=int, default=6, help='Number of heads')
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
parser.add_argument('--topk', type=int, default=64, help='Number of kv blocks each q block attends to')
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[29120], help='Sequence lengths to benchmark')
parser.add_argument('--num_iterations', type=int, default=50, help='Number of test iterations to run')
args = parser.parse_args()
main(args)
-136
View File
@@ -1,136 +0,0 @@
import torch
from tqdm import tqdm
import matplotlib.pyplot as plt
import numpy as np
def pytorch_test(Q, K, V, dO):
q_ = Q.to(torch.float64).requires_grad_()
k_ = K.to(torch.float64).requires_grad_()
v_ = V.to(torch.float64).requires_grad_()
dO_ = dO.to(torch.float64)
# manual pytorch implementation of scaled dot product attention
QK = torch.matmul(q_, k_.transpose(-2, -1))
QK /= (q_.size(-1) ** 0.5)
# Causal mask removed since causal is always false
QK = torch.nn.functional.softmax(QK, dim=-1)
output = torch.matmul(QK, v_)
output.backward(dO_)
q_grad = q_.grad
k_grad = k_.grad
v_grad = v_.grad
return output, q_grad, k_grad, v_grad
def fa2_test(Q, K, V, dO):
Q.requires_grad = True
K.requires_grad = True
V.requires_grad = True
output = torch.nn.functional.scaled_dot_product_attention(Q, K, V, is_causal=False)
output.backward(dO)
return output, Q.grad, K.grad, V.grad
def generate_tensor(shape, mean, std, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
return scaled_tensor.contiguous()
def check_correctness(b, h, n, d, mean, std, num_iterations=100, error_mode='all', test_mode='forward_backward'):
results = {
'FA2 vs PT': {'sum_diff': 0, 'sum_abs': 0, 'max_diff': 0},
}
for _ in range(num_iterations):
torch.manual_seed(0)
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
dO = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, dO)
fa2_o, fa2_qg, fa2_kg, fa2_vg = fa2_test(Q, K, V, dO)
if test_mode == 'forward_only':
tensors_fa2_pt = [(pt_o, fa2_o)]
else: # 'forward_backward'
if error_mode == 'output':
tensors_fa2_pt = [(pt_o, fa2_o)]
elif error_mode == 'backward':
tensors_fa2_pt = [(pt_qg, fa2_qg),
(pt_kg, fa2_kg),
(pt_vg, fa2_vg)]
else: # 'all'
tensors_fa2_pt = [(pt_o, fa2_o),
(pt_qg, fa2_qg),
(pt_kg, fa2_kg),
(pt_vg, fa2_vg)]
for pt, fa2 in tensors_fa2_pt:
diff = pt - fa2
abs_diff = torch.abs(diff)
results['FA2 vs PT']['sum_diff'] += torch.sum(abs_diff).item()
results['FA2 vs PT']['sum_abs'] += torch.sum(torch.abs(pt)).item()
results['FA2 vs PT']['max_diff'] = max(results['FA2 vs PT']['max_diff'], torch.max(abs_diff).item())
torch.cuda.empty_cache()
# Calculate total elements based on test mode and error mode
if test_mode == 'forward_only':
total_elements = b * h * n * d * num_iterations
else: # 'forward_backward'
total_elements = b * h * n * d * num_iterations * (1 if error_mode == 'output' else 3 if error_mode == 'backward' else 4)
for name, data in results.items():
avg_diff = data['sum_diff'] / total_elements
max_diff = data['max_diff']
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
return results
def generate_error_tables(b, h, d, mean, std, error_mode='all', test_mode='forward_backward'):
seq_lengths = [768 * (2**i) for i in range(1)]
print(f"\n{'='*80}")
print(f"ATTENTION ERROR COMPARISON TABLE (b={b}, h={h}, d={d}, mean={mean}, std={std})")
print(f"Mode: {error_mode}, Test: {test_mode}")
print(f"{'='*80}")
# Print header
print(f"{'Seq Length':<12} | {'FA2 vs PT Avg':<15} | {'FA2 vs PT Max':<15}")
print(f"{'-'*12} | {'-'*15} | {'-'*15}")
for n in seq_lengths:
results = check_correctness(b, h, n, d, mean, std, error_mode=error_mode, test_mode=test_mode)
fa2_pt_avg = results['FA2 vs PT']['avg_diff']
fa2_pt_max = results['FA2 vs PT']['max_diff']
# Print row
print(f"{n:<12} | {fa2_pt_avg:<15.6e} | {fa2_pt_max:<15.6e}")
print(f"{'='*80}\n")
# fix random seed
torch.manual_seed(0)
# Example usage
b, h, d = 2, 2, 64
mean = 1e-1
std = 10
# Test forward only
generate_error_tables(b, h, d, mean, std, error_mode='output', test_mode='forward_only')
# Test forward and backward
generate_error_tables(b, h, d, mean, std, error_mode='all', test_mode='forward_backward')
print("Attention error comparison completed.")
-175
View File
@@ -1,175 +0,0 @@
import torch
from flash_attn_interface import flash_attn_func
from st_attn import mha_forward, mha_backward
import random
from tqdm import tqdm
import matplotlib.pyplot as plt
import numpy as np
def pytorch_test(Q, K, V, dO):
q_ = Q.to(torch.float64).requires_grad_()
k_ = K.to(torch.float64).requires_grad_()
v_ = V.to(torch.float64).requires_grad_()
dO_ = dO.to(torch.float64)
# manual pytorch implementation of scaled dot product attention
QK = torch.matmul(q_, k_.transpose(-2, -1))
QK /= (q_.size(-1) ** 0.5)
# Causal mask removed since causal is always false
QK = torch.nn.functional.softmax(QK, dim=-1)
output = torch.matmul(QK, v_)
output.backward(dO_)
q_grad = q_.grad
k_grad = k_.grad
v_grad = v_.grad
return output, q_grad, k_grad, v_grad
def fa2_test(Q, K, V, dO):
Q.requires_grad = True
K.requires_grad = True
V.requires_grad = True
output = torch.nn.functional.scaled_dot_product_attention(Q, K, V, is_causal=False)
output.backward(dO)
return output, Q.grad, K.grad, V.grad
def mha_kernel_test(Q, K, V, dO, mode):
Q.requires_grad = True
K.requires_grad = True
V.requires_grad = True
o, l_vec = mha_forward(Q, K, V)
if mode == 'forward_only':
return o, None, None, None
else: # 'forward_backward'
qg, kg, vg = mha_backward(Q, K, V, o, l_vec, dO)
return o, qg, kg, vg
def generate_tensor(shape, mean, std, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
return scaled_tensor.contiguous()
def check_correctness(b, h, n, d, mean, std, num_iterations=100, error_mode='all', test_mode='forward_backward'):
results = {
'MHA vs PT': {'sum_diff': 0, 'sum_abs': 0, 'max_diff': 0},
'FA2 vs PT': {'sum_diff': 0, 'sum_abs': 0, 'max_diff': 0},
}
for _ in range(num_iterations):
torch.manual_seed(0)
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
dO = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, dO)
fa2_o, fa2_qg, fa2_kg, fa2_vg = fa2_test(Q, K, V, dO)
if test_mode == 'forward_only':
mha_o, _, _, _ = mha_kernel_test(Q, K, V, dO, 'forward_only')
tensors_mha_pt = [(pt_o, mha_o)]
tensors_fa2_pt = [(pt_o, fa2_o)]
else: # 'forward_backward'
mha_o, mha_qg, mha_kg, mha_vg = mha_kernel_test(Q, K, V, dO, 'forward_backward')
if error_mode == 'output':
tensors_mha_pt = [(pt_o, mha_o)]
tensors_fa2_pt = [(pt_o, fa2_o)]
elif error_mode == 'backward':
tensors_mha_pt = [(pt_qg, mha_qg),
(pt_kg, mha_kg),
(pt_vg, mha_vg)]
tensors_fa2_pt = [(pt_qg, fa2_qg),
(pt_kg, fa2_kg),
(pt_vg, fa2_vg)]
else: # 'all'
tensors_mha_pt = [(pt_o, mha_o),
(pt_qg, mha_qg),
(pt_kg, mha_kg),
(pt_vg, mha_vg)]
tensors_fa2_pt = [(pt_o, fa2_o),
(pt_qg, fa2_qg),
(pt_kg, fa2_kg),
(pt_vg, fa2_vg)]
for pt, mha in tensors_mha_pt:
diff = pt - mha
abs_diff = torch.abs(diff)
results['MHA vs PT']['sum_diff'] += torch.sum(abs_diff).item()
results['MHA vs PT']['sum_abs'] += torch.sum(torch.abs(pt)).item()
results['MHA vs PT']['max_diff'] = max(results['MHA vs PT']['max_diff'], torch.max(abs_diff).item())
for pt, fa2 in tensors_fa2_pt:
diff = pt - fa2
abs_diff = torch.abs(diff)
results['FA2 vs PT']['sum_diff'] += torch.sum(abs_diff).item()
results['FA2 vs PT']['sum_abs'] += torch.sum(torch.abs(pt)).item()
results['FA2 vs PT']['max_diff'] = max(results['FA2 vs PT']['max_diff'], torch.max(abs_diff).item())
torch.cuda.empty_cache()
# Calculate total elements based on test mode and error mode
if test_mode == 'forward_only':
total_elements = b * h * n * d * num_iterations
else: # 'forward_backward'
total_elements = b * h * n * d * num_iterations * (1 if error_mode == 'output' else 3 if error_mode == 'backward' else 4)
for name, data in results.items():
avg_diff = data['sum_diff'] / total_elements
max_diff = data['max_diff']
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
return results
def generate_error_tables(b, h, d, mean, std, error_mode='all', test_mode='forward_backward'):
seq_lengths = [768 * (2**i) for i in range(1)]
print(f"\n{'='*80}")
print(f"MHA ERROR COMPARISON TABLE (b={b}, h={h}, d={d}, mean={mean}, std={std})")
print(f"Mode: {error_mode}, Test: {test_mode}")
print(f"{'='*80}")
# Print header
print(f"{'Seq Length':<12} | {'MHA vs PT Avg':<15} | {'MHA vs PT Max':<15} | {'FA2 vs PT Avg':<15} | {'FA2 vs PT Max':<15}")
print(f"{'-'*12} | {'-'*15} | {'-'*15} | {'-'*15} | {'-'*15}")
for n in seq_lengths:
results = check_correctness(b, h, n, d, mean, std, error_mode=error_mode, test_mode=test_mode)
mha_pt_avg = results['MHA vs PT']['avg_diff']
mha_pt_max = results['MHA vs PT']['max_diff']
fa2_pt_avg = results['FA2 vs PT']['avg_diff']
fa2_pt_max = results['FA2 vs PT']['max_diff']
# Print row
print(f"{n:<12} | {mha_pt_avg:<15.6e} | {mha_pt_max:<15.6e} | {fa2_pt_avg:<15.6e} | {fa2_pt_max:<15.6e}")
print(f"{'='*80}\n")
# fix random seed
torch.manual_seed(0)
# Example usage
b, h, d = 2, 2, 64
mean = 1e-1
std = 10
# Test forward only
generate_error_tables(b, h, d, mean, std, error_mode='output', test_mode='forward_only')
# Test forward and backward
generate_error_tables(b, h, d, mean, std, error_mode='all', test_mode='forward_backward')
print("MHA attention error comparison completed.")
Submodule csrc/attn/tk deleted from 1719fb7264
-27
View File
@@ -1,27 +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_BLOCK_SPARSE
extern std::vector<torch::Tensor> block_sparse_attention_forward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q2k_block_sparse_index, torch::Tensor q2k_block_sparse_num
);
extern std::vector<torch::Tensor> block_sparse_attention_backward(
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, torch::Tensor l_vec, torch::Tensor og, torch::Tensor k2q_block_sparse_index, torch::Tensor k2q_block_sparse_num
);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "Video Sparse Attention Kernels"; // optional module docstring
#ifdef TK_COMPILE_BLOCK_SPARSE
m.def("block_sparse_fwd", torch::wrap_pybind_function(block_sparse_attention_forward), "block sparse attention");
m.def("block_sparse_bwd", torch::wrap_pybind_function(block_sparse_attention_backward), "block sparse attention backward");
#endif
}
-470
View File
@@ -1,470 +0,0 @@
import math
import torch
from torch.utils.checkpoint import detach_variable
from typing import Tuple
try:
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
except ImportError:
block_sparse_fwd = None
block_sparse_bwd = None
BLOCK_M = 64
BLOCK_N = 64
def video_sparse_attn(q, k, v, topk, block_size, compress_attn_weight=None):
"""
q: [batch_size, num_heads, seq_len, head_dim]
k: [batch_size, num_heads, seq_len, head_dim]
v: [batch_size, num_heads, seq_len, head_dim]
topk: int
block_size: int or tuple of 3 ints
video_shape: tuple of (T, H, W)
compress_attn_weight: [batch_size, num_heads, seq_len, head_dim]
select_attn_weight: [batch_size, num_heads, seq_len, head_dim]
V1 of sparse attention. Include compress attn and sparse attn branch, use average pooling to compress.
Assume q, k, v is flattened in this way: [batch_size, num_heads, T//block_size[0], H//block_size[1], W//block_size[2], block_size[0], block_size[1], block_size[2]]
"""
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
block_elements = block_size[0] * block_size[1] * block_size[2]
assert block_elements % 64 == 0 and block_elements >= 64
assert q.shape[2] % block_elements == 0
batch_size, num_heads, seq_len, head_dim = q.shape
# compress attn
q_compress = q.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).mean(dim=3)
k_compress = k.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).mean(dim=3)
v_compress = v.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).mean(dim=3)
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
v_compress)
output_compress = output_compress.view(batch_size, num_heads,
seq_len // block_elements, 1,
head_dim)
output_compress = output_compress.repeat(1, 1, 1, block_elements,
1).view(batch_size, num_heads,
seq_len, head_dim)
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num = generate_topk_block_sparse_pattern(
block_attn_score, topk)
output_select = block_sparse_attn(q, k, v, q2k_block_sparse_index,
q2k_block_sparse_num,
k2q_block_sparse_index,
k2q_block_sparse_num)
if compress_attn_weight is not None:
final_output = output_compress * compress_attn_weight + output_select
else:
final_output = output_compress + output_select
return final_output
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
QK = torch.matmul(q, k.transpose(-2, -1))
QK /= (q.size(-1)**0.5)
# Causal mask removed since causal is always false
QK = torch.nn.functional.softmax(QK, dim=-1)
output = torch.matmul(QK, v)
return output, QK
def generate_topk_block_sparse_pattern(block_attn_score: torch.Tensor,
topk: int):
"""
Generate a block sparse pattern where each q block attends to exactly topk kv blocks,
based on the provided attention scores.
Args:
block_attn_score: [bs, h, num_q_blocks, num_kv_blocks]
Attention scores between query and key blocks
topk: int
Number of kv blocks each q block attends to
Returns:
q2k_block_sparse_index: [bs, h, num_q_blocks, topk]
Contains the indices of kv blocks that each q block attends to.
q2k_block_sparse_num: [bs, h, num_q_blocks]
Contains the number of kv blocks that each q block attends to (all equal to topk).
k2q_block_sparse_index: [bs, h, num_kv_blocks, max_q_per_kv]
Contains the indices of q blocks that attend to each kv block.
k2q_block_sparse_num: [bs, h, num_kv_blocks]
Contains the number of q blocks that attend to each kv block.
"""
device = block_attn_score.device
# Extract dimensions from block_attn_score
bs, h, num_q_blocks, num_kv_blocks = block_attn_score.shape
sorted_result = torch.sort(block_attn_score, dim=-1, descending=True)
sorted_indice = sorted_result.indices
q2k_block_sparse_index, _ = torch.sort(sorted_indice[:, :, :, :topk],
dim=-1)
q2k_block_sparse_index = q2k_block_sparse_index.to(dtype=torch.int32)
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks),
topk,
device=device,
dtype=torch.int32)
block_map = topk_index_to_map(q2k_block_sparse_index,
num_kv_blocks,
transpose_map=True)
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(
block_map.transpose(2, 3))
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
@torch._dynamo.disable
def block_sparse_attn(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
"""
Differentiable block sparse attention function.
Args:
q: Query tensor [batch_size, num_heads, seq_len_q, head_dim]
k: Key tensor [batch_size, num_heads, seq_len_kv, head_dim]
v: Value tensor [batch_size, num_heads, seq_len_kv, head_dim]
q2k_block_sparse_index: Indices for query-to-key sparse blocks
q2k_block_sparse_num: Number of sparse blocks for each query block
k2q_block_sparse_index: Indices for key-to-query sparse blocks (for backward pass)
k2q_block_sparse_num: Number of sparse blocks for each key block (for backward pass)
Returns:
output: Attention output tensor [batch_size, num_heads, seq_len_q, head_dim]
"""
return BlockSparseAttentionFunction.apply(
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
)
def block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num):
"""
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks].
[*, *, i, j] = 1 means the i-th q block should attend to the j-th kv block.
"""
# assert all elements in q2k_block_sparse_num can be devisible by 2
o, lse = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
return o, lse
def block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num):
grad_output = grad_output.contiguous()
grad_q, grad_k, grad_v = block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
return grad_q, grad_k, grad_v
## pytorch sdpa version of block sparse ##
import triton
import triton.language as tl
@triton.jit
def index_to_mask_kernel(
q2k_block_sparse_index_ptr,
q2k_block_sparse_num_ptr,
mask_ptr,
batch_size: tl.constexpr,
num_heads: tl.constexpr,
num_q_blocks: tl.constexpr,
num_k_blocks: tl.constexpr,
max_kv_blocks: tl.constexpr,
BLOCK_Q: tl.constexpr,
BLOCK_K: tl.constexpr,
):
bh, q, id = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64), tl.program_id(2).to(tl.int64)
b = bh // num_heads
h = bh % num_heads
num_valid_blocks = tl.load(q2k_block_sparse_num_ptr + b * num_heads * num_q_blocks + h * num_q_blocks + q)
if num_valid_blocks <= id:
return
k = tl.load(q2k_block_sparse_index_ptr + b * num_heads * num_q_blocks * max_kv_blocks + h * num_q_blocks * max_kv_blocks + q * max_kv_blocks + id)
full_mask = (tl.arange(0, BLOCK_Q)[:, None] < BLOCK_Q) & (tl.arange(0, BLOCK_K)[None, :] < BLOCK_K)
q_lengths = num_q_blocks * BLOCK_Q
k_lengths = num_k_blocks * BLOCK_K
mask_ptr_base = mask_ptr + b * num_heads * q_lengths * k_lengths + h * q_lengths * k_lengths + q * BLOCK_Q * k_lengths + k * BLOCK_K
tl.store(mask_ptr_base + tl.arange(0, BLOCK_Q)[:, None] * k_lengths + tl.arange(0, BLOCK_K)[None, :], full_mask)
def index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, BLOCK_Q, BLOCK_K, num_k_blocks):
"""
Convert block sparse indices to a mask.
Args:
q2k_block_sparse_index: Indices for query-to-key sparse blocks
q2k_block_sparse_num: Number of sparse blocks for each query block
Returns:
mask: Block sparse mask tensor
"""
batch_size, num_heads, num_q_blocks, max_kv_blocks = q2k_block_sparse_index.shape
assert q2k_block_sparse_num.shape == (batch_size, num_heads, num_q_blocks)
mask = torch.zeros((batch_size, num_heads, num_q_blocks * BLOCK_Q, num_k_blocks * BLOCK_K), dtype=torch.bool, device=q2k_block_sparse_index.device)
grid = (batch_size * num_heads, num_q_blocks, max_kv_blocks)
index_to_mask_kernel[grid](
q2k_block_sparse_index,
q2k_block_sparse_num,
mask,
batch_size,
num_heads,
num_q_blocks,
num_k_blocks,
max_kv_blocks,
BLOCK_Q=BLOCK_Q,
BLOCK_K=BLOCK_K,
)
return mask
@triton.jit
def topk_index_to_map_kernel(
map_ptr,
index_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
topk: tl.constexpr,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
for i in tl.static_range(topk):
index = tl.load(index_ptr_base + i * index_kv_stride)
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
@triton.jit
def map_to_index_kernel(
map_ptr,
index_ptr,
index_num_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
index_num_bs_stride,
index_num_h_stride,
index_num_q_stride,
num_kv_blocks: tl.constexpr,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
num = 0
for i in tl.static_range(num_kv_blocks):
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
if map_entry:
tl.store(index_ptr_base + num * index_kv_stride, i)
num += 1
tl.store(
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
q * index_num_q_stride, num)
def topk_index_to_map(index: torch.Tensor,
num_kv_blocks: int,
transpose_map: bool = False):
"""
Convert topk indices to a map.
Args:
index: [bs, h, num_q_blocks, topk]
The topk indices tensor.
num_kv_blocks: int
The number of key-value blocks in the block_map returned
transpose_map: bool
If True, the block_map will be transposed on the final two dimensions.
Returns:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
A binary map where 1 indicates that the q block attends to the kv block.
"""
bs, h, num_q_blocks, topk = index.shape
if transpose_map is False:
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
dtype=torch.bool,
device=index.device)
else:
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
dtype=torch.bool,
device=index.device)
block_map = block_map.transpose(2, 3)
grid = (bs, h, num_q_blocks)
topk_index_to_map_kernel[grid](
block_map,
index,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
topk=topk,
)
return block_map
def map_to_index(block_map: torch.Tensor):
"""
Convert a block map to indices and counts.
Args:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
The block map tensor.
Returns:
index: [bs, h, num_q_blocks, num_kv_blocks]
The indices of the blocks.
index_num: [bs, h, num_q_blocks]
The number of blocks for each q block.
"""
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
index = torch.full((block_map.shape),
-1,
dtype=torch.int32,
device=block_map.device)
index_num = torch.empty((bs, h, num_q_blocks),
dtype=torch.int32,
device=block_map.device)
grid = (bs, h, num_q_blocks)
map_to_index_kernel[grid](
block_map,
index,
index_num,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
index_num.stride(0),
index_num.stride(1),
index_num.stride(2),
num_kv_blocks=num_kv_blocks,
)
return index, index_num
class BlockSparseAttentionFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
o, lse = block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
ctx.save_for_backward(q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num)
return o
@staticmethod
def backward(ctx, grad_output):
q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num = ctx.saved_tensors
grad_q, grad_k, grad_v = block_sparse_attention_backward(
q, k, v, o, lse, grad_output, k2q_block_sparse_index, k2q_block_sparse_num
)
return grad_q, grad_k, grad_v, None, None, None, None
class DummyOperator(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
return x
@staticmethod
def backward(ctx, grad_output):
return grad_output
class CheckpointSDPA(torch.autograd.Function):
@staticmethod
def forward(ctx, obj, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k):
"""Forward pass."""
with torch.no_grad():
mask = index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k, k.shape[2] // block_k)
outputs = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
ctx.save_for_backward(*detach_variable((q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)))
ctx.block_q = block_q
ctx.block_k = block_k
# the obj is passed in, then it can access the saved input
# tensors later for recomputation
obj.ctx = ctx
return outputs
@staticmethod
def backward(ctx, grad_output):
"""Backward pass."""
inputs = ctx.saved_tensors
output = ctx.output
torch.autograd.backward(output, grad_output)
ctx.output = None
grads = tuple(inp.grad for inp in inputs)
return (None, ) + grads + (None, None)
class BlockSparseAttnTorch:
def __init__(self):
self.ctx = None
def recompute_mask(self, _):
recomputed_mask = index_to_mask(self.q2k_block_sparse_index, self.q2k_block_sparse_num, self.block_q, self.block_k, self.num_kv_blocks)
mask_size = recomputed_mask.untyped_storage().size()
self.mask.untyped_storage().resize_(mask_size)
self.mask.untyped_storage().copy_(recomputed_mask.untyped_storage())
def recompute(self, _):
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num = self.ctx.saved_tensors
block_q = self.ctx.block_q
block_k = self.ctx.block_k
mask = index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k, k.shape[2] // block_k)
with torch.enable_grad():
output = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
self.ctx.output = output
self.ctx = None
@torch._dynamo.disable
def forward(self, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k):
"""
Differentiable block sparse attention function using PyTorch.
Args:
q: Query tensor [batch_size, num_heads, seq_len_q, head_dim]
k: Key tensor [batch_size, num_heads, seq_len_kv, head_dim]
v: Value tensor [batch_size, num_heads, seq_len_kv, head_dim]
q2k_block_sparse_index: Indices for query-to-key sparse blocks
q2k_block_sparse_num: Number of sparse blocks for each query block
block_q: Block size for query
block_k: Block size for key-value
Returns:
output: Attention output tensor [batch_size, num_heads, seq_len_q, head_dim]
"""
output = CheckpointSDPA.apply(
self, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k
)
o = DummyOperator.apply(output)
o.register_hook(self.recompute)
return o
+39 -21
View File
@@ -1,7 +1,9 @@
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
FROM nvidia/cuda:12.8.0-cudnn-devel-ubuntu22.04
ENV DEBIAN_FRONTEND=noninteractive
SHELL ["/bin/bash", "-c"]
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
@@ -9,17 +11,25 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
git \
ca-certificates \
openssh-server \
zsh \
vim \
curl \
gcc-11 \
g++-11 \
clang-11 \
&& rm -rf /var/lib/apt/lists/*
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
rm Miniconda3-latest-Linux-x86_64.sh
# Set up C++20 compilers for ThunderKittens
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
ENV PATH=/opt/conda/bin:$PATH
# Set CUDA environment variables
ENV CUDA_HOME=/usr/local/cuda-12.8
ENV PATH=${CUDA_HOME}/bin:${PATH}
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
RUN conda create --name fastvideo-dev python=3.10.0 -y
SHELL ["/bin/bash", "-c"]
# Install uv and source its environment
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject.toml ./
@@ -27,22 +37,30 @@ COPY pyproject.toml ./
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
conda clean -afy
# Create and activate virtual environment with specific Python version and seed
RUN source $HOME/.local/bin/env && \
uv venv --python 3.10 --seed /opt/venv && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp310-cp310-linux_x86_64.whl
COPY . .
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
# Install dependencies using uv and set up shell configuration
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir -e .[dev] && \
git config --unset-all http.https://github.com/.extraheader || true && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Remove authentication headers
RUN git config --unset-all http.https://github.com/.extraheader || true
# Install FastVideo Unified Kernel
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd fastvideo-kernel && \
git submodule update --init --recursive && \
./build.sh
# Set up automatic conda environment activation for all shells
RUN echo 'source /opt/conda/etc/profile.d/conda.sh' >> /root/.bashrc && \
echo 'conda activate fastvideo-dev' >> /root/.bashrc && \
# Ensure .bashrc is sourced for SSH login shells
echo 'if [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
EXPOSE 22
EXPOSE 22
+39 -21
View File
@@ -1,7 +1,9 @@
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
FROM nvidia/cuda:12.8.0-cudnn-devel-ubuntu22.04
ENV DEBIAN_FRONTEND=noninteractive
SHELL ["/bin/bash", "-c"]
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
@@ -9,17 +11,25 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
git \
ca-certificates \
openssh-server \
zsh \
vim \
curl \
gcc-11 \
g++-11 \
clang-11 \
&& rm -rf /var/lib/apt/lists/*
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
rm Miniconda3-latest-Linux-x86_64.sh
# Set up C++20 compilers for ThunderKittens
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
ENV PATH=/opt/conda/bin:$PATH
# Set CUDA environment variables
ENV CUDA_HOME=/usr/local/cuda-12.8
ENV PATH=${CUDA_HOME}/bin:${PATH}
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
RUN conda create --name fastvideo-dev python=3.11.11 -y
SHELL ["/bin/bash", "-c"]
# Install uv and source its environment
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject.toml ./
@@ -27,22 +37,30 @@ COPY pyproject.toml ./
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
conda clean -afy
# Create and activate virtual environment with specific Python version and seed
RUN source $HOME/.local/bin/env && \
uv venv --python 3.11 --seed /opt/venv && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp311-cp311-linux_x86_64.whl
COPY . .
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
# Install dependencies using uv and set up shell configuration
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir -e .[dev] && \
git config --unset-all http.https://github.com/.extraheader || true && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Remove authentication headers
RUN git config --unset-all http.https://github.com/.extraheader || true
# Install FastVideo Unified Kernel
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd fastvideo-kernel && \
git submodule update --init --recursive && \
./build.sh
# Set up automatic conda environment activation for all shells
RUN echo 'source /opt/conda/etc/profile.d/conda.sh' >> /root/.bashrc && \
echo 'conda activate fastvideo-dev' >> /root/.bashrc && \
# Ensure .bashrc is sourced for SSH login shells
echo 'if [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
EXPOSE 22
EXPOSE 22
+5 -12
View File
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir flash-attn==2.8.0.post2 --no-build-isolation
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp312-cp312-linux_x86_64.whl
COPY . .
@@ -55,18 +55,11 @@ RUN source $HOME/.local/bin/env && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install STA (Sliding Tile Attention)
# Install FastVideo Unified Kernel
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn && \
cd fastvideo-kernel && \
git submodule update --init --recursive && \
python setup_sta.py install
./build.sh
# Install VSA
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn && \
git submodule update --init --recursive && \
python setup_vsa.py install
EXPOSE 22
EXPOSE 22
+66
View File
@@ -0,0 +1,66 @@
FROM nvidia/cuda:12.9.1-cudnn-devel-ubuntu22.04
ENV DEBIAN_FRONTEND=noninteractive
SHELL ["/bin/bash", "-c"]
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
wget \
git \
ca-certificates \
openssh-server \
zsh \
vim \
curl \
gcc-11 \
g++-11 \
clang-11 \
&& rm -rf /var/lib/apt/lists/*
# Set up C++20 compilers for ThunderKittens
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Set CUDA environment variables
ENV CUDA_HOME=/usr/local/cuda-12.9
ENV PATH=${CUDA_HOME}/bin:${PATH}
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
# Install uv and source its environment
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject.toml ./
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
# Create and activate virtual environment with specific Python version and seed
RUN source $HOME/.local/bin/env && \
uv venv --python 3.12 --seed /opt/venv && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
COPY . .
# Install dependencies using uv and set up shell configuration
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir -e .[dev] && \
git config --unset-all http.https://github.com/.extraheader || true && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install FastVideo Unified Kernel
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd fastvideo-kernel && \
git submodule update --init --recursive && \
./build.sh
EXPOSE 22
+58
View File
@@ -0,0 +1,58 @@
FROM rocm/pytorch:rocm7.1_ubuntu22.04_py3.10_pytorch_release_2.9.1
ENV DEBIAN_FRONTEND=noninteractive
SHELL ["/bin/bash", "-c"]
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
wget \
git \
ca-certificates \
openssh-server \
zsh \
vim \
curl \
gcc-11 \
g++-11 \
clang-11 \
&& rm -rf /var/lib/apt/lists/*
# Set up C++20 compilers for ThunderKittens
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Install uv and source its environment
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject_other.toml ./pyproject.toml
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
# Create and activate virtual environment with specific Python version and seed
RUN source $HOME/.local/bin/env && \
uv venv --python 3.10 --seed /opt/venv && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip
COPY . .
# Install dependencies using uv and set up shell configuration
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir -e .[rocm] && \
git config --unset-all http.https://github.com/.extraheader || true && \
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
# Install FastVideo Unified Kernel
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd fastvideo-kernel && \
git submodule update --init --recursive && \
./build.sh --rocm
EXPOSE 22
-25
View File
@@ -1,25 +0,0 @@
# Minimal makefile for Sphinx documentation
#
# You can set these variables from the command line, and also
# from the environment for the first two.
SPHINXOPTS ?=
SPHINXBUILD ?= sphinx-build
SOURCEDIR = source
BUILDDIR = build
# Put it first so that "make" without argument is like "make help".
help:
@$(SPHINXBUILD) -M help "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
.PHONY: help Makefile
# Catch-all target: route all unknown targets to Sphinx using the new
# "make mode" option. $(O) is meant as a shortcut for $(SPHINXOPTS).
%: Makefile
@$(SPHINXBUILD) -M $@ "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
clean:
@$(SPHINXBUILD) -M clean "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
rm -rf "$(SOURCEDIR)/getting_started/examples"
rm -rf "$(SOURCEDIR)/inference/examples"
+29 -10
View File
@@ -1,20 +1,39 @@
# FastVideo documents
# FastVideo Documentation
## Build the docs
This directory contains the FastVideo documentation built with MkDocs.
## Build the docs locally
```bash
# Install dependencies.
pip install -r requirements-docs.txt
# Install dependencies
pip install -r requirements-mkdocs.txt
# Build the docs.
make clean
make html
# Serve docs with live reload (recommended for development)
mkdocs serve
# Or build static site
mkdocs build
```
## Open the docs with your browser
## View the docs
### Development server (with live reload)
```bash
python -m http.server -d build/html/
mkdocs serve
```
Launch your browser and open localhost:8000.
Then open your browser to: http://127.0.0.1:8000
### Static build
```bash
mkdocs build
python -m http.server -d site/
```
Then open your browser to: http://localhost:8000
## Automatic Deployment
Documentation is automatically built and deployed to GitHub Pages when changes are pushed to the `main` branch via the `.github/workflows/docs.yml` workflow.
+248
View File
@@ -0,0 +1,248 @@
# FastVideo API Reference
This page contains the complete API reference for the FastVideo library.
## fastvideo
### Modules
| Name | Description |
|------|-------------|
| [attention](#fastvideoattention) | Attention mechanisms and backends for video generation |
| [configs](#fastvideoconfigs) | Configuration classes for pipelines, models, and sampling |
| [distributed](#fastvideodistributed) | Distributed execution and communication utilities |
| [entrypoints](#fastvideoentrypoints) | Main API entry points for video generation |
| [models](#fastvideomodels) | Model implementations (transformers, VAEs, schedulers) |
| [pipelines](#fastvideopipelines) | Core pipeline classes for video diffusion |
| [training](#fastvideotraining) | Training utilities and helpers |
| [workflow](#fastvideoworkflow) | Workflow management and orchestration |
| [dataset](#fastvideodataset) | Dataset handling and preprocessing |
| [layers](#fastvideolayers) | Custom neural network layers |
| [platforms](#fastvideoplatforms) | Platform-specific implementations |
| [utils](#fastvideoutils) | Utility functions and helpers |
| [worker](#fastvideoworker) | Execution workers for video generation |
## fastvideo.attention
::: fastvideo.attention
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
show_submodules: true
heading_level: 3
## fastvideo.configs
::: fastvideo.configs
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
show_submodules: true
heading_level: 3
### Submodules
#### fastvideo.configs.pipelines
::: fastvideo.configs.pipelines
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
heading_level: 4
#### fastvideo.configs.models
::: fastvideo.configs.models
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
heading_level: 4
#### fastvideo.configs.sample
::: fastvideo.configs.sample
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
heading_level: 4
## fastvideo.distributed
::: fastvideo.distributed
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
show_submodules: true
heading_level: 3
## fastvideo.entrypoints
::: fastvideo.entrypoints
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
show_submodules: true
heading_level: 3
## fastvideo.models
::: fastvideo.models
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
show_submodules: true
heading_level: 3
### Submodules
#### fastvideo.models.registry
::: fastvideo.models.registry
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
heading_level: 4
#### fastvideo.models.loader
::: fastvideo.models.loader
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
heading_level: 4
## fastvideo.pipelines
::: fastvideo.pipelines
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
show_submodules: true
heading_level: 3
### Submodules
#### fastvideo.pipelines.composed_pipeline_base
::: fastvideo.pipelines.composed_pipeline_base
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
heading_level: 4
#### fastvideo.pipelines.lora_pipeline
::: fastvideo.pipelines.lora_pipeline
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
heading_level: 4
#### fastvideo.pipelines.pipeline_batch_info
::: fastvideo.pipelines.pipeline_batch_info
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
heading_level: 4
#### fastvideo.pipelines.pipeline_registry
::: fastvideo.pipelines.pipeline_registry
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
heading_level: 4
#### fastvideo.pipelines.stages
::: fastvideo.pipelines.stages
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
heading_level: 4
## fastvideo.training
::: fastvideo.training
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
show_submodules: true
heading_level: 3
## fastvideo.workflow
::: fastvideo.workflow
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
show_submodules: true
heading_level: 3
## fastvideo.dataset
::: fastvideo.dataset
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
show_submodules: true
heading_level: 3
## fastvideo.layers
::: fastvideo.layers
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
show_submodules: true
heading_level: 3
## fastvideo.platforms
::: fastvideo.platforms
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
show_submodules: true
heading_level: 3
## fastvideo.utils
::: fastvideo.utils
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
heading_level: 3
## fastvideo.worker
::: fastvideo.worker
options:
show_source: true
show_root_heading: true
show_root_toc_entry: true
show_submodules: true
heading_level: 3
+27
View File
@@ -0,0 +1,27 @@
# API Summary
This page provides a quick overview of the main FastVideo API components.
## Video Generator
::: fastvideo.VideoGenerator
options:
show_root_heading: false
show_source: false
heading_level: 3
## Initialization Configuration
::: fastvideo.PipelineConfig
options:
show_root_heading: false
show_source: false
heading_level: 3
## Sampling Configuration
::: fastvideo.SamplingParam
options:
show_root_heading: false
show_source: false
heading_level: 3
+43
View File
@@ -0,0 +1,43 @@
.vertical-table-header th.head:not(.stub) {
writing-mode: sideways-lr;
white-space: nowrap;
max-width: 0;
}
/* Keep header cell paragraph content tight (avoid CSS nesting for compatibility) */
.vertical-table-header th.head:not(.stub) p {
margin: 0;
}
/* Image sizing classes */
.image-small {
max-width: 200px;
height: auto;
}
.image-medium {
max-width: 400px;
height: auto;
}
.image-large {
max-width: 600px;
height: auto;
}
.image-full {
max-width: 100%;
height: auto;
}
/* Responsive images */
img {
max-width: 100%;
height: auto;
}
/* Center images */
.image-center {
display: block;
margin: 0 auto;
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 98 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 122 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 194 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 378 KiB

Before

Width:  |  Height:  |  Size: 303 KiB

After

Width:  |  Height:  |  Size: 303 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 575 KiB

Before

Width:  |  Height:  |  Size: 18 KiB

After

Width:  |  Height:  |  Size: 18 KiB

Before

Width:  |  Height:  |  Size: 27 KiB

After

Width:  |  Height:  |  Size: 27 KiB

Before

Width:  |  Height:  |  Size: 40 KiB

After

Width:  |  Height:  |  Size: 40 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

+186
View File
@@ -0,0 +1,186 @@
# Adding a New Attention Backend
FastVideo allows integrating new attention mechanisms easily. This guide walks you through adding a new backend (e.g., `MyNewAttn`).
## 1. Implement the Backend (Python)
Create a new file in `fastvideo/attention/backends/` (e.g., `mynew_attn.py`).
Your implementation should inherit from `AttentionBackend` defined in `abstract.py`.
```python
# fastvideo/attention/backends/mynew_attn.py
import torch
from .abstract import AttentionBackend
# Import the context manager to access metadata (optional)
from fastvideo.forward_context import get_forward_context
# Import compiled kernel if applicable (see Section 2)
try:
# Import from the top-level package
from fastvideo_kernel import my_compiled_attn_func
except ImportError:
my_compiled_attn_func = None
class MyNewAttnBackend(AttentionBackend):
def process_inputs(self, q, k, v, **kwargs):
# Pre-process inputs if necessary
return q, k, v
def forward(self, q, k, v, **kwargs):
# Optional: Access extra metadata passed via ForwardContext
# Only needed if your backend requires global state (e.g. window_size)
try:
context = get_forward_context()
metadata = context.attn_metadata
# Example: window_size = metadata.window_size
except (AssertionError, AttributeError):
# Handle case where context is not set (e.g. standard inference)
pass
if my_compiled_attn_func is not None:
return my_compiled_attn_func(q, k, v)
else:
# Fallback implementation (e.g., Triton or pure PyTorch)
return self.fallback_impl(q, k, v)
```
## 2. Passing Extra Information via ForwardContext (Optional)
FastVideo uses a `ForwardContext` to pass global metadata (like current timestep, batch info, or custom attention configurations) to attention backends without changing the `forward` signature of every layer. **This is optional and only required if your backend needs dynamic per-step information.**
To use this:
1. **Set Context**: In your pipeline or generation loop, use the `set_forward_context` context manager.
2. **Access Context**: Inside your attention backend, use `get_forward_context()`.
See [`docs/attention/sta/index.md`](../sta/index.md) (Sliding Tile Attention) for an example of how complex configuration (window sizes) is passed this way.
## 3. Adding Compiled Kernels (C++/CUDA)
If your backend requires custom CUDA kernels, you need to add them to the `fastvideo-kernel` package.
### A. Add Source Files
Place your kernel implementation files in `fastvideo-kernel/csrc/attention/`.
* `mynew_attn.cu` (CUDA implementation)
* `mynew_attn.h` (Optional headers)
### B. Register in Extension
Update `fastvideo-kernel/csrc/common_extension.cpp` to expose your function to Python.
```cpp
// 1. Declare external function
#ifdef COMPILE_MYNEW_ATTN
extern torch::Tensor mynew_attn_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v);
#endif
// 2. Register in module
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
// ... other kernels ...
#ifdef COMPILE_MYNEW_ATTN
m.def("mynew_attn_fwd", torch::wrap_pybind_function(mynew_attn_forward), "My New Attention Forward");
#endif
}
```
### C. Update CMakeLists.txt
Update `fastvideo-kernel/CMakeLists.txt` to compile your new files.
**Case 1: General CUDA Kernel (Runs on all GPUs)**
Add your source file directly to `EXTENSION_SOURCES` and define the compilation flag.
```cmake
# Add to EXTENSION_SOURCES
list(APPEND EXTENSION_SOURCES csrc/attention/mynew_attn.cu)
# Add compilation definition for common_extension.cpp
list(APPEND COMPILE_DEFS COMPILE_MYNEW_ATTN)
```
**Case 2: ThunderKittens Kernel (Hopper H100 Only)**
If your kernel uses ThunderKittens (TK), it requires specific architecture flags (`sm_90a`). Add it inside the `ENABLE_TK_KERNELS` block.
```cmake
if(ENABLE_TK_KERNELS)
# Add source only if TK is enabled
list(APPEND EXTENSION_SOURCES csrc/attention/mynew_attn_tk.cu)
# Add definition to guard registration
list(APPEND COMPILE_DEFS TK_COMPILE_MYNEW_ATTN)
endif()
```
### D. Expose in Python Ops
Update `fastvideo-kernel/python/fastvideo_kernel/ops.py` to make the function importable and handle fallbacks gracefully.
```python
# fastvideo-kernel/python/fastvideo_kernel/ops.py
# Try to load C++ extension symbols
try:
from fastvideo_kernel._C import fastvideo_kernel_ops
mynew_attn_fwd = getattr(fastvideo_kernel_ops, "mynew_attn_fwd", None)
except ImportError:
mynew_attn_fwd = None
def my_compiled_attn_func(q, k, v):
# Runtime check: use C++ kernel if available, else fallback
if mynew_attn_fwd is not None:
return mynew_attn_fwd(q, k, v)
else:
# Call Triton/Python fallback
return mynew_attn_triton(q, k, v)
```
### E. Expose in Package Init
Update `fastvideo-kernel/python/fastvideo_kernel/__init__.py` to export the function.
```python
from fastvideo_kernel.ops import (
my_compiled_attn_func,
# ...
)
__all__ = [
"my_compiled_attn_func",
# ...
]
```
## 4. Register the Backend
Update `fastvideo/attention/backends/__init__.py` to export your new class.
```python
from .mynew_attn import MyNewAttnBackend
```
## 5. Platform Integration
If your backend requires specific platform checks (e.g., checking for H100 support), handle that in `fastvideo/platforms/cuda.py` or within your backend's `__init__`.
## 6. Add Documentation
Create a new documentation page for your backend to explain its usage, installation (if custom kernels are needed), and features.
1. **Create Directory**: `docs/attention/mynew_attn/`
2. **Create Index**: `docs/attention/mynew_attn/index.md`
3. **Update Navigation**: Add an entry to `mkdocs.yml` under the "Attention" tab.
## Checklist
* [ ] Created `fastvideo/attention/backends/mynew_attn.py`.
* [ ] (Optional) Added CUDA kernels in `fastvideo-kernel/csrc/attention/`.
* [ ] (Optional) Updated `common_extension.cpp` and `CMakeLists.txt`.
* [ ] (Optional) Exposed kernel in `fastvideo-kernel/python/fastvideo_kernel/ops.py`.
* [ ] (Optional) Exported kernel in `fastvideo-kernel/python/fastvideo_kernel/__init__.py`.
* [ ] Implemented `forward` method respecting the standard signature.
* [ ] Added unit tests in `tests/`.
* [ ] Added documentation in `docs/attention/` and updated `mkdocs.yml`.
+53
View File
@@ -0,0 +1,53 @@
# FastVideo Attention Kernels
FastVideo provides highly optimized custom attention kernels to accelerate video generation.
## Supported Kernels
* **[Video Sparse Attention (VSA)](vsa/index.md)**: Sparse attention mechanism selecting top-k blocks.
* **[Sliding Tile Attention (STA)](sta/index.md)**: Optimized attention for window-based video generation.
## General Build Instructions
These instructions apply to building the `fastvideo-kernel` package from source, which includes both STA and VSA kernels.
### Prerequisites
* **PyTorch**: 2.5.0+
* **CUDA**: 12.4+ (12.8 recommended for best performance)
* **C++ Compiler**: GCC 11+ (C++20 support required for ThunderKittens)
Install system dependencies:
```bash
sudo apt update
sudo apt install -y gcc-11 g++-11 clang-11 ninja-build
# Set gcc-11 as default
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
```
Set up your CUDA environment variables (adjust version as needed):
```bash
export CUDA_HOME=/usr/local/cuda-12.8
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
```
### Compile and Install
Clone the repository and build the kernel:
```bash
# Clone recursively to get ThunderKittens submodule
git clone https://github.com/hao-ai-lab/FastVideo.git
cd FastVideo/fastvideo-kernel
# Build and install
./build.sh
```
The build script automatically detects your GPU architecture:
* **H100 (sm_90a)**: Compiles optimized C++ ThunderKittens kernels.
* **Other (A100, etc.)**: Skips C++ compilation; installs Python package with Triton kernels.
+36
View File
@@ -0,0 +1,36 @@
# Sliding Tile Attention (STA)
Optimized attention for window-based video generation (e.g., HunyuanVideo).
## Installation
STA is included in the `fastvideo-kernel` package. See the [main Attention page](../index.md) for build instructions.
## Usage
```python
from fastvideo_kernel import sliding_tile_attention
# q, k, v: [batch_size, num_heads, seq_length, head_dim]
# window_size: List of (t, h, w) tiles. Tile size is (6, 8, 8).
# text_length: Number of text tokens (0-256)
out = sliding_tile_attention(
q, k, v,
window_size=[(3, 3, 3)], # Example window
text_length=256
)
```
## Citation
If you use Sliding Tile Attention in your research, please cite:
```bibtex
@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}
}
```
+38
View File
@@ -0,0 +1,38 @@
# Video Sparse Attention (VSA)
Sparse attention mechanism selecting top-k blocks.
## Installation
VSA is included in the `fastvideo-kernel` package. See the [main Attention page](../index.md) for build instructions.
## Usage
```python
from fastvideo_kernel import video_sparse_attn
# q, k, v: [batch_size, num_heads, seq_len, head_dim]
# variable_block_sizes: Number of valid tokens per block
# q_variable_block_sizes: Number of valid tokens per q block (can differ from KV for q/k of different lengths)
# topk: Number of blocks to attend
output = video_sparse_attn(
q, k, v,
block_sizes,
block_sizes,
topk=32
)
```
## Citation
If you use Video Sparse Attention in your research, please cite:
```bibtex
@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}
}
```
@@ -1,14 +1,14 @@
(docker)=
# 🐳 Using the FastVideo Docker Image
If you prefer a containerized development environment or want to avoid managing dependencies manually, you can use our prebuilt Docker image:
**Image:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
**Images:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
## Starting the container
```bash
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest
```
This will:
@@ -3,11 +3,3 @@
# 🧰 Developer Environment
Accelerate your FastVideo development workflow by leveraging Docker images and cloud GPUs for efficient experimentation and reproducible environments.
:::{toctree}
:caption: Contents
:maxdepth: 1
docker
runpod
:::
@@ -1,4 +1,3 @@
(runpod)=
# 📦 Developing FastVideo on RunPod
@@ -6,11 +5,11 @@ You can easily use the FastVideo Docker image as a custom container on [RunPod](
## Creating a new pod
Choose a GPU that supports CUDA 12.4
Choose a GPU that supports CUDA 12.8
Pick 1 or 2 L40S GPU(s)
![RunPod CUDA selection](../../_static/images/runpod_cuda.png)
![RunPod CUDA selection](../../assets/images/runpod_cuda.png)
When creating your pod template, use this image:
@@ -24,11 +23,11 @@ Paste Container Start Command to support SSH ([RunPod Docs](https://docs.runpod.
bash -c "apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-server -y;mkdir -p ~/.ssh;cd $_;chmod 700 ~/.ssh;echo \"$PUBLIC_KEY\" >> authorized_keys;chmod 700 authorized_keys;service ssh start;sleep infinity"
```
![RunPod template configuration](../../_static/images/runpod_template.png)
![RunPod template configuration](../../assets/images/runpod_template.png)
After deploying, the pod will take a few minutes to pull the image and start the SSH service.
![RunPod ssh](../../_static/images/runpod_ssh.png)
![RunPod ssh](../../assets/images/runpod_ssh.png)
## Working with the pod
@@ -1,4 +1,3 @@
(developer-overview)=
# 🛠️ Contributing to FastVideo
@@ -7,7 +6,7 @@ Thank you for your interest in contributing to FastVideo. We want to make the pr
Our community is open to everyone and welcomes any contributions no matter how large or small.
# Developer Environment:
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only support Linux and CUDA GPUs, but we hope to support other platforms in the future.
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only supports Linux and CUDA GPUs, but we hope to support other platforms in the future.
We recommend using a fresh Python 3.10 Conda environment to develop FastVideo:
@@ -22,10 +21,20 @@ source ~/.bashrc
Create and activate a Conda environment for FastVideo:
```
conda create -n fastvideo python=3.10 -y
conda create -n fastvideo python=3.12 -y
conda activate fastvideo
```
Install `uv` (optional, but recommended):
From instructions on [uv](https://astral.sh/uv/):
```
curl -LsSf https://astral.sh/uv/install.sh | sh
# or
wget -qO- https://astral.sh/uv/install.sh | sh
```
Clone the FastVideo repository and go to the FastVideo directory:
```
@@ -36,10 +45,10 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
Now you can install FastVideo and setup git hooks for running linting. By using `pre-commit`, the linters will run and have to pass before you'll be able to make a commit.
```bash
pip install -e .[dev]
uv pip install -e .[dev]
# Can also install flash-attn (optional)
pip install flash-attn==2.7.4.post1 --no-build-isolation
uv pip install flash-attn --no-build-isolation
# Linting, formatting and static type checking
pre-commit install --hook-type pre-commit --hook-type commit-msg
@@ -50,3 +59,18 @@ pre-commit run --all-files
# Unit tests
pytest tests/
```
If you are on a Hopper GPU, you should also install [FA3](https://github.com/Dao-AILab/flash-attention) for much better performance:
```
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention/hopper
# make sure you have ninja installed
uv pip install ninja
python setup.py install
```
## Testing
Please refer to the [Testing Guide](testing.md) for more information on how to add and run tests in FastVideo.
+53
View File
@@ -0,0 +1,53 @@
# Profiling FastVideo
!!! warning
Profiling is only intended for FastVideo developers and maintainers to understand the proportion of time spent in different parts of the codebase. **FastVideo end-users should never turn on profiling** as it will significantly slow down inference.
## Profiling with PyTorch
FastVideo exposes a process-wide torch profiler that you can enable via environment variables. Set `FASTVIDEO_TORCH_PROFILER_DIR` to an absolute directory path to start collecting traces, and specify the regions you want recorded with `FASTVIDEO_TORCH_PROFILE_REGIONS`:
```bash
FASTVIDEO_TORCH_PROFILER_DIR=/mnt/traces/fastvideo \
FASTVIDEO_TORCH_PROFILE_REGIONS="profiler_region_model_loading,profiler_region_training_step"
```
All profiled regions must be registered in `fastvideo.profiler`; the current list includes:
- `profiler_region_model_loading` — pipeline/module loading
- `profiler_region_inference_pre_denoising`
- `profiler_region_inference_denoising`
- `profiler_region_inference_post_denoising`
- `profiler_region_training_checkpoint_saving`
- `profiler_region_training_dit`
- `profiler_region_training_validation`
- `profiler_region_training_epoch`
- `profiler_region_training_step`
- `profiler_region_training_backward`
- `profiler_region_training_optimizer`
- `profiler_region_distillation_teacher_forward`
- `profiler_region_distillation_student_forward`
- `profiler_region_distillation_loss`
- `profiler_region_distillation_update`
While profiling is enabled, FastVideo records additional annotations:
- `fastvideo.region::<name>` spans are emitted when entering a region.
- `fastvideo.profiler.enable_collection` / `fastvideo.profiler.disable_collection` events mark when torch profiler collection is toggled on or off.
Only one profiler instance is created per process; subsequent pipelines reuse the same controller. If you set `FASTVIDEO_TORCH_PROFILE_REGIONS` incorrectly (e.g. misspelled name), FastVideo logs a warning and ignores that entry.
Additional knobs:
- `FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES`
- `FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY`
- `FASTVIDEO_TORCH_PROFILER_WITH_STACK`
- `FASTVIDEO_TORCH_PROFILER_WITH_FLOPS`
Traces can be visualized using <https://ui.perfetto.dev/>.
### Best Practices
- Keep the profiled step count small; traces can be large and slow down job shutdown while the profiler flushes data.
- After profiling, clean up trace directories to avoid filling disk storage.
- When adding new regions, register them in `fastvideo.profiler` and wrap the corresponding code block with `with self.profiler_controller.region("your_region"):` or the `@profile_region` decorator.
+131
View File
@@ -0,0 +1,131 @@
# Testing in FastVideo
This guide explains how to add and run tests in FastVideo. The testing suite is divided into several categories to ensure correctness across components, training workflows, and inference quality.
## Test Types
* **Unit Tests**: Located in `fastvideo/tests/dataset`, `fastvideo/tests/entrypoints`, and `fastvideo/tests/workflow`. These test individual functions and classes.
* **Component Tests**: Located in `fastvideo/tests/encoders`, `fastvideo/tests/transformers`, and `fastvideo/tests/vaes`. These verify the loading and basic functionality of model components.
* **SSIM Tests**: Located in `fastvideo/tests/ssim`. These are regression tests that compare generated videos against reference videos using the Structural Similarity Index Measure (SSIM) to detect quality degradation.
* **Training Tests**: Located in `fastvideo/tests/training`. These validate training loops, loss calculations, and specific training techniques like LoRA, Distillation, and VSA.
* **Inference Tests**: Located in `fastvideo/tests/inference`. These test specialized inference pipelines and optimizations (e.g., STA, V-MoBA).
For now, we will focus on **SSIM Tests**.
## SSIM Tests
SSIM tests are located in `fastvideo/tests/ssim`. These tests generate videos using specific models and parameters, and compare them against reference videos to ensure that changes in the codebase do not degrade generation quality or alter the output unexpectedly.
!!! note
If you are adding an SSIM test, this serves as a safeguard. Any future code changes that break or cause errors with the specific arguments and configurations you defined will trigger a failure. Therefore, it is important to include multiple settings and arguments that cover the core features of your new pipeline to ensure robust regression testing.
### Directory Structure
```
fastvideo/tests/ssim/
├── <GPU>_reference_videos/ # Reference videos organized by GPU type (e.g., L40S_reference_videos)
│ ├── <Model_Name>/
│ │ ├── <Backend>/ # e.g., FLASH_ATTN, TORCH_SDPA
│ │ │ └── <Video_File>
├── test_causal_similarity.py
├── test_inference_similarity.py
├── update_reference_videos.sh
└── ...
```
### Adding a New SSIM Test
To add a new SSIM test, follow these steps:
1. **Create or Update a Test File**: You can add a new test function to an existing file (like `test_inference_similarity.py`) or create a new one if testing a distinct category of models.
2. **Define Model Parameters**: Define the configuration for the model you want to test. This includes model path, dimensions, inference steps, and other generation parameters. **Note:** Consider using lower `num_inference_steps` or reduced resolution (e.g., 480p instead of 720p) to keep test execution time reasonable, provided it doesn't compromise the test's ability to detect regression.
```python
MY_MODEL_PARAMS = {
"num_gpus": 1,
"model_path": "organization/model-name",
"height": 480,
"width": 832,
"num_frames": 45,
"num_inference_steps": 20,
# ... other parameters
}
```
3. **Implement the Test Function**:
* Use `pytest.mark.parametrize` to run the test with different prompts, backends, and models.
* Set the attention backend environment variable.
* Initialize the `VideoGenerator`.
* Generate the video.
* Compare the generated video with the reference video using `compute_video_ssim_torchvision`.
Example structure:
```python
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
def test_my_model_similarity(prompt, ATTENTION_BACKEND):
# Setup output directories
# ...
# Initialize Generator
generator = VideoGenerator.from_pretrained(...)
generator.generate_video(prompt, ...)
# Compare with Reference
ssim_values = compute_video_ssim_torchvision(
reference_path, generated_path, use_ms_ssim=True
)
assert ssim_values[0] >= 0.98 # Threshold
```
4. **Reference Videos**:
* When running the test for the first time (or when updating the reference), the test will fail because the reference video is missing. The generated video will be saved in `fastvideo/tests/ssim/generated_videos`.
* Inspect the generated video to ensure it meets quality expectations.
* Move the generated video to the appropriate reference folder: `fastvideo/tests/ssim/<GPU>_reference_videos/<Model>/<Backend>/`.
* You can use the helper script `update_reference_videos.sh` to automate copying videos from `generated_videos` to `L40S_reference_videos`. Note: Check the script to ensure paths match your environment (it defaults to `L40S_reference_videos`).
### Running Tests Locally
To run the SSIM tests locally:
```bash
pytest fastvideo/tests/ssim/ -vs
```
Ensure you have the necessary GPUs available as defined in your test parameters.
## Modal Workflow
FastVideo uses [Modal](https://modal.com/) for running tests in a CI environment. The workflow scripts are located in `fastvideo/tests/modal/`.
### `pr_test.py`
The main entry point for CI tests is `fastvideo/tests/modal/pr_test.py`. This script defines Modal functions that execute the pytest suites on specific hardware (e.g., L40S, H100).
### Updating Modal Configuration
If you add a new test that requires:
* **Different GPU Hardware**: You may need to change the `@app.function(gpu=...)` decorator.
* **Longer Execution Time**: Increase the `timeout` parameter.
* **New Environment Variables/Secrets**: Add them to `secrets=[...]` or the image environment. For example, if your model is gated on Hugging Face, ensure `HF_API_KEY` is passed.
For SSIM tests, the `run_ssim_tests` function in `pr_test.py` currently runs:
```python
@app.function(gpu="L40S:2", image=image, timeout=2700, secrets=[modal.Secret.from_dict({"HF_API_KEY": os.environ.get("HF_API_KEY", "")})])
def run_ssim_tests():
run_test("hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/ssim -vs")
```
If your new test file is inside `fastvideo/tests/ssim`, it will automatically be picked up by this command. However, ensure that the `gpu="L40S:2"` configuration is sufficient for your model. If your model requires more GPUs (e.g., 4 or 8), you might need to create a separate Modal function or update the existing one.
### Workflow Scripts
The shell script that triggers these tests in the CI pipeline is located at `.buildkite/scripts/pr_test.sh`. If you add a new test category (e.g., a new folder outside of `ssim`), you will need to:
1. Add a new function in `fastvideo/tests/modal/pr_test.py`.
2. Add a new case in `.buildkite/scripts/pr_test.sh` to handle the new test type.
!!! note
If you are a maintainer, you'll need to finally manually update the workflow script in Buildkite. Otherwise, a maintainer will help you update.
@@ -1,46 +1,48 @@
# 🔍 FastVideo Overview
This document outlines FastVideo's architecture for developers interested in framework internals or contributions. It serves as an onboarding guide for new contributors by providing an overview of the most important directories and files within the `fastvideo/v1/` codebase.
This document outlines FastVideo's architecture for developers interested in framework internals or contributions. It serves as an onboarding guide for new contributors by providing an overview of the most important directories and files within the `fastvideo/` codebase.
## Table of Contents - V1 Directory Structure and Files
## Table of Contents - Directory Structure and Files
- [`fastvideo/v1/pipelines/`](#design-pipeline-system) - Core diffusion pipeline components
- [`fastvideo/v1/models/`](#design-model-components) - Model implementations
- [`dits/`](#design-transformer-models) - Transformer-based diffusion models
- [`vaes/`](#design-vae-variational-auto-encoder) - Variational autoencoders
- [`encoders/`](#design-text-and-image-encoders) - Text and image encoders
- [`schedulers/`](#design-schedulers) - Diffusion schedulers
- [`fastvideo/v1/attention/`](#design-optimized-attention) - Optimized attention implementations
- [`fastvideo/v1/distributed/`](#design-distributed-processing) - Distributed computing utilities
- [`fastvideo/v1/layers/`](#design-tensor-parallelism) - Custom neural network layers
- [`fastvideo/v1/platforms/`](#design-platforms) - Hardware platform abstractions
- [`fastvideo/v1/worker/`](#design-executor-and-worker-abstractions) - Multi-GPU process management
- [`fastvideo/v1/fastvideo_args.py`](#design-fastvideo-args) - Argument handling
- [`fastvideo/v1/forward_context.py`](#design-forwardcontext) - Forward pass context management
- `fastvideo/v1/utils.py` - Utility functions
- [`fastvideo/v1/logger.py`](#design-logger) - Logging infrastructure
- [`fastvideo/pipelines/`](#pipeline-system) - Core diffusion pipeline components
- [`fastvideo/models/`](#model-components) - Model implementations
- [`dits/`](#transformer-models) - Transformer-based diffusion models
- [`vaes/`](#vae-variational-auto-encoder) - Variational autoencoders
- [`encoders/`](#text-and-image-encoders) - Text and image encoders
- [`schedulers/`](#schedulers) - Diffusion schedulers
- [`fastvideo/attention/`](#optimized-attention) - Optimized attention implementations
- [`fastvideo/distributed/`](#distributed-processing) - Distributed computing utilities
- [`fastvideo/layers/`](#tensor-parallelism) - Custom neural network layers
- [`fastvideo/platforms/`](#platforms) - Hardware platform abstractions
- [`fastvideo/worker/`](#executor-and-worker-system) - Multi-GPU process management
- [`fastvideo/fastvideo_args.py`](#fastvideoargs) - Argument handling
- [`fastvideo/forward_context.py`](#forward-context-management) - Forward pass context management
- `fastvideo/utils.py` - Utility functions
- [`fastvideo/logger.py`](#logger) - Logging infrastructure
## Core Architecture
FastVideo separates model components from execution logic with these principles:
- **Component Isolation**: Models (encoders, VAEs, transformers) are isolated from execution (pipelines, stages, distributed processing)
- **Modular Design**: Components can be independently replaced
- **Distributed Execution**: Supports various parallelism strategies (Tensor, Sequence)
- **Custom Attention Backends**: Components can support and use different Attention implementations
- **Pipeline Abstraction**: Consistent interface across diffusion models
(design-fastvideo-args)=
## FastVideoArgs
The `FastVideoArgs` class in `fastvideo/v1/fastvideo_args.py` serves as the central configuration system for FastVideo. It contains all parameters needed to control model loading, inference configuration, performance optimization settings, and more.
The `FastVideoArgs` class in `fastvideo/fastvideo_args.py` serves as the central configuration system for FastVideo. It contains all parameters needed to control model loading, inference configuration, performance optimization settings, and more.
Key features include:
- **Command-line Interface**: Automatic conversion between CLI arguments and dataclass fields
- **Configuration Groups**: Organized by functional areas (model loading, video params, optimization settings)
- **Context Management**: Global access to current settings via `get_current_fastvideo_args()`
- **Parameter Validation**: Ensures valid combinations of settings
Common configuration areas:
- **Model paths and loading options**: `model_path`, `trust_remote_code`, `revision`
- **Distributed execution settings**: `num_gpus`, `tp_size`, `sp_size`
- **Video generation parameters**: `height`, `width`, `num_frames`, `num_inference_steps`
@@ -61,7 +63,6 @@ with set_current_fastvideo_args(fastvideo_args):
result = generate_video()
```
(design-pipeline-system)=
## Pipeline System
### `ComposedPipelineBase`
@@ -92,7 +93,9 @@ class MyCustomPipeline(ComposedPipelineBase):
```
### Pipeline Stages
Each stage handles a specific diffusion process component:
- **Input Validation**: Parameter verification
- **Text Encoding**: CLIP, LLaMA, or T5-based encoding
- **Image Encoding**: Image input processing
@@ -108,10 +111,11 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward
return batch
```
(design-forwardbatch)=
![Pipeline execution and data flow](../assets/images/pipeline.png)
### ForwardBatch
Defined in `fastvideo/v1/pipelines/pipeline_batch_info.py`, `ForwardBatch` encapsulates the data payload passed between pipeline stages. It typically holds:
Defined in `fastvideo/pipelines/pipeline_batch_info.py`, `ForwardBatch` encapsulates the data payload passed between pipeline stages. It typically holds:
- **Input Data**: Prompts, images, generation parameters
- **Intermediate State**: Embeddings, latents, timesteps, accumulated during stage execution
@@ -120,22 +124,21 @@ Defined in `fastvideo/v1/pipelines/pipeline_batch_info.py`, `ForwardBatch` encap
This structure facilitates clear state transitions between stages.
(design-model-components)=
## Model Components
The `fastvideo/v1/models/` directory contains implementations of the core neural network models used in video diffusion:
The `fastvideo/models/` directory contains implementations of the core neural network models used in video diffusion:
(design-transformer-models)=
### Transformer Models
Transformer networks perform the actual denoising during diffusion:
- **Location**: `fastvideo/v1/models/dits/`
- **Location**: `fastvideo/models/dits/`
- **Examples**:
- `WanTransformer3DModel`
- `HunyuanVideoTransformer3DModel`
Features include:
- Text/image conditioning
- Standardized interface for model-specific optimizations
@@ -152,12 +155,11 @@ def forward(
return noise_pred # Predicted noise residual
```
(design-vae-variational-auto-encoder)=
### VAE (Variational Auto-Encoder)
VAEs handle conversion between pixel space and latent space:
- **Location**: `fastvideo/v1/models/vaes/`
- **Location**: `fastvideo/models/vaes/`
- **Examples**:
- `AutoencoderKLWan`
- `AutoencoderKLHunyuanVideo`
@@ -165,17 +167,17 @@ VAEs handle conversion between pixel space and latent space:
These models compress image/video data to a more efficient latent representation (typically 4x-8x smaller in each dimension).
FastVideo's VAE implementations include:
- Efficient video batch processing
- Memory optimization
- Optional tiling for large frames
- Distributed weight support
(design-text-and-image-encoders)=
### Text and Image Encoders
Encoders process conditioning inputs into embeddings:
- **Location**: `fastvideo/v1/models/encoders/`
- **Location**: `fastvideo/models/encoders/`
- **Text Encoders**:
- `CLIPTextModel`
- `LlamaModel`
@@ -184,21 +186,22 @@ Encoders process conditioning inputs into embeddings:
- `CLIPVisionModel`
FastVideo implements optimizations such as:
- Vocab parallelism for distributed processing
- Caching for common prompts
- Precision-tuned computation
(design-schedulers)=
### Schedulers
Schedulers manage the diffusion sampling process:
- **Location**: `fastvideo/v1/models/schedulers/`
- **Location**: `fastvideo/models/schedulers/`
- **Examples**:
- `UniPCMultistepScheduler`
- `FlowMatchEulerDiscreteScheduler`
These components control:
- Diffusion timestep sequences
- Noise prediction to latent update conversions
- Quality/speed trade-offs
@@ -216,13 +219,18 @@ def step(
return prev_sample
```
(design-optimized-attention)=
This diagram shows how models are discovered, validated, and loaded across entrypoints, executors, pipelines, and model loaders.
![Model loading flow](../assets/images/load_models.png)
## Optimized Attention
The `fastvideo/v1/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
The `fastvideo/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
### Attention Backends
Multiple implementations with automatic selection:
- **FLASH_ATTN**: Optimized for supporting hardware
- **TORCH_SDPA**: Built-in PyTorch scaled dot-product attention
- **SLIDING_TILE_ATTN**: For very long sequences
@@ -240,17 +248,19 @@ self.attn = LocalAttention(
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
```
![Attention backend selector design](../assets/images/attention_backend.png)
### Attention Patterns
Supports various patterns with memory optimization techniques:
- **Cross/Self/Temporal/Global-Local Attention**
- Chunking, progressive computation, optimized masking
(design-distributed-processing)=
## Distributed Processing
The `fastvideo/v1/distributed/` directory contains implementations for distributed model execution:
The `fastvideo/distributed/` directory contains implementations for distributed model execution:
(design-tensor-parallelism)=
### Tensor Parallelism
Tensor parallelism splits model weights across devices:
@@ -260,7 +270,7 @@ Tensor parallelism splits model weights across devices:
```python
# Tensor-parallel layers in a transformer block
from fastvideo.v1.layers.linear import ColumnParallelLinear, RowParallelLinear
from fastvideo.layers.linear import ColumnParallelLinear, RowParallelLinear
# Split along output dimension
self.qkv_proj = ColumnParallelLinear(
@@ -288,7 +298,7 @@ Sequence parallelism splits sequences across devices:
```python
# Distributed attention for long sequences
from fastvideo.v1.attention import DistributedAttention
from fastvideo.attention import DistributedAttention
self.attn = DistributedAttention(
num_heads=num_heads,
@@ -299,6 +309,7 @@ self.attn = DistributedAttention(
```
### Communication Primitives
Efficient distributed operations via AllGather, AllReduce, and synchronization mechanisms.
Efficient communication primitives minimize distributed overhead:
@@ -307,17 +318,17 @@ Efficient communication primitives minimize distributed overhead:
- **Tensor-Parallel AllReduce**: Combines partial results
- **Distributed Synchronization**: Coordinates execution
(design-forwardcontext)=
## Forward Context Management
### ForwardContext
Defined in `fastvideo/v1/forward_context.py`, `ForwardContext` manages execution-specific state *within* a forward pass, particularly for low-level optimizations. It is accessed via `get_forward_context()`.
Defined in `fastvideo/forward_context.py`, `ForwardContext` manages execution-specific state *within* a forward pass, particularly for low-level optimizations. It is accessed via `get_forward_context()`.
- **Attention Metadata**: Configuration for optimized attention kernels (`attn_metadata`)
- **Profiling Data**: Potential hooks for performance metrics collection
This context-based approach enables:
- Dynamic optimization based on execution state (e.g., attention backend selection)
- Step-specific customizations within model components
@@ -330,10 +341,9 @@ with set_forward_context(current_timestep, attn_metadata, fastvideo_args):
output = model(inputs)
```
(design-executor-and-worker-abstractions)=
## Executor and Worker System
The `fastvideo/v1/worker/` directory contains the distributed execution framework:
The `fastvideo/worker/` directory contains the distributed execution framework:
### Executor Abstraction
@@ -344,12 +354,14 @@ FastVideo implements a flexible execution model for distributed processing:
- **GPU Workers**: Handle actual model execution on individual GPUs
The MultiProcExecutor implementation:
1. Spawns worker processes for each GPU
2. Establishes communication channels via pipes
3. Coordinates distributed operations across workers
4. Handles graceful startup and shutdown of the process group
Each GPU worker:
1. Initializes the distributed environment
2. Builds the pipeline for the specified model
3. Executes requested operations on its assigned GPU
@@ -357,19 +369,20 @@ Each GPU worker:
This design allows FastVideo to efficiently utilize multiple GPUs while providing a simple, unified interface for model execution.
(design-platforms)=
## Platforms
The `fastvideo/v1/platforms/` directory provides hardware platform abstractions that enable FastVideo to run efficiently on different hardware configurations:
The `fastvideo/platforms/` directory provides hardware platform abstractions that enable FastVideo to run efficiently on different hardware configurations:
### Platform Abstraction
FastVideo's platform abstraction layer enables:
- **Hardware Detection**: Automatic detection of available hardware
- **Backend Selection**: Appropriate selection of compute kernels
- **Memory Management**: Efficient utilization of hardware-specific memory features
The primary components include:
- **Platform Interface**: Defines the common API for all platform implementations
- **CUDA Platform**: Optimized implementation for NVIDIA GPUs
- **Backend Enum**: Used throughout the codebase for feature selection
@@ -377,7 +390,7 @@ The primary components include:
Usage example:
```python
from fastvideo.v1.platforms import current_platform, _Backend
from fastvideo.platforms import current_platform, _Backend
# Check hardware capabilities
if current_platform.supports_backend(_Backend.FLASH_ATTN):
@@ -388,8 +401,8 @@ else:
The platform system is designed to be extensible for future hardware targets.
(design-logger)=
## Logger
See [PR](https://github.com/hao-ai-lab/FastVideo/pull/356)
*TODO*: (help wanted) Add an environment variable that disables process-aware logging.
@@ -398,12 +411,13 @@ See [PR](https://github.com/hao-ai-lab/FastVideo/pull/356)
If you're a new contributor, here are some common areas to explore:
1. **Adding a new model**: Implement new model types in the appropriate subdirectory of `fastvideo/v1/models/`
1. **Adding a new model**: Implement new model types in the appropriate subdirectory of `fastvideo/models/`
2. **Optimizing performance**: Look at attention implementations or memory management
3. **Adding a new pipeline**: Create a new pipeline subclass in `fastvideo/v1/pipelines/`
3. **Adding a new pipeline**: Create a new pipeline subclass in `fastvideo/pipelines/`
4. **Hardware support**: Extend the `platforms` module for new hardware targets
When adding code, follow these practices:
- Use type hints for better code readability
- Add appropriate docstrings
- Maintain the separation between model components and execution logic
+42
View File
@@ -0,0 +1,42 @@
# 🧱 Data Preprocess for Distillation
For distillation, we use the same data preprocessing pipeline as training. Please refer to the [Training Data Preprocess](../training/data_preprocess.md) for general preprocessing steps.
## Distillation-Specific Datasets
### FastVideo 480P Synthetic Wan Dataset
For Wan2.1 T2V distillation, we use the **FastVideo 480P Synthetic Wan dataset** ([FastVideo/Wan-Syn_77x448x832_600k](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x448x832_600k)) which contains 600k synthetic latents.
```bash
# Download the preprocessed dataset
python scripts/huggingface/download_hf.py \
--repo_id "FastVideo/Wan-Syn_77x448x832_600k" \
--local_dir "FastVideo/Wan-Syn_77x448x832_600k" \
--repo_type "dataset"
```
### Crush Smol Dataset
For Wan2.2 TI2V distillation, we use the crush_smol dataset which includes both raw videos and preprocessed latents.
```bash
# Download dataset
python scripts/huggingface/download_hf.py \
--repo_id=FastVideo/mini_i2v_dataset \
--local_dir=data/mini_i2v_dataset \
--repo_type=dataset
```
## Preprocessing for Distillation
The preprocessing steps are identical to training. Run the appropriate preprocessing script based on your model:
```bash
# For Wan2.1 T2V
bash scripts/preprocess/v1_preprocess_wan_data_t2v
# For Wan2.2 TI2V
bash examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/preprocess_wan_data_ti2v_5b.sh
```
+87
View File
@@ -0,0 +1,87 @@
# 🎯 Distillation
We introduce a new finetuning strategy - **Sparse-distill**, which jointly integrates **[DMD](https://arxiv.org/abs/2405.14867)** and **[VSA](https://arxiv.org/abs/2505.13389)** in a single training process. This approach combines the benefits of both distillation to shorten diffusion steps and sparse attention to reduce attention computation, enabling much faster video generation.
## 📊 Model Overview
We provide two distilled models:
- **[FastWan2.1-T2V-1.3B-Diffusers](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers)**: 3-step inference, up to **16 FPS** on H100 GPU
- **[FastWan2.1-T2V-14B-480P-Diffusers](https://huggingface.co/FastVideo/FastWan2.1-T2V-14B-480P-Diffusers)**: 3-step inference, up to **60x speed up** at 480P, **90x speed up** at 720P for denoising loop
- **[FastWan2.2-TI2V-5B-FullAttn-Diffusers](https://huggingface.co/FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers)**: 3-step inference, up to **50x speed up** at 720P for denoising loop
Both models are trained on **61×448×832** resolution but support generating videos with **any resolution** (1.3B model mainly support 480P, 14B model support 480P and 720P, quality may degrade for different resolutions).
## ⚙️ Inference
First install [VSA](../attention/vsa/index.md). Set `MODEL_BASE` to your own model path and run:
```bash
bash scripts/inference/v1_inference_wan_dmd.sh
```
## 🗂️ Dataset
We use the **FastVideo 480P Synthetic Wan dataset** ([FastVideo/Wan-Syn_77x448x832_600k](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x448x832_600k)) for distillation, which contains 600k synthetic latents.
### Download Dataset
```bash
# Download the preprocessed dataset
python scripts/huggingface/download_hf.py \
--repo_id "FastVideo/Wan-Syn_77x448x832_600k" \
--local_dir "FastVideo/Wan-Syn_77x448x832_600k" \
--repo_type "dataset"
```
## 🚀 Training Scripts
### Wan2.1 1.3B Model Sparse-Distill
For the 1.3B model, we use **4 nodes with 32 H200 GPUs** (8 GPUs per node):
```bash
# Multi-node training (8 nodes, 64 GPUs total)
sbatch examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P/distill_dmd_VSA_t2v_1.3B.slurm
```
**Key Configuration:**
- Global batch size: 64
- Gradient accumulation steps: 2
- Learning rate: 1e-5
- VSA attention sparsity: 0.8
- Training steps: 4000 (~12 hours)
### Wan2.1 14B Model Sparse-Distill
For the 14B model, we use **8 nodes with 64 H200 GPUs** (8 GPUs per node):
```bash
# Multi-node training (8 nodes, 64 GPUs total)
sbatch examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P/distill_dmd_VSA_t2v_14B.slurm
```
**Key Configuration:**
- Global batch size: 64
- Sequence parallel size: 4
- Gradient accumulation steps: 4
- Learning rate: 1e-5
- VSA attention sparsity: 0.9
- Training steps: 3000 (~52 hours)
- HSDP shard dim: 8
### Wan2.2 5B Model Sparse-Distill
For the 5B model, we use **8 nodes with 64 H200 GPUs** (8 GPUs per node):
```bash
# Multi-node training (8 nodes, 64 GPUs total)
sbatch examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/distill_dmd_t2v_5B.sh
```
**Key Configuration:**
- Global batch size: 64
- Sequence parallel size: 1
- Gradient accumulation steps: 1
- Learning rate: 2e-5
- Training steps: 3000 (~12 hours)
- HSDP shard dim: 1
+12
View File
@@ -0,0 +1,12 @@
# 💡 Examples
A collection of examples demonstrating usage of FastVideo.
All documented examples are autogenerated using [generate_examples.py](https://github.com/hao-ai-lab/FastVideo/blob/main/docs/generate_examples.py) from examples found in the [examples](https://github.com/hao-ai-lab/FastVideo/tree/main/examples) directory.
## Examples
- [Examples Distillation Index](distillation/examples/examples_distillation_index.md)
- [Examples Training Index](training/examples/examples_training_index.md)
- [Examples Inference Index](inference/examples/examples_inference_index.md)
+553
View File
@@ -0,0 +1,553 @@
# SPDX-License-Identifier: Apache-2.0
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/docs/source/generate_examples.py
import itertools
import re
from dataclasses import dataclass, field
from pathlib import Path
ROOT_DIR = Path(__file__).parent.parent.resolve()
ROOT_DIR_RELATIVE = '../..'
EXAMPLE_DIR = ROOT_DIR / "examples"
EXAMPLE_DOC_DIR = ROOT_DIR / "docs/getting_started/examples"
GITHUB_REPO = "hao-ai-lab/FastVideo" # Update this to your repo
def fix_case(text: str) -> str:
subs = {
"api": "API",
"cli": "CLI",
"cpu": "CPU",
"llm": "LLM",
"tpu": "TPU",
"aqlm": "AQLM",
"gguf": "GGUF",
"lora": "LoRA",
"rlhf": "RLHF",
"vllm": "vLLM",
"openai": "OpenAI",
"multilora": "MultiLoRA",
"mlpspeculator": "MLPSpeculator",
"finetune": "Finetune",
"distillation": "Distillation",
"wan": "Wan",
"i2v": "I2V",
"t2v": "T2V",
"1.3b": "1.3B",
"14b": "14B",
"480p": "480P",
"720p": "720P",
r"fp\d+": lambda x: x.group(0).upper(), # e.g. fp16, fp32
r"int\d+": lambda x: x.group(0).upper(), # e.g. int8, int16
}
for pattern, repl in subs.items():
text = re.sub(rf'\b{pattern}\b', repl, text,
flags=re.IGNORECASE) # type: ignore[call-overload]
return text
@dataclass
class Index:
"""
Index class to generate a structured document index.
Attributes:
path (Path): The path save the index file to.
title (str): The title of the index.
description (str): A brief description of the index.
caption (str): An optional caption for the table of contents.
maxdepth (int): The maximum depth of the table of contents. Defaults to 1.
documents (list[str]): A list of document paths to include in the index. Defaults to an empty list.
Methods:
generate() -> str:
Generates the index content as a string in the specified format.
""" # noqa: E501
path: Path
title: str
description: str
caption: str
maxdepth: int = 1
documents: list[str] = field(default_factory=list)
def generate(self) -> str:
content = f"# {self.title}\n\n{self.description}\n\n"
if self.caption:
content += f"## {self.caption}\n\n"
# Generate a simple list of links for MkDocs
for doc in self.documents:
# Convert document path to proper link
doc_link = doc.replace("\\", "/")
# Get just the filename for the link text
doc_title = fix_case(Path(doc).stem.replace("_", " ").title())
content += f"- [{doc_title}]({doc_link}.md)\n"
content += "\n"
return content
@dataclass
class Example:
"""
Example class for generating documentation content from a given path.
Attributes:
path (Path): The path to the main directory or file.
category (str): The category of the document.
main_file (Path): The main file in the directory.
other_files (list[Path]): list of other files in the directory.
title (str): The title of the document.
Methods:
__post_init__(): Initializes the main_file, other_files, and title attributes.
determine_main_file() -> Path: Determines the main file in the given path.
determine_other_files() -> list[Path]: Determines other files in the directory excluding the main file.
determine_title() -> str: Determines the title of the document.
generate() -> str: Generates the documentation content.
""" # noqa: E501
path: Path
category: str | None = None
main_file: Path = field(init=False)
other_files: list[Path] = field(init=False)
title: str = field(init=False)
def __post_init__(self):
self.main_file = self.determine_main_file()
self.other_files = self.determine_other_files()
self.title = self.determine_title()
def determine_main_file(self) -> Path:
"""
Determines the main file in the given path.
If the path is a file, it returns the path itself. Otherwise, it searches
for Markdown files (*.md) in the directory and returns the first one found.
Returns:
Path: The main file path, either the original path if it's a file or the first
Markdown file found in the directory.
Raises:
IndexError: If no Markdown files are found in the directory.
""" # noqa: E501
return self.path if self.path.is_file() else list(
self.path.glob("*.md")).pop()
def determine_other_files(self) -> list[Path]:
"""
Determine other files in the directory excluding the main file.
This method checks if the given path is a file. If it is, it returns an empty list.
Otherwise, it recursively searches through the directory and returns a list of all
files that are not the main file.
Returns:
list[Path]: A list of Path objects representing the other files in the directory.
""" # noqa: E501
if self.path.is_file():
return []
is_other_file = lambda file: file.is_file() and file != self.main_file
return [file for file in self.path.rglob("*")
if is_other_file(file)] # type: ignore[no-untyped-call]
def determine_title(self) -> str:
return fix_case(self.path.stem.replace("_", " ").title())
def generate(self) -> str:
# Create GitHub link to source
github_path = str(self.path.relative_to(ROOT_DIR)).replace("\\", "/")
github_url = f"https://github.com/{GITHUB_REPO}/blob/main/{github_path}"
content = f"**Source:** [{github_path}]({github_url})\n\n"
# Add title for code files
if self.main_file.suffix != ".md":
content += f"# {self.title}\n\n"
# Include main file content
if self.main_file.suffix == ".md":
# For markdown files, include the content directly
with open(self.main_file, encoding='utf-8') as f:
content += f.read() + "\n\n"
else:
# For code files, use code blocks
language = self.main_file.suffix[1:] if self.main_file.suffix else ""
with open(self.main_file, encoding='utf-8') as f:
file_content = f.read()
content += f"```{language}\n{file_content}\n```\n\n"
if not self.other_files:
return content
content += "## Additional Files\n\n"
# Define binary/non-text file extensions to skip
binary_extensions = {
'.mp4', '.avi', '.mov', '.mkv', '.gif', '.jpg', '.jpeg', '.png',
'.webp', '.bmp', '.pdf', '.zip', '.tar', '.gz', '.mp3', '.wav'
}
for file in sorted(self.other_files):
# Skip binary files
if file.suffix.lower() in binary_extensions:
continue
file_rel_path = file.relative_to(self.path)
# Use collapsible admonition syntax for MkDocs
content += f"??? note \"{file_rel_path}\"\n\n"
try:
if file.suffix == ".md":
# Include markdown content with indentation
with open(file, encoding='utf-8') as f:
for line in f:
content += f" {line}"
else:
# Include code with proper formatting
language = file.suffix[1:] if file.suffix else ""
with open(file, encoding='utf-8') as f:
file_content = f.read()
# Indent the code block for the admonition
content += f" ```{language}\n"
for line in file_content.split('\n'):
content += f" {line}\n"
content += " ```\n"
content += "\n"
except UnicodeDecodeError:
# Skip files that can't be decoded as UTF-8
continue
return content
@dataclass
class NestedStructure:
"""Helper class to manage nested documentation structures for training/distillation."""
category: str
method: str
model: str
dataset: str
example: Example
@property
def filename(self) -> str:
return f"{self.model}_{self.dataset}"
@property
def title(self) -> str:
return fix_case(self.dataset.replace('_', ' '))
@property
def description(self) -> str:
category_name = self.category.title()
return f"{category_name} example using the {self.dataset} dataset with the {self.model} model."
def create_category_indices() -> dict[str, Index]:
"""Create category indices with their respective configurations."""
main_index_dir = ROOT_DIR / "docs/examples"
if not main_index_dir.exists():
main_index_dir.mkdir(parents=True)
category_indices = {
"inference":
Index(
path=ROOT_DIR /
"docs/inference/examples/examples_inference_index.md",
title="🚀 Examples",
description=
"Inference examples demonstrate how to use FastVideo inference. We recommend starting with [basic.md](basic.md).",
caption="Examples",
maxdepth=1,
),
"training":
Index(
path=ROOT_DIR / "docs/training/examples/examples_training_index.md",
title="🚀 Examples",
description=
"Training examples demonstrate how to use FastVideo training.",
caption="Examples",
maxdepth=3,
),
"distillation":
Index(
path=ROOT_DIR /
"docs/distillation/examples/examples_distillation_index.md",
title="🚀 Examples",
description=
"Distillation examples demonstrate how to use FastVideo distillation.",
caption="Examples",
maxdepth=3,
),
}
# Ensure all category doc directories exist
for index in category_indices.values():
if not index.path.parent.exists():
index.path.parent.mkdir(parents=True)
return category_indices
def find_examples(category_indices: dict[str, Index],
generate_main_index: bool) -> list[Example]:
"""Find all examples from the examples directory."""
examples = []
glob_patterns = ["*.py", "*.md", "*.sh"]
# Map category names to actual directory names
category_dir_mapping = {
"distillation": "distill", # examples/distill/ -> distillation category
}
# Find categorised examples
for category in category_indices:
# Use mapped directory name if available, otherwise use category name
dir_name = category_dir_mapping.get(category, category)
category_dir = EXAMPLE_DIR / dir_name
# Skip if directory doesn't exist
if not category_dir.exists():
continue
globs = [category_dir.glob(pattern) for pattern in glob_patterns]
for path in itertools.chain(*globs):
examples.append(Example(path, category))
# Find examples in subdirectories (recursively)
for path in category_dir.glob("**/*.md"):
examples.append(Example(path.parent, category))
# Find uncategorised examples only if we're generating a main index
if generate_main_index:
globs = [EXAMPLE_DIR.glob(pattern) for pattern in glob_patterns]
for path in itertools.chain(*globs):
examples.append(Example(path))
# Find examples in subdirectories
for path in EXAMPLE_DIR.glob("*/*.md"):
# Skip categorised examples
if path.parent.name in category_indices:
continue
examples.append(Example(path.parent))
return examples
def create_nested_structures(
examples: list[Example]
) -> dict[str, dict[str, dict[str, dict[str, NestedStructure]]]]:
"""Create nested structures for training and distillation categories."""
nested_structures: dict[str, dict[str, dict[str,
dict[str,
NestedStructure]]]] = {}
# Map category names to actual directory names
category_dir_mapping = {
"distillation": "distill",
}
for example in examples:
if example.category not in ["training", "distillation"]:
continue
# Use mapped directory name if available
dir_name = category_dir_mapping.get(example.category, example.category)
category_dir = EXAMPLE_DIR / dir_name
relative_path = example.path.relative_to(category_dir)
path_parts = relative_path.parts
if example.category == "training":
# For training examples like finetune/wan_i2v_14b_480p/crush_smol
if len(path_parts) >= 3:
method = path_parts[0] # e.g., "finetune"
model = path_parts[1] # e.g., "wan_i2v_14b_480p"
dataset = path_parts[2] # e.g., "crush_smol"
# Initialize nested structure
if example.category not in nested_structures:
nested_structures[example.category] = {}
if method not in nested_structures[example.category]:
nested_structures[example.category][method] = {}
if model not in nested_structures[example.category][method]:
nested_structures[example.category][method][model] = {}
# Store the nested structure
nested_structures[
example.category][method][model][dataset] = NestedStructure(
category=example.category,
method=method,
model=model,
dataset=dataset,
example=example)
elif example.category == "distillation" and len(path_parts) >= 2:
# For distillation examples like Wan2.1-T2V/Wan-Syn-Data-480P
model = path_parts[0] # e.g., "Wan2.1-T2V"
dataset = path_parts[1] # e.g., "Wan-Syn-Data-480P"
method = "DMD" # Default method for distillation
# Initialize nested structure
if example.category not in nested_structures:
nested_structures[example.category] = {}
if method not in nested_structures[example.category]:
nested_structures[example.category][method] = {}
if model not in nested_structures[example.category][method]:
nested_structures[example.category][method][model] = {}
# Store the nested structure
nested_structures[
example.category][method][model][dataset] = NestedStructure(
category=example.category,
method=method,
model=model,
dataset=dataset,
example=example)
return nested_structures
def generate_flat_examples(examples: list[Example],
category_indices: dict[str, Index],
examples_index: Index | None,
generate_main_index: bool) -> None:
"""Generate documentation for flat structure examples (inference, etc.)."""
for example in examples:
if example.category in ["training", "distillation"]:
continue # Skip nested structure examples
# Determine which index to use for this example
if example.category is not None and example.category in category_indices:
index = category_indices[example.category]
elif generate_main_index:
assert examples_index is not None
index = examples_index
else:
continue
# Generate the example documentation
doc_path = index.path.parent / f"{example.path.stem}.md"
with open(doc_path, "w+") as f:
f.write(example.generate())
index.documents.append(example.path.stem)
def generate_nested_examples(nested_structures: dict[str, dict[str, dict[
str, dict[str, NestedStructure]]]], category_indices: dict[str,
Index]) -> None:
"""Generate documentation for nested structure examples (training, distillation)."""
for category_name in ["training", "distillation"]:
if category_name not in category_indices or category_name not in nested_structures:
continue
category_index = category_indices[category_name]
category_base_dir = category_index.path.parent
for method, models in nested_structures[category_name].items():
# Create method-level index
method_index = Index(path=category_base_dir / f"{method}.md",
title=fix_case(method),
description=f"Examples using {method}.",
caption=f"{fix_case(method)} Examples",
maxdepth=2)
for model, datasets in models.items():
# Generate dataset examples using the Example class
for dataset, nested_struct in datasets.items():
doc_path = category_base_dir / f"{nested_struct.filename}.md"
with open(doc_path, "w+") as f:
f.write(nested_struct.example.generate())
# Create model-level index
model_index = Index(
path=category_base_dir / f"{model}.md",
title=fix_case(model.replace('_', ' ')),
description=f"Examples for the {model} model.",
caption=f"{fix_case(model.replace('_', ' '))} Datasets",
maxdepth=1)
# Add dataset indices to model index
for dataset, nested_struct in datasets.items():
model_index.documents.append(nested_struct.filename)
# Write model index
with open(model_index.path, "w+") as f:
f.write(model_index.generate())
# Add model to method index
method_index.documents.append(model)
# Write method index
with open(method_index.path, "w+") as f:
f.write(method_index.generate())
# Add method to main category index
category_index.documents.append(method)
def generate_examples(generate_main_index: bool = False) -> None:
"""
Generate example documentation.
Args:
generate_main_index (bool): Whether to generate the main examples index.
If False, only category-specific indices will be generated.
"""
# Create category indices
category_indices = create_category_indices()
# Create the main examples index only if requested
examples_index = None
if generate_main_index:
main_index_dir = ROOT_DIR / "docs/examples"
examples_index = Index(
path=main_index_dir / "examples_index.md",
title="💡 Examples",
description=
"A collection of examples demonstrating usage of FastVideo.\n\n"
f"All documented examples are autogenerated using [generate_examples.py](https://github.com/{GITHUB_REPO}/blob/main/docs/generate_examples.py) "
f"from examples found in the [examples](https://github.com/{GITHUB_REPO}/tree/main/examples) directory.",
caption="Examples",
maxdepth=2)
# Find all examples
examples = find_examples(category_indices, generate_main_index)
# Create nested structures for training and distillation
nested_structures = create_nested_structures(examples)
# Generate flat structure examples (inference, etc.)
generate_flat_examples(examples, category_indices, examples_index,
generate_main_index)
# Generate nested structure examples (training, distillation)
generate_nested_examples(nested_structures, category_indices)
# Generate the index files for categories
for category_index in category_indices.values():
if category_index.documents:
# Add to main index if it exists
if generate_main_index and examples_index:
main_index_dir = examples_index.path.parent
rel_path = category_index.path.relative_to(
main_index_dir.parent)
examples_index.documents.insert(
0,
str(rel_path).replace(".md", ""))
# Write the category index file
with open(category_index.path, "w+") as f:
f.write(category_index.generate())
# Write the main index file if requested
if generate_main_index and examples_index:
with open(examples_index.path, "w+") as f:
f.write(examples_index.generate())
def on_pre_build_hook(config, **kwargs):
"""
MkDocs hook to generate examples before building the documentation.
This function is called automatically by the mkdocs-simple-hooks plugin.
"""
print("Generating example documentation...")
generate_examples(generate_main_index=True)
print("Example documentation generated successfully!")
if __name__ == "__main__":
print("Generating example documentation...")
generate_examples(generate_main_index=True)
print("Example documentation generated successfully!")
+45
View File
@@ -0,0 +1,45 @@
# 🔧 Installation
FastVideo supports the following hardware platforms:
- [NVIDIA CUDA](installation/gpu.md)
- [Apple silicon](installation/mps.md)
## Quick Installation
### Using pip
```bash
# Create and activate a new conda environment
conda create -n fastvideo python=3.12
conda activate fastvideo
pip install fastvideo
```
### From source
```bash
git clone https://github.com/hao-ai-lab/FastVideo.git
cd FastVideo
pip install -e .
```
Also optionally install flash-attn:
```bash
pip install flash-attn --no-build-isolation
```
## Hardware Requirements
- **NVIDIA GPUs**: CUDA 11.8+ with compute capability 7.0+
- **Apple Silicon**: macOS 12.0+ with M1/M2/M3 chips
- **CPU**: x86_64 architecture (for CPU-only inference)
## Next Steps
- [Quick Start Guide](quick_start.md) - Get started with your first video generation
- [Configuration](../inference/configuration.md) - Learn about configuration options
- [Examples](../inference/examples/examples_inference_index.md) - Explore example scripts and notebooks
@@ -1,14 +1,12 @@
(fastvideo-installation)=
# NVIDIA GPU
# 🔧 Installation
FastVideo currently only supports Linux and NVIDIA CUDA GPUs.
Instructions to install FastVideo for NVIDIA CUDA GPUs.
## Requirements
- **OS: Linux**
- **OS: Linux or Windows WSL**
- **Python: 3.10-3.12**
- **CUDA 12.4**
- **CUDA 12.8**
- **At least 1 NVIDIA GPU**
## Set up using Python
@@ -32,16 +30,8 @@ conda create -n fastvideo python=3.12 -y
conda activate fastvideo
```
:::{note}
[PyTorch has deprecated the conda release channel](https://github.com/pytorch/pytorch/issues/138506). If you use `conda`, please only use it to create Python environment rather than installing packages.
:::
#### uv
:::{tip}
We highly recommend using `uv` to install FastVideo. In our experience, `uv` speeds up installation by at least 3x.
:::
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
```console
@@ -62,7 +52,7 @@ uv pip install fastvideo
Also optionally install flash-attn:
```bash
pip install flash-attn==2.7.4.post1 --no-build-isolation
pip install flash-attn --no-build-isolation
```
### Installation from Source
@@ -89,22 +79,22 @@ uv pip install -e .
#### Flash Attention
```bash
pip install flash-attn==2.7.4.post1 --no-build-isolation
pip install flash-attn --no-build-isolation
```
## Set up using Docker
We also have prebuilt docker images with FastVideo dependencies pre-installed:
[Docker Images](#docker)
[Docker Images](../../contributing/developer_env/docker.md)
## Development Environment Setup
If you're planning to contribute to FastVideo please see the following page:
[Contributor Guide](#developer-overview)
[Contributor Guide](../../contributing/overview.md)
## Hardware Requirements
### For Basic Inference
- NVIDIA GPU with CUDA 12.4 support
- NVIDIA GPU with CUDA 12.8 support
### For Lora Finetuning
- 40GB GPU memory each for 2 GPUs with lora
+93
View File
@@ -0,0 +1,93 @@
# MPS (Apple Silicon)
Instructions to install FastVideo for Apple Silicon.
## Requirements
- **OS: MacOS**
- **Python: 3.12.4**
## Set up using Python
### Create a new Python environment
#### Conda
You can create a new python environment using [Conda](https://docs.conda.io/projects/conda/en/stable/user-guide/getting-started.html)
##### 1. Install Miniconda (if not already installed)
```bash
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-MacOSX-arm64.sh
bash Miniconda3-latest-MacOSX-arm64.sh
source ~/.zshrc
```
##### 2. Create and activate a Conda environment for FastVideo
```bash
# (Recommended) Create a new conda environment.
conda create -n fastvideo python=3.12.4 -y
conda activate fastvideo
```
#### uv
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
```console
# (Recommended) Create a new uv environment. Use `--seed` to install `pip` and `setuptools` in the environment.
uv venv --python 3.12 --seed
source .venv/bin/activate
```
### Dependencies
```
brew install ffmpeg
```
### Installation
```bash
pip install fastvideo
# or if you are using uv
uv pip install fastvideo
```
### Installation from Source
#### 1. Clone the FastVideo repository
```bash
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
```
#### 2. Install FastVideo
Basic installation:
```bash
pip install -e .
# or if you are using uv
uv pip install -e .
```
## Development Environment Setup
If you're planning to contribute to FastVideo please see the following page:
[Contributor Guide](../../contributing/overview.md)
## Hardware Requirements
### For Basic Inference
- Mac M1, M2, M3, or M4 (at least 32 GB RAM is preferable for high quality video generation)
## Troubleshooting
If you encounter any issues during installation, please open an issue on our [GitHub repository](https://github.com/hao-ai-lab/FastVideo).
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ) for additional support.
+83
View File
@@ -0,0 +1,83 @@
# 🚀 Quick Start
Get up and running with FastVideo in minutes!
## Installation
First, install FastVideo:
```bash
# Create and activate a new conda environment
conda create -n fastvideo python=3.12
conda activate fastvideo
# Install FastVideo
pip install fastvideo
```
Also optionally install flash-attn:
```bash
pip install flash-attn --no-build-isolation
```
## Basic Usage
### Text-to-Video Generation
```python
from fastvideo import VideoGenerator
def main():
# Create a video generator with a pre-trained model
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1, # Adjust based on your hardware
)
# Define a prompt for your video
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
# Generate the video
video = generator.generate_video(
prompt,
return_frames=True, # Also return frames from this call (defaults to False)
output_path="my_videos/", # Controls where videos are saved
save_video=True
)
if __name__ == '__main__':
main()
```
### Image-to-Video Generation
```python
from fastvideo import VideoGenerator, SamplingParam
def main():
# Create the generator
model_name = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
generator = VideoGenerator.from_pretrained(model_name, num_gpus=1)
# Set up parameters with an initial image
sampling_param = SamplingParam.from_pretrained(model_name)
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
sampling_param.num_frames = 107
# Generate video based on the image
prompt = "A photograph coming to life with gentle movement"
generator.generate_video(prompt, sampling_param=sampling_param,
output_path="my_videos/",
save_video=True)
if __name__ == '__main__':
main()
```
## Next Steps
- [Installation Guide](installation.md) - Detailed installation instructions
- [Configuration](../inference/configuration.md) - Learn about configuration options
- [Examples](../inference/examples/) - Explore more examples
- [Optimizations](../inference/optimizations.md) - Performance optimization tips

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