Compare commits

...
Author SHA1 Message Date
JerryZhou54 140fb9f20e Fix test for forward_train 2025-09-15 23:23:33 +00:00
JerryZhou54 2078876b98 Ensure 0 numerical diff for forward_train 2025-09-15 22:49:52 +00:00
JerryZhou54 a953f46bd6 Add test for _forward_train 2025-09-15 22:37:50 +00:00
JerryZhou54 adae957008 Fix numerical diff between causal_wanvideo.py and SF's causal wan 2025-09-14 08:09:57 +00:00
RandNMR73 93ebd15a0d text preprocessing ready 2025-09-10 11:12:48 +00:00
JerryZhou54 1110474065 checkpoint 2025-09-10 08:57:54 +00:00
JerryZhou54 80baffd540 Enable timestep warping & using SelfForcing scheduler 2025-09-09 23:30:55 +00:00
JerryZhou54 918180048e Stop backprop through kv_cache 2025-09-09 10:02:03 +00:00
RandNMR73 b7dbd7cb9e new branch 2025-09-09 10:02:00 +00:00
RandNMR73 71159b6416 inference works after changes added 2025-09-09 10:01:31 +00:00
William LinandRandNMR73 ac11127397 [Self-forcing] [1/n] Handle extra dim in time embedding and add timestep warping (#792)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
2025-09-09 02:52:02 -07:00
Eric LiangandEricLiang e028dcc7c0 [Backend][Vmoba] Add implementation of VMoba (#778)
Co-authored-by: EricLiang <https://github.com/EricLina>
2025-09-08 23:53:25 -07:00
Wenxuan Tanandgemini-code-assist[bot] 076f45c1ee [Feature] Support Lora for DMD (#755)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2025-09-08 14:18:21 -07:00
85eb7265db fix: lora_B init zeros (#781)
Co-authored-by: zbchu2 <zbchu2@iflytek.com>
Co-authored-by: Wenxuan Tan <wenxuan.tan@wisc.edu>
2025-09-05 22:56:52 -07:00
William Lin d3ceb67e66 [misc] Update Slack invite link (#786) 2025-09-05 12:16:18 -07:00
Zhang Peiyuan 7ac153a5ca Update WeChat Link 2025-09-05 11:40:47 -07:00
William Lin d1e7aa0abd [CI] Add ssim test for causal inference (#784) 2025-09-05 01:23:01 -07:00
William Lin 2d846c55a1 [misc] Improve text encoding stage (#774) 2025-09-04 17:51:27 -07:00
Jinzhe Pan b318063c0a [Preprocess][Fix] video quality issue (#773) 2025-09-03 20:47:33 -07:00
Jinzhe Pan 4aa307be55 [Preprocess][Feat] support torchvision to load video in new preprocessing (#761) 2025-09-01 23:37:01 -07:00
William Lin 055e52e5ea [misc] [VSA] [STA] fix tk_root in setup.py for VSA and STA (#772) 2025-08-29 01:13:37 -07:00
William Lin 7d2069596b [bugfix] [VSA] [STA] Fix MANIFEST.in for VSA and STA; Move tk into both directories (#771) 2025-08-29 00:51:05 -07:00
William Lin c45009c9a4 [bugfix] fix STA install setup.py import (#770) 2025-08-28 23:02:53 -07:00
William LinandPeiyuan Zhang b91020b407 [VSA] [STA] Fix directory structure for pypi publishing (#769)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-08-28 22:34:03 -07:00
William Lin 2dcc5ea4f6 [chore] Release 0.1.6 (#768) 2025-08-28 20:56:21 -07:00
Wei ZhouandSolitaryThinker 359151d9a0 [Feature] Add wan2.2 5b i2v (#760)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-08-28 18:15:59 -07:00
Wei ZhouandSolitaryThinker ce67cd3729 [Feat] Support Self-Forcing's Causal Inference for Wan2.1 T2V 1.3B (#766)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-08-28 16:47:49 -07:00
Zhang Peiyuan 7c554e5da8 Update Community Link (#765) 2025-08-27 16:12:47 -07:00
William Lin 663ea33ff1 [bugfix] Fix wrong HF model string for FastWan2.2 5B (#763) 2025-08-26 22:05:40 -07:00
William Lin 3ef04f1654 [misc] [docs] Various fixes for logging and docs (#758) 2025-08-23 21:13:50 -07:00
Jinzhe Pan 0eced76a41 [Feat][Preprocess] support multi-gpus (#753) 2025-08-23 11:34:42 +08:00
Jinzhe Pan 3ab6470d1a [Feat][Preprocess] support merged dataset (#752) 2025-08-22 15:29:33 -07:00
Wenxuan Tan 989a03532c Optionally use unmerged weights for inference (#745) 2025-08-22 15:20:31 -07:00
William Lin fa15369a02 [bugfix] Check that model_index.json module is in required_modules list before removing (#756) 2025-08-22 14:36:44 -07:00
Zhang Peiyuan 78a9cb88d8 [Fix] fix seed in dmd denoising loop (#736) 2025-08-21 18:06:16 -07:00
Peng Xiaoand肖鹏 a0bff12746 [bugfix] [dmd] Align backward simulation with dmd2 sample back (#744)
Co-authored-by: 肖鹏 <xiaopeng1@aishi.ai>
2025-08-20 22:25:33 -07:00
William Lin 98f2af94e5 [bugfix] Missing Docker file for cuda12.9 (#750) 2025-08-20 15:34:31 -07:00
William Lin 46f7b6d574 [Docker] add 12.9 docker image and also fix py3.10 and py3.11 dockerfile (#749) 2025-08-20 15:31:15 -07:00
Jinzhe Pan 911a6a6a35 [Feat][Preprocessing] i2v preprocessing workflow (#737) 2025-08-14 20:47:25 -07:00
Zhang Peiyuan 38c7949d5c Update WeChat group link (#739) 2025-08-14 15:03:35 -07:00
Jinzhe Pan 7e7a0dba9d feat: preprocess validation dataset only when exist (#734) 2025-08-12 02:16:31 -07:00
Zhang Peiyuan f62e210ae6 Fix vsa backward gQ (#735) 2025-08-11 21:43:13 -07:00
William Lin 6ceb4942a0 [bugfix] [dmd] Fix backward simulation and also naming in wan_i2v_dmd_pipeline (#731) 2025-08-10 21:13:30 -07:00
William LinandRandNMR73 8cae5e4708 [feature] add Gradio live serving demo code (#727)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
2025-08-10 15:34:03 -07:00
William Lin 2a773fa34e [bugfix] [distill] remove i2v validation schema import in distill (#728) 2025-08-09 20:47:42 -07:00
Wenxuan Tan 5357f63327 Fix LoRA load from training checkpoint (#719) 2025-08-09 20:46:00 -05:00
William Lin 60f61c8101 [bugfix] fix pyproject install and VSA precision test (#726) 2025-08-08 18:45:03 -07:00
Jiali Chen 3d75ba8251 update version selection for VSA workflow (#725) 2025-08-08 13:05:16 -07:00
Wenxuan Tan 6c6bcd914d Remove all empty_cache (#713) 2025-08-07 22:50:38 -07:00
Jiali Chen f79b08de81 add cicd workflow for publishing VSA kernel (#723) 2025-08-07 18:53:05 -07:00
Jinzhe Pan f2bc037fff [Fix] training pipeline pin_cpu_memory issue (#692) 2025-08-07 02:31:20 -07:00
Jinzhe Pan 86604a684b [3/3][Preprocess] add preprocessing workflows (#645) 2025-08-07 01:49:07 -07:00
Zhang Peiyuan 47bd1e0178 [Misc] change installation logic of vsa (#721) 2025-08-06 21:54:09 -07:00
Wei ZhouandSolitaryThinker c41305ad18 [Feat] Add Wan2.2 14B MoE (#688)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-08-06 20:31:03 -07:00
Zhang Peiyuan 98ce9034f0 [Chore] Include our demo in the readme. (#720) 2025-08-06 19:29:40 -07:00
William Lin 0ceff110da [chore] Release 0.1.5 (#717) 2025-08-06 13:07:52 -07:00
Yongqi Chen 1d018acb3e [Feature]Add Data-free distillation readme (#710) 2025-08-05 14:27:39 -04:00
Yongqi Chen 7d8cf38dbe Fix typo (#709) 2025-08-04 20:21:21 -07:00
Yongqi Chen 8d483fe4aa [Bugfix] Fix neg_prompt bug when training from local cp (#708) 2025-08-04 15:54:06 -07:00
Zhang Peiyuan c1191250bf Add WeChat group link (#707) 2025-08-04 15:19:01 -07:00
Wenxuan Tan 4b7266349a [misc] Remove allow_tf32 in scripts (#705) 2025-08-04 15:37:56 -05:00
Yongqi Chen 22f9b7681f [Feature]Update Wan2.2+DMD doc example (#706) 2025-08-04 16:14:22 -04:00
Yongqi Chen 589d32cc39 [Feature] Update Readme and scripts (#703) 2025-08-04 15:02:32 -04:00
Hao Zhang 89199837db Update readme pre-release (#704) 2025-08-04 11:54:31 -07:00
William Lin d6ebaf1b49 [Docs] Fix README (#701) 2025-08-04 11:27:27 -07:00
Yongqi Chen fac927777c [Feature] Update readme (#702) 2025-08-04 14:27:18 -04:00
William Lin ecbd697dae [misc] Readme fixes (#699) 2025-08-04 10:20:57 -07:00
Yongqi Chen 7d4acef64d [Feature] Update sparse distill readme and doc (#700) 2025-08-04 10:16:20 -07:00
William Lin 9f0ce517cf [Docs] Update README and docs for FastWan (#698) 2025-08-04 09:05:18 -07:00
Yongqi Chen c718e56b0d [Feature] Remove unused args (#695) 2025-08-03 23:01:54 -04:00
Yongqi ChenandSolitaryThinker b65f0316d1 [Feature] Add Wan2.2 DMD example files; Update lr scheduler (#694)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-08-03 22:55:01 -04:00
William Lin 8d8bcb76b0 [config] Add config for FastWan2.2 ti2v 5B (#693) 2025-08-03 19:09:30 -07:00
Yongqi ChenandSolitaryThinker 5f42748ed1 [Feature] Add Wan2.2-TI2V-5B Sparse Distill (#690)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-08-03 01:29:57 -04:00
Yongqi Chen c9005045dc [Feature[[Readme] Add VSA/DMD doc (#673) 2025-08-02 02:35:46 -04:00
Wenxuan Tan 6c81befc87 [Feature] Optionally enable torch compile (#684) 2025-08-01 20:17:40 -07:00
Yongqi Chen dfe0b288e1 [Bugfix] Add i2v vae loading (#686) 2025-08-01 23:15:37 -04:00
Wenxuan Tanandgemini-code-assist[bot] 31200fbb83 [Misc] Fix training scripts (#683)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2025-08-01 15:54:31 -05:00
Yongqi Chen 9185978c55 [Bugfix] Fix multi-gpu training lr_scheduler (#682) 2025-08-01 15:54:03 -04:00
MartinPernus fcba463553 [Bugfix] fix _normalize_dit_input (#681) 2025-08-01 05:10:27 -04:00
Yongqi Chen 2c53d3eecf [Feature]Add DMD visualization for debugging (#674) 2025-07-31 05:54:47 -04:00
Zhang Peiyuan 516ecd374a [Misc] Update examples/ and other misc (#672) 2025-07-30 19:10:27 -07:00
Wei Zhou 3b1b54a74d Modify args to make sure the scripts are runnable on 4090 (#671) 2025-07-30 14:55:08 -07:00
Yongqi Chen 6914e7c904 [Bugfix]Fix DMD pipeline registry (#670) 2025-07-30 13:21:35 -07:00
Sopiko Kurdadze 5452369749 [Feature] [Inference]Add ROCm platform support for single-gpu inference (#669) 2025-07-30 12:57:02 -07:00
Yongqi Chen a113311e77 [Bugfix][Training]Fix Wan2.2 training vae config issue (#668) 2025-07-30 12:24:45 -07:00
Kevin Lin 44da97da92 [chore] Release 0.1.4 (#667) 2025-07-30 01:02:36 -07:00
Yongqi Chen f759980a58 [Feature]Add VSA slurm training example scripts (#666) 2025-07-30 01:27:54 -04:00
Zhang Peiyuan 37e0f8c236 [BUG] Fix distillation + vsa (#665) 2025-07-29 19:48:11 -07:00
Kevin Lin 51711d5906 [ComfyUI] Add __init__.py for node discovery (#663) 2025-07-29 18:21:09 -07:00
William LinandJerryZhou54 6375223b16 [Feature] Add wan2.2 5B T2V (#658)
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
2025-07-29 17:16:03 -07:00
Yongqi Chen 4cb046768d [Feature]Add DMD distillation training resume checkpoint; Update DMD CI test (#662) 2025-07-29 19:06:05 -04:00
Yongqi Chen 3322542444 [Feature] Add DMD CI test (#661) 2025-07-29 03:14:00 -04:00
Yongqi Chen 65f707354b [Bugfix]Fix mdoel inference checkpoint saving when enabling HSDP (#660) 2025-07-28 22:39:06 -07:00
Yongqi Chen 109e2e7e9d [Bugfix]Fix DMD wan pipeline (#659) 2025-07-28 21:50:44 -07:00
Zhang Peiyuan cbc3a6bb9d [Feat] Support VSA with any resolution. (#650) 2025-07-28 20:14:40 -07:00
Yongqi Chen 2fa8d4ae6d [Feature][Distill]Add 14B 480p T2V distill example scripts (#655) 2025-07-28 18:36:31 -04:00
Jinzhe Pan 7b6c8aee99 [2/3][Preprocess] refactor pipeline registry & file structure (#639) 2025-07-27 23:30:17 -07:00
Yongqi Chen 6284eaa363 [Feature][Distill]Add DMD+VSA joint training example (#654) 2025-07-27 18:18:01 -04:00
Yongqi Chen 636524e87f [Feature] Add Wan-14B-T2V-VSA CLI inference; add master port args (#653) 2025-07-27 07:13:44 -04:00
Yongqi Chen 202b2f3972 [Feature] Ignore [union-attr] and [override] mypy check and remove from training (#652) 2025-07-27 04:34:38 -04:00
Yongqi Chen 247fe273d8 [Feature] Add DMD T2V training pipeline (#651) 2025-07-27 03:35:51 -04:00
William Lin cb320dfa3a [bugfix] VideoGenerator improperly extracts output_video_name (#649) 2025-07-26 19:46:26 -07:00
Kevin Lin d8bb5abc46 [CI] Fix ComfyUI publisher ID (#648) 2025-07-25 19:02:52 -07:00
Kevin Lin cc703eca51 [CI] Add publish workflow for ComfyUI (#647) 2025-07-25 18:39:05 -07:00
William Lin 81c9df629c [core] Add offloading for vae and image encoder and rename offloading args (#643) 2025-07-25 17:55:03 -07:00
Yongqi Chen d3c0c52208 [Feature] Add prompt_txt support for CLI inference; Add DMD CLI inference (#646) 2025-07-25 19:44:10 -04:00
William Lin 744e0555c0 [misc] Use FASTVIDEO_STAGE_LOGGING for perf timing of stage (#644) 2025-07-25 16:20:53 -07:00
Jinzhe Pan 3a38f7dfdc [1/3][Preprocess] refactor preprocessing configs (#638) 2025-07-25 14:27:12 -07:00
William Lin f572319bd9 [Feature] Remove V1 folder (#642) 2025-07-24 22:43:12 -07:00
Wenxuan TanandWei Feng 48528f468c [Feature] Multi-lora inference (#640)
Co-authored-by: Wei (Will) Feng <134637289+weifengpy@users.noreply.github.com>
2025-07-24 21:01:28 -07:00
Yongqi ChenandSolitaryThinker 4264a80ca9 [Feature] Add DMD inference pipeline (#637)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-07-24 21:00:47 -07:00
William Lin 8573d4f05e [Docs] Docs update for Training and MPS (#641) 2025-07-24 19:23:58 -07:00
Wenxuan Tan 210a733515 [Bugfix] Fix LoRA trainable params and training ckpt loading (#630) 2025-07-23 20:01:40 -07:00
William Lin 0aef0e6f63 [bugfix] Fix preprocessing pipelines and nightly tests (#633) 2025-07-22 22:44:25 -07:00
Kevin Lin dd022ad9be [CI] Fix CI for pull request targets other than main (#632) 2025-07-22 21:00:03 -07:00
William Lin 832ad61e5b [bugfix] fa3 no longer returns lse (#631) 2025-07-22 18:31:34 -07:00
Wenxuan Tan 9419c04ee3 Fix lora train steps (#627) 2025-07-21 23:29:30 -05:00
Zhang Peiyuanandroot a37b39d83c Py/add triton block sparse (#593)
Co-authored-by: root <a1286225768@gmail,com>
2025-07-17 16:50:40 -07:00
Wenxuan Tan bb8c769c8e [LoRA] Support v1 LoRA training (#576) 2025-07-17 15:52:09 -05:00
William Lin 576c214f28 [v0] Remove V0 code (#621) 2025-07-15 22:05:16 -07:00
RandNMR73 b79d1fc15b video gen working on apple silicon (addressed issues from prior pr) (#595) 2025-07-15 22:04:27 -07:00
William Lin eb66e1c18d [chore] release 0.1.2 (#622) 2025-07-15 14:42:54 -07:00
William Lin 616d43c1cf [bugfix] [training] use separate generator for validation (#610) 2025-07-15 13:19:38 -07:00
Wenxuan Tan 7244a4b27f [CI] Add LoRA inference tests (#546) 2025-07-15 15:06:44 -05:00
Yongqi Chen 7e5ebb4582 [Feature][Training]Update example fine-tuning scripts to enable gradient checkpointing (#618) 2025-07-15 11:22:56 -07:00
Wenxuan Tan 65ed588570 Set encoder TP size to 1 by default (#569) 2025-07-09 17:44:24 -05:00
Wenxuan Tan 14adfe2edc Remove all unnecessary torch.cuda.empty_cache (#606) 2025-07-09 16:45:20 -05:00
William Lin 6198c6a640 [docs] update dev guide runpod image to py3.12 (#602) 2025-07-07 19:56:07 -05:00
William Lin e6b71b531b [docs] Update slack invite (#601) 2025-07-07 14:20:42 -05:00
William Lin ae1d112c6a [bugfix] [training] fix deadlock in latent datasets and init error in multi-node training (#598) 2025-07-06 01:34:15 -05:00
William Lin bf4de1f38f [chore] Upgrade min Python version from 3.8 to 3.10 (#597) 2025-07-04 22:08:31 -05:00
William Lin 66fdcc8e76 [Training] Use inference pipeline for training validation (#585) 2025-07-04 17:23:55 -05:00
Wenxuan Tan ed1e8d6bad [Feature] Offload all text encoders by default (#594) 2025-07-03 19:30:14 -05:00
Kevin Lin b9423ca3f8 Add ComfyUI custom node for inference (#596) 2025-07-03 16:09:03 -05:00
Wenxuan Tan ad16289871 [LoRA] Fix lora merge weights (#579) 2025-07-02 00:12:52 -05:00
Wenxuan Tan 2a41da1e6b Fix VAE precisions (#588) 2025-07-01 14:49:04 -05:00
William Lin 32133171da [chore] Release 0.1.1 (#592) 2025-07-01 01:43:12 -05:00
Kevin Lin 19674c6f29 [CI] Fix fork builds (#590) 2025-07-01 01:03:34 -05:00
Yongqi Chen 508afb7002 [docs] Update Readme (#591) 2025-06-30 23:48:45 -05:00
Yongqi Chen 288ea88105 [Feat][Training] Rename weight conversion function and update gradient checkpoint in scripts (#589) 2025-07-01 00:20:02 -04:00
Jinzhe Pan eb0f1318f3 [Feat] activation checkpointing (#584) 2025-06-30 15:24:29 -05:00
William Lin ce9b5910cc [Training] add caption to validation log (#582) 2025-06-30 02:42:17 -05:00
William Lin d0e5a6214a [misc] [training] Add --video_length_tolerance_range 10 to preprocessing scripts (#581) 2025-06-30 02:21:22 -05:00
Wenxuan Tan 834562b2db [CI] Fix pre-commit CI (#578) 2025-06-29 16:52:29 -05:00
Wei (Will) Feng 060cc7b9ba fully_shard usage on RMSNorm (#577) 2025-06-29 16:35:24 -05:00
Yongqi Chen 6c58a5ba62 [Bugfix]Fix VSA sp for training/inference (#574) 2025-06-29 13:44:33 -05:00
William Lin 48d9f61f86 [ci] [misc] fix training test threshold (#573) 2025-06-28 22:17:18 -05:00
William Lin 5f938b5844 [Revert] "[Feature] Load weights from distributed" (#571) 2025-06-28 20:55:14 -05:00
Wenxuan Tan 74da2a7370 Fix CLIP config (#568) 2025-06-28 19:01:23 -05:00
637 changed files with 38857 additions and 32323 deletions
+107 -54
View File
@@ -1,148 +1,201 @@
env:
IMAGE_VERSION: "py3.12-latest"
BUILDKITE_CLEAN_CHECKOUT: true
steps:
- label: "pre-commit"
command: ".buildkite/scripts/pre_commit.sh"
agents:
queue: "default"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- wait
- 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 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Encoder Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- 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 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "VAE Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- 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:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Transformer Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=transformer
agents:
queue: "default"
- path:
- "fastvideo/v1/**/*.py"
- "fastvideo/**/*.py"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 60m .buildkite/scripts/pr_test.sh"
command: "timeout 45m .buildkite/scripts/pr_test.sh"
label: "SSIM Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=ssim
agents:
queue: "default"
- path:
- "fastvideo/v1/**"
- "fastvideo/tests/lora/**"
- "fastvideo/models/loader/**"
- "fastvideo/tests/transformers/**"
- "fastvideo/pipelines/**"
- "fastvideo/layers/lora/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "LoRA Inference Tests"
env:
- TEST_TYPE=inference_lora
agents:
queue: "default"
- path:
- "fastvideo/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Training Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=training
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 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Distillation DMDTests"
env:
- TEST_TYPE=distillation_dmd
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/**"
- "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"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Training Tests VSA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=training_vsa
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/**"
- "csrc/attn/sliding_tile_attn/**"
- "csrc/attn/sliding_tile_attn/setup.py"
- "csrc/attn/sliding_tile_attn/config_sta.py"
- "csrc/attn/sliding_tile_attn/st_attn.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Inference Tests STA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=inference_sta
agents:
queue: "default"
- path:
- "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/setup.py"
- "csrc/attn/sliding_tile_attn/config_sta.py"
- "csrc/attn/sliding_tile_attn/st_attn.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests STA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=precision_sta
agents:
queue: "default"
- path:
- "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/tests/test_vsa.py"
- "csrc/attn/video_sparse_attn/setup.py"
- "csrc/attn/video_sparse_attn/config_vsa.py"
- "csrc/attn/video_sparse_attn/vsa.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VSA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=precision_vsa
agents:
queue: "default"
- path:
- "csrc/attn/vmoba_attn/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VMoBA"
env:
- TEST_TYPE=precision_vmoba
agents:
queue: "default"
- path:
- "csrc/attn/vmoba_attn/vmoba/**"
- "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"
+23 -2
View File
@@ -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"
@@ -58,7 +58,7 @@ if [ -z "${TEST_TYPE:-}" ]; then
fi
log "Test type: $TEST_TYPE"
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT IMAGE_VERSION=$IMAGE_VERSION"
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$BUILDKITE_PULL_REQUEST IMAGE_VERSION=$IMAGE_VERSION"
case "$TEST_TYPE" in
"encoder")
@@ -81,6 +81,10 @@ case "$TEST_TYPE" in
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"
@@ -97,6 +101,23 @@ case "$TEST_TYPE" in
log "Running precision VSA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_VSA"
;;
"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
"inference_vmoba")
log "Running V-MoBA inference tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_vmoba"
;;
"precision_vmoba")
log "Running V-MoBA precision tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_vmoba"
;;
*)
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
+2 -2
View File
@@ -8,14 +8,14 @@ on:
- main
paths:
- "docs/**/*.md"
- "fastvideo/v1/examples/**/*.py"
- "fastvideo/examples/**/*.py"
pull_request:
branches:
- main
types: [opened, ready_for_review, synchronize, reopened]
paths:
- "docs/**/*.md"
- "fastvideo/v1/examples/**/*.py"
- "fastvideo/examples/**/*.py"
# Allows you to run this workflow manually from the Actions tab
workflow_dispatch:
+1 -1
View File
@@ -13,4 +13,4 @@
]
}
]
}
}
+37 -36
View File
@@ -104,48 +104,49 @@ jobs:
- 'pyproject.toml'
- '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:
@@ -167,7 +168,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 +186,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 +204,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 +230,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 +249,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 }}
@@ -264,11 +265,11 @@ jobs:
with:
job_id: "training-test-VSA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 1
gpu_count: 2
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/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 }}
@@ -284,11 +285,11 @@ jobs:
with:
job_id: "inference-test-STA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 1
gpu_count: 2
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/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,7 +327,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] && 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 }}
@@ -343,7 +344,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/nightly/test_e2e_overfit_single_sample.py -vs"
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 }}
+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/
+7 -1
View File
@@ -20,6 +20,7 @@ samples/
data/
outputs/
outputs_video
checkpoints/
sbatch.sh
*.out
env
@@ -40,6 +41,8 @@ eggs/
docs/_build/
docs/source/getting_started/examples/
docs/source/inference/examples/
docs/source/training/examples/
docs/source/distillation/examples/
# VSCode
.vscode/
@@ -55,7 +58,10 @@ 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
!comfyui/assets/**/*.png
!comfyui/assets/**/*.gif
dmd_t2v_output/
+6 -2
View File
@@ -1,3 +1,7 @@
[submodule "csrc/attn/tk"]
path = csrc/attn/tk
[submodule "csrc/attn/video_sparse_attn/tk"]
path = csrc/attn/video_sparse_attn/tk
url = https://github.com/HazyResearch/ThunderKittens.git
[submodule "csrc/attn/sliding_tile_attn/tk"]
path = csrc/attn/sliding_tile_attn/tk
url = https://github.com/HazyResearch/ThunderKittens.git
+4 -3
View File
@@ -3,7 +3,7 @@ default_stages:
- manual # Run in CI
exclude: |
(?x)(
fastvideo/v1/third_party/.*|
fastvideo/third_party/.*|
csrc/.*|
assets/.*|
tests/.*|
@@ -22,6 +22,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
)
@@ -60,7 +61,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 +70,7 @@ repos:
entry: bash
args:
- -c
- 'git ls-files | grep -v "^fastvideo/v1/tests/ssim/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
- 'git ls-files | grep -v "^fastvideo/tests/ssim/" | grep -v "^fastvideo/tests/inference/lora/L40S_reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
language: system
always_run: true
pass_filenames: false
+71 -50
View File
@@ -1,37 +1,41 @@
<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.**
**FastVideo is a unified post-training and inference 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.
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
<p align="center">
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> <b>Slack</b> </a> |
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</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.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/S7HLCSTh" target="_blank"> <b> WeChat </b> </a> |
</p>
<div align="center">
<img src=assets/perf.png width="90%"/>
<img src=assets/fastwan.png width="90%"/>
</div>
## NEWS
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
- ```2025/02/18```: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
## Key Features
FastVideo has the following features:
- End-to-end post-training support:
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 to achineve >50x denoising speedup
- Data preprocessing pipeline for video data
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs
- State-of-the-art performance optimizations for inference
- [Video Sparse Attention](https://arxiv.org/pdf/2505.13389)
- [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.
- Diverse hardware and OS support
- Support H100, A100, 4090
- Support Linux, Windows, MacOS
## Getting Started
We recommend using an environment manager such as `Conda` to create a clean environment:
@@ -47,17 +51,31 @@ pip install fastvideo
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) 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.html) 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.html). 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
)
@@ -90,60 +108,63 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html)
## 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.html)
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
## 📑 Development Plan
<!-- - More distillation methods -->
<!-- - [ ] Add Distribution Matching Distillation -->
- More models support
<!-- - [ ] Add CogvideoX model -->
- [x] Add StepVideo to V1
- Optimization features
- [x] Teacache in V1
- [x] SageAttention in V1
- Code updates
- [x] V1 Configuration API
- [ ] Support Training in V1
More FastWan Models Coming Soon!
- [ ] Add FastWan2.1-T2V-14B
- [ ] Add FastWan2.2-T2V-14B
- [ ] Add FastWan2.2-I2V-14B
<!-- - Optimization features
- Code updates -->
<!-- - [ ] fp8 support -->
<!-- - [ ] faster load model and save model support -->
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/468).
## 🤝 Contributing
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview.html)
## 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{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},
@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{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{zhang2025vsa,
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
author={Zhang, Peiyuan and Huang, Haofeng and Chen, Yongqi and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
journal={arXiv preprint arXiv:2505.13389},
year={2025}
}
@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

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

-24
View File
@@ -1,24 +0,0 @@
# Configuration for Cog ⚙️
# Reference: https://cog.run/yaml
build:
gpu: true
cuda: "12.1"
python_version: "3.10"
python_packages:
- "torch==2.4.0"
- "torchvision"
- "ninja==1.11.1.3"
- "transformers==4.46.1"
- "git+https://github.com/huggingface/diffusers.git@bf64b32652a63a1865a0528a73a13652b201698b"
- "accelerate==1.0.1"
- "safetensors==0.4.5"
- "peft==0.13.2"
- "packaging==24.2"
- "git+https://github.com/hao-ai-lab/FastVideo"
run:
- FLASH_ATTENTION_SKIP_CUDA_BUILD=TRUE pip install flash-attn --no-build-isolation
- curl -o /usr/local/bin/pget -L "https://github.com/replicate/pget/releases/latest/download/pget_$(uname -s)_$(uname -m)" && chmod +x /usr/local/bin/pget
predict: "predict.py:Predictor"
@@ -16,7 +16,7 @@ import sys
# Run it with `python collect_env.py` or `python -m torch.utils.collect_env`
from collections import namedtuple
from fastvideo.v1.envs import environment_variables
from fastvideo.envs import environment_variables
try:
import torch
@@ -62,6 +62,7 @@ SystemEnv = namedtuple(
DEFAULT_CONDA_PATTERNS = {
"torch",
"numpy",
"mypy"
"cudatoolkit",
"soumith",
"mkl",
+138
View File
@@ -0,0 +1,138 @@
# ComfyUI-FastVideo
A custom node suite for ComfyUI that provides accelerated video generation using [FastVideo](https://github.com/hao-ai-labs/FastVideo). See the [blog post](https://hao-ai-lab.github.io/blogs/fastvideo/) about FastVideo V1 to learn more.
## Multi-GPU Parallel Inference
One of the key features ComfyUI-FastVideo brings to ComfyUI is its ability to distribute the generation workload across multiple GPUs, resulting in significantly faster inference times.
![Wan2.1-I2V-14B-480P-Diffusers](./assets/wani2v.gif)
Example of Wan2.1-I2V-14B-480P-Diffusers model running on 4 GPUs.
## Features
- Generate high-quality videos from text prompts and images
- Configurable video parameters (prompt, resolution, frame count, FPS)
- Support for multiple GPUs with tensor and sequence parallelism
- Advanced configuration options for VAE, Text Encoder, and DIT components
- Interruption/cancellation support for long-running generations
## Installation
### Requirements
- [ComfyUI](https://github.com/comfyanonymous/ComfyUI)
- CUDA-capable GPU(s) with sufficient VRAM
### Install using ComfyUI Manager
Coming soon!
### Manual Installation
#### Copy the FastVideo `comfyui` directory into your ComfyUI custom_nodes directory:
```bash
cp -r /path/to/FastVideo/comfyui /path/to/ComfyUI/custom_nodes/FastVideo
```
#### Install dependencies:
Currently, the only dependency is `fastvideo`, which can be installed using pip.
```bash
pip install fastvideo
```
#### Install missing custom nodes:
`ComfyUI-VideoHelperSuite`:
```bash
cd /path/to/ComfyUI/custom_nodes
git clone https://github.com/Kosinkadink/ComfyUI-VideoHelperSuite.git
```
If you're seeing `ImportError: libGL.so.1: cannot open shared object file: No such file or directory`,
you may need to install ffmpeg
```bash
apt-get update && apt-get install ffmpeg
```
## Usage
After installation, the following nodes will be available in the ComfyUI interface under the "fastvideo" category:
- **Video Generator**: The main node for generating videos from prompts
- **Inference Args**: Configure video generation parameters
- **VAE Config**
- **Text Encoder Config**
- **DIT Config**
- **Load Image Path**: Load images for potential conditioning
You may have noticed many arguments on the nodes have 'auto' as the default value. This is because FastVideo will automatically detect the best values for these parameters based on the model and the hardware. However, you can also manually configure these parameters to get the best performance for your specific use case. We plan on releasing more optimized workflow files for different models and hardware configurations in the future.
You can see what some of the default configurations are by looking at the FastVideo repo:
- [Wan2.1-I2V-14B-480P-Diffusers](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/configs/wan_14B_i2v_480p_pipeline.json)
- [FastHunyuan-diffusers](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/configs/fasthunyuan_t2v.json)
### Node Configuration
#### Video Generator
- **prompt**: Text description of the video to generate
- **output_path**: Directory where generated videos will be saved
- **num_gpus**: Number of GPUs to use for generation
- **model_path**: Path to the FastVideo model
- **embedded_cfg_scale**: Classifier-free guidance scale
- **sp_size**: Sequence parallelism size (usually should match num_gpus)
- **tp_size**: Tensor parallelism size (usually should match num_gpus)
- **precision**: Model precision (fp16 or bf16)
`model_path takes either a model id from huggingface or a local path to a model. Models by default will be downloaded to ~/.cache/huggingface/hub/ and cached for subsequent runs.`
#### Inference Args
- **height/width**: Resolution of the output video
- **num_frames**: Number of frames to generate
- **num_inference_steps**: Number of diffusion steps per frame
- **guidance_scale**: Classifier-free guidance scale
- **flow_shift**: Frame flow shift parameter
- **seed**: Random seed for reproducible generation
- **fps**: Frames per second of the output video
- **image_path**: Optional path to input image for conditioning (for i2v models)
## Memory Management
Models will remain loaded in GPU memory between runs when you only change inference arguments (such as prompt, resolution, frame count, FPS, guidance scale, etc.) or the prompt text. This allows for faster subsequent generations since the model doesn't need to be reloaded.
However, if you need to change the following parameters, you will need to restart the ComfyUI server:
- **Number of GPUs** (`num_gpus`)
- **Model path** (`model_path`)
- **Tensor parallelism size** (`tp_size`)
- **Sequence parallelism size** (`sp_size`)
These parameters affect the model's distribution across GPUs and require a complete reinitialization of the model pipeline.
## Example workflows
### Text to Video
FastVideo-FastHunyuan-diffusers
![FastVideo-FastHunyuan-diffusers](./assets/fasthunyuan.png)
- [FastHunyuan-diffusers.json](./examples/FastHunyuan-diffusers.json)
### Image to Video
Wan2.1-I2V-14B-480P-Diffusers
![Wan2.1-I2V-14B-480P-Diffusers](./assets/wani2v.png)
- [Wan2.1-I2V-14B-480P-Diffusers.json](./examples/Wan2.1-I2V-14B-480P-Diffusers.json)
## License
This project is licensed under Apache 2.0.
+5
View File
@@ -0,0 +1,5 @@
from .video_generator.nodes import (NODE_CLASS_MAPPINGS,
NODE_DISPLAY_NAME_MAPPINGS)
WEB_DIRECTORY = "./web"
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY']
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.3 MiB

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

After

Width:  |  Height:  |  Size: 691 B

Binary file not shown.

After

Width:  |  Height:  |  Size: 8.7 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 769 KiB

+645
View File
@@ -0,0 +1,645 @@
{
"id": "23a1f065-bbba-4a8f-b144-944e1318fcbf",
"revision": 0,
"last_node_id": 8,
"last_link_id": 7,
"nodes": [
{
"id": 4,
"type": "VAEConfig",
"pos": [
374.2159423828125,
554.85888671875
],
"size": [
334.080078125,
322
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "vae_config",
"type": "VAE_CONFIG",
"links": [
2
]
}
],
"properties": {
"Node name for S&R": "VAEConfig"
},
"widgets_values": [
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999
],
"auto_widget_states": {
"load_encoder": {
"isAuto": true,
"value": -99999,
"cachedValue": true
},
"load_decoder": {
"isAuto": true,
"value": -99999,
"cachedValue": true
},
"tile_sample_min_height": {
"isAuto": true,
"value": -99999,
"cachedValue": 256
},
"tile_sample_min_width": {
"isAuto": true,
"value": -99999,
"cachedValue": 256
},
"tile_sample_min_num_frames": {
"isAuto": true,
"value": -99999,
"cachedValue": 16
},
"tile_sample_stride_height": {
"isAuto": true,
"value": -99999,
"cachedValue": 192
},
"tile_sample_stride_width": {
"isAuto": true,
"value": -99999,
"cachedValue": 192
},
"tile_sample_stride_num_frames": {
"isAuto": true,
"value": -99999,
"cachedValue": 12
},
"blend_num_frames": {
"isAuto": true,
"value": -99999,
"cachedValue": 0
},
"use_tiling": {
"isAuto": true,
"value": -99999,
"cachedValue": true
},
"use_temporal_tiling": {
"isAuto": true,
"value": -99999,
"cachedValue": true
},
"use_parallel_tiling": {
"isAuto": true,
"value": -99999,
"cachedValue": true
}
}
},
{
"id": 5,
"type": "TextEncoderConfig",
"pos": [
416.4937744140625,
953.6171875
],
"size": [
270,
106
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "text_encoder_config",
"type": "TEXT_ENCODER_CONFIG",
"links": [
7
]
}
],
"properties": {
"Node name for S&R": "TextEncoderConfig"
},
"widgets_values": [
-99999,
-99999,
-99999
],
"auto_widget_states": {
"prefix": {
"isAuto": true,
"value": -99999,
"cachedValue": ""
},
"quant_config": {
"isAuto": true,
"value": -99999,
"cachedValue": ""
},
"lora_config": {
"isAuto": true,
"value": -99999,
"cachedValue": ""
}
}
},
{
"id": 6,
"type": "DITConfig",
"pos": [
415.1928405761719,
1154.1573486328125
],
"size": [
270,
82
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "dit_config",
"type": "DIT_CONFIG",
"links": [
6
]
}
],
"properties": {
"Node name for S&R": "DITConfig"
},
"widgets_values": [
-99999,
-99999
],
"auto_widget_states": {
"prefix": {
"isAuto": true,
"value": -99999,
"cachedValue": ""
},
"quant_config": {
"isAuto": true,
"value": -99999,
"cachedValue": ""
}
}
},
{
"id": 1,
"type": "VideoGenerator",
"pos": [
818.804931640625,
348.9299621582031
],
"size": [
400,
436
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "inference_args",
"shape": 7,
"type": "INFERENCE_ARGS",
"link": 3
},
{
"name": "vae_config",
"shape": 7,
"type": "VAE_CONFIG",
"link": 2
},
{
"name": "text_encoder_config",
"shape": 7,
"type": "TEXT_ENCODER_CONFIG",
"link": 7
},
{
"name": "dit_config",
"shape": 7,
"type": "DIT_CONFIG",
"link": 6
}
],
"outputs": [
{
"name": "video_path",
"type": "STRING",
"links": [
4
]
}
],
"properties": {
"Node name for S&R": "VideoGenerator"
},
"widgets_values": [
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest. The playful yet serene atmosphere is complemented by soft natural light filtering through the petals. Mid-shot, warm and cheerful tones.",
"/workspace/ComfyUI/outputs_video/",
2,
"FastVideo/FastHunyuan-diffusers",
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999,
-99999
],
"auto_widget_states": {
"embedded_cfg_scale": {
"isAuto": true,
"value": -99999,
"cachedValue": 6
},
"sp_size": {
"isAuto": true,
"value": -99999,
"cachedValue": 2
},
"tp_size": {
"isAuto": true,
"value": -99999,
"cachedValue": 2
},
"vae_precision": {
"isAuto": true,
"value": -99999,
"cachedValue": "fp16"
},
"vae_tiling": {
"isAuto": true,
"value": -99999,
"cachedValue": true
},
"vae_sp": {
"isAuto": true,
"value": -99999,
"cachedValue": true
},
"text_encoder_precision": {
"isAuto": true,
"value": -99999,
"cachedValue": "fp16"
},
"precision": {
"isAuto": true,
"value": -99999,
"cachedValue": "fp16"
},
"dit_cpu_offload": {
"isAuto": true,
"value": -99999,
"cachedValue": true
}
}
},
{
"id": 3,
"type": "VHS_LoadVideoPath",
"pos": [
1350.136962890625,
331.20361328125
],
"size": [
231.8896484375,
286
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "meta_batch",
"shape": 7,
"type": "VHS_BatchManager",
"link": null
},
{
"name": "vae",
"shape": 7,
"type": "VAE",
"link": null
},
{
"name": "video",
"type": "STRING",
"widget": {
"name": "video"
},
"link": 4
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
5
]
},
{
"name": "frame_count",
"type": "INT",
"links": null
},
{
"name": "audio",
"type": "AUDIO",
"links": null
},
{
"name": "video_info",
"type": "VHS_VIDEOINFO",
"links": null
}
],
"properties": {
"Node name for S&R": "VHS_LoadVideoPath"
},
"widgets_values": {
"video": "",
"force_rate": 0,
"custom_width": 0,
"custom_height": 0,
"frame_load_cap": 0,
"skip_first_frames": 0,
"select_every_nth": 1,
"format": "Wan",
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "",
"type": "path",
"format": "video/",
"force_rate": 0,
"custom_width": 0,
"custom_height": 0,
"frame_load_cap": 0,
"skip_first_frames": 0,
"select_every_nth": 1
}
}
}
},
{
"id": 2,
"type": "InferenceArgs",
"pos": [
411.46307373046875,
178.18182373046875
],
"size": [
278.73828125,
298
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "inference_args",
"type": "INFERENCE_ARGS",
"links": [
3
]
}
],
"properties": {
"Node name for S&R": "InferenceArgs"
},
"widgets_values": [
720,
1280,
45,
6,
-99999,
-99999,
1025,
"fixed",
24,
-99999,
-99999
],
"auto_widget_states": {
"height": {
"isAuto": false,
"value": 720,
"cachedValue": 720
},
"width": {
"isAuto": false,
"value": 1280,
"cachedValue": 1280
},
"num_frames": {
"isAuto": false,
"value": 45,
"cachedValue": 45
},
"num_inference_steps": {
"isAuto": false,
"value": 6,
"cachedValue": 6
},
"guidance_scale": {
"isAuto": true,
"value": -99999,
"cachedValue": 1
},
"flow_shift": {
"isAuto": true,
"value": -99999,
"cachedValue": 17
},
"seed": {
"isAuto": false,
"value": 1025,
"cachedValue": 1024
},
"fps": {
"isAuto": false,
"value": 24,
"cachedValue": 24
},
"image_path": {
"isAuto": true,
"value": -99999,
"cachedValue": "X://insert/path/here.mp4"
},
"enable_teacache": {
"isAuto": true,
"value": -99999,
"cachedValue": true
}
}
},
{
"id": 8,
"type": "VHS_VideoCombine",
"pos": [
1668.3499755859375,
328.22625732421875
],
"size": [
507.507080078125,
622.2227172851562
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 5
},
{
"name": "audio",
"shape": 7,
"type": "AUDIO",
"link": null
},
{
"name": "meta_batch",
"shape": 7,
"type": "VHS_BatchManager",
"link": null
},
{
"name": "vae",
"shape": 7,
"type": "VAE",
"link": null
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 24,
"loop_count": 0,
"filename_prefix": "",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 19,
"save_metadata": true,
"trim_to_audio": false,
"pingpong": false,
"save_output": false,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "._00003.mp4",
"subfolder": "",
"type": "temp",
"format": "video/h264-mp4",
"frame_rate": 24,
"workflow": "._00003.png",
"fullpath": "/workspace/ComfyUI/temp/._00003.mp4"
}
}
}
}
],
"links": [
[
2,
4,
0,
1,
1,
"VAE_CONFIG"
],
[
3,
2,
0,
1,
0,
"INFERENCE_ARGS"
],
[
4,
1,
0,
3,
2,
"STRING"
],
[
5,
3,
0,
8,
0,
"IMAGE"
],
[
6,
6,
0,
1,
3,
"DIT_CONFIG"
],
[
7,
5,
0,
1,
2,
"TEXT_ENCODER_CONFIG"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.9090909090909091,
"offset": [
112.86678372727341,
-71.45635903989245
]
},
"frontendVersion": "1.20.4",
"VHS_latentpreview": false,
"VHS_latentpreviewrate": 0,
"VHS_MetadataImage": true,
"VHS_KeepIntermediate": true
},
"version": 0.4
}
@@ -0,0 +1,697 @@
{
"id": "23a1f065-bbba-4a8f-b144-944e1318fcbf",
"revision": 0,
"last_node_id": 8,
"last_link_id": 7,
"nodes": [
{
"id": 7,
"type": "LoadImagePath",
"pos": [
33.15385437011719,
191.2037353515625
],
"size": [
270,
334
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "image_path",
"type": "STRING",
"links": [
1
]
},
{
"name": "IMAGE",
"type": "IMAGE",
"links": null
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"Node name for S&R": "LoadImagePath"
},
"widgets_values": [
"woman.jpg",
"image"
]
},
{
"id": 3,
"type": "VHS_LoadVideoPath",
"pos": [
1350.136962890625,
331.20361328125
],
"size": [
231.8896484375,
286
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "meta_batch",
"shape": 7,
"type": "VHS_BatchManager",
"link": null
},
{
"name": "vae",
"shape": 7,
"type": "VAE",
"link": null
},
{
"name": "video",
"type": "STRING",
"widget": {
"name": "video"
},
"link": 4
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
5
]
},
{
"name": "frame_count",
"type": "INT",
"links": null
},
{
"name": "audio",
"type": "AUDIO",
"links": null
},
{
"name": "video_info",
"type": "VHS_VIDEOINFO",
"links": null
}
],
"properties": {
"Node name for S&R": "VHS_LoadVideoPath"
},
"widgets_values": {
"video": "",
"force_rate": 0,
"custom_width": 0,
"custom_height": 0,
"frame_load_cap": 0,
"skip_first_frames": 0,
"select_every_nth": 1,
"format": "Wan",
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "",
"type": "path",
"format": "video/",
"force_rate": 0,
"custom_width": 0,
"custom_height": 0,
"frame_load_cap": 0,
"skip_first_frames": 0,
"select_every_nth": 1
}
}
}
},
{
"id": 8,
"type": "VHS_VideoCombine",
"pos": [
1668.3499755859375,
328.22625732421875
],
"size": [
214.7587890625,
334
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 5
},
{
"name": "audio",
"shape": 7,
"type": "AUDIO",
"link": null
},
{
"name": "meta_batch",
"shape": 7,
"type": "VHS_BatchManager",
"link": null
},
{
"name": "vae",
"shape": 7,
"type": "VAE",
"link": null
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 24,
"loop_count": 0,
"filename_prefix": "",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 19,
"save_metadata": true,
"trim_to_audio": false,
"pingpong": false,
"save_output": false,
"videopreview": {
"hidden": false,
"paused": false,
"params": {}
}
}
},
{
"id": 4,
"type": "VAEConfig",
"pos": [
374.2159423828125,
554.85888671875
],
"size": [
334.080078125,
322
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "vae_config",
"type": "VAE_CONFIG",
"links": [
2
]
}
],
"properties": {
"Node name for S&R": "VAEConfig"
},
"widgets_values": [
true,
true,
256,
256,
16,
192,
192,
12,
0,
true,
true,
true
],
"auto_widget_states": {
"load_encoder": {
"isAuto": true,
"value": true,
"cachedValue": true
},
"load_decoder": {
"isAuto": true,
"value": true,
"cachedValue": true
},
"tile_sample_min_height": {
"isAuto": true,
"value": 256,
"cachedValue": 256
},
"tile_sample_min_width": {
"isAuto": true,
"value": 256,
"cachedValue": 256
},
"tile_sample_min_num_frames": {
"isAuto": true,
"value": 16,
"cachedValue": 16
},
"tile_sample_stride_height": {
"isAuto": true,
"value": 192,
"cachedValue": 192
},
"tile_sample_stride_width": {
"isAuto": true,
"value": 192,
"cachedValue": 192
},
"tile_sample_stride_num_frames": {
"isAuto": true,
"value": 12,
"cachedValue": 12
},
"blend_num_frames": {
"isAuto": true,
"value": 0,
"cachedValue": 0
},
"use_tiling": {
"isAuto": true,
"value": true,
"cachedValue": true
},
"use_temporal_tiling": {
"isAuto": true,
"value": true,
"cachedValue": true
},
"use_parallel_tiling": {
"isAuto": true,
"value": true,
"cachedValue": true
}
}
},
{
"id": 2,
"type": "InferenceArgs",
"pos": [
411.46307373046875,
178.18182373046875
],
"size": [
278.73828125,
298
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "image_path",
"shape": 7,
"type": "STRING",
"widget": {
"name": "image_path"
},
"link": 1
}
],
"outputs": [
{
"name": "inference_args",
"type": "INFERENCE_ARGS",
"links": [
3
]
}
],
"properties": {
"Node name for S&R": "InferenceArgs"
},
"widgets_values": [
832,
480,
45,
20,
1,
17,
1024,
"fixed",
24,
"X://insert/path/here.mp4",
true
],
"auto_widget_states": {
"height": {
"isAuto": false,
"value": 832,
"cachedValue": 720
},
"width": {
"isAuto": false,
"value": 480,
"cachedValue": 1280
},
"num_frames": {
"isAuto": false,
"value": 45,
"cachedValue": 45
},
"num_inference_steps": {
"isAuto": false,
"value": 20,
"cachedValue": 6
},
"guidance_scale": {
"isAuto": true,
"value": 1,
"cachedValue": 1
},
"flow_shift": {
"isAuto": true,
"value": 17,
"cachedValue": 17
},
"seed": {
"isAuto": false,
"value": 1024,
"cachedValue": 1024
},
"fps": {
"isAuto": false,
"value": 24,
"cachedValue": 24
},
"image_path": {
"isAuto": true,
"value": "X://insert/path/here.mp4",
"cachedValue": "X://insert/path/here.mp4"
},
"enable_teacache": {
"isAuto": true,
"value": true,
"cachedValue": true
}
}
},
{
"id": 1,
"type": "VideoGenerator",
"pos": [
818.804931640625,
348.9299621582031
],
"size": [
400,
436
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "inference_args",
"shape": 7,
"type": "INFERENCE_ARGS",
"link": 3
},
{
"name": "vae_config",
"shape": 7,
"type": "VAE_CONFIG",
"link": 2
},
{
"name": "text_encoder_config",
"shape": 7,
"type": "TEXT_ENCODER_CONFIG",
"link": 7
},
{
"name": "dit_config",
"shape": 7,
"type": "DIT_CONFIG",
"link": 6
}
],
"outputs": [
{
"name": "video_path",
"type": "STRING",
"links": [
4
]
}
],
"properties": {
"Node name for S&R": "VideoGenerator"
},
"widgets_values": [
"A woman crying from laughter.",
"/workspace/ComfyUI/outputs_video/",
4,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
6,
2,
2,
"fp16",
true,
true,
"fp16",
"fp16",
true
],
"auto_widget_states": {
"embedded_cfg_scale": {
"isAuto": true,
"value": 6,
"cachedValue": 6
},
"sp_size": {
"isAuto": true,
"value": 2,
"cachedValue": 2
},
"tp_size": {
"isAuto": true,
"value": 2,
"cachedValue": 2
},
"vae_precision": {
"isAuto": true,
"value": "fp16",
"cachedValue": "fp16"
},
"vae_tiling": {
"isAuto": true,
"value": true,
"cachedValue": true
},
"vae_sp": {
"isAuto": true,
"value": true,
"cachedValue": true
},
"text_encoder_precision": {
"isAuto": true,
"value": "fp16",
"cachedValue": "fp16"
},
"precision": {
"isAuto": true,
"value": "fp16",
"cachedValue": "fp16"
},
"dit_cpu_offload": {
"isAuto": true,
"value": true,
"cachedValue": true
}
}
},
{
"id": 5,
"type": "TextEncoderConfig",
"pos": [
416.4937744140625,
953.6171875
],
"size": [
270,
106
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "text_encoder_config",
"type": "TEXT_ENCODER_CONFIG",
"links": [
7
]
}
],
"properties": {
"Node name for S&R": "TextEncoderConfig"
},
"widgets_values": [
"",
"",
""
],
"auto_widget_states": {
"prefix": {
"isAuto": true,
"value": "",
"cachedValue": ""
},
"quant_config": {
"isAuto": true,
"value": "",
"cachedValue": ""
},
"lora_config": {
"isAuto": true,
"value": "",
"cachedValue": ""
}
}
},
{
"id": 6,
"type": "DITConfig",
"pos": [
415.1928405761719,
1154.1573486328125
],
"size": [
270,
82
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "dit_config",
"type": "DIT_CONFIG",
"links": [
6
]
}
],
"properties": {
"Node name for S&R": "DITConfig"
},
"widgets_values": [
"",
""
],
"auto_widget_states": {
"prefix": {
"isAuto": true,
"value": "",
"cachedValue": ""
},
"quant_config": {
"isAuto": true,
"value": "",
"cachedValue": ""
}
}
}
],
"links": [
[
1,
7,
0,
2,
0,
"STRING"
],
[
2,
4,
0,
1,
1,
"VAE_CONFIG"
],
[
3,
2,
0,
1,
0,
"INFERENCE_ARGS"
],
[
4,
1,
0,
3,
2,
"STRING"
],
[
5,
3,
0,
8,
0,
"IMAGE"
],
[
6,
6,
0,
1,
3,
"DIT_CONFIG"
],
[
7,
5,
0,
1,
2,
"TEXT_ENCODER_CONFIG"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.8264462809917354,
"offset": [
646.7950212991898,
66.17259910028655
]
},
"frontendVersion": "1.20.4",
"VHS_latentpreview": false,
"VHS_latentpreviewrate": 0,
"VHS_MetadataImage": true,
"VHS_KeepIntermediate": true
},
"version": 0.4
}
+31
View File
@@ -0,0 +1,31 @@
class DITConfig:
@classmethod
def INPUT_TYPES(cls):
return {
"optional": {
"prefix": ("STRING", {
"default": ""
}),
"quant_config": ("STRING", {
"default": ""
}),
}
}
@classmethod
def VALIDATE_INPUTS(cls, **kwargs):
return True
RETURN_TYPES = ("DIT_CONFIG", )
RETURN_NAMES = ("dit_config", )
FUNCTION = "set_args"
CATEGORY = "fastvideo"
def set_args(self, prefix, quant_config):
raw_args = {"prefix": prefix, "quant_config": quant_config}
# Filter out keys where value is -99999
args = {k: v for k, v in raw_args.items() if str(int(v)) != str(-99999)}
return (args, )
+89
View File
@@ -0,0 +1,89 @@
class InferenceArgs:
@classmethod
def INPUT_TYPES(cls):
return {
"optional": {
"height": ("INT", {
"default": 720
}),
"width": ("INT", {
"default": 1280
}),
"num_frames": ("INT", {
"default": 45
}),
"num_inference_steps": ("INT", {
"default": 6
}),
"guidance_scale": ("FLOAT", {
"default": 1.0
}),
"flow_shift": ("INT", {
"default": 17
}),
"seed": ("INT", {
"default": 1024
}),
"fps": ("INT", {
"default": 24
}),
"image_path": ("STRING", {
"default": "X://insert/path/here.mp4"
}),
"enable_teacache": ([True, False], {
"default": False
}),
}
}
@classmethod
def VALIDATE_INPUTS(cls, **kwargs):
return True
RETURN_TYPES = ("INFERENCE_ARGS", )
RETURN_NAMES = ("inference_args", )
FUNCTION = "set_args"
CATEGORY = "fastvideo"
def set_args(
self,
height,
width,
num_frames,
num_inference_steps,
guidance_scale,
flow_shift,
seed,
fps,
image_path,
enable_teacache,
):
raw_args = {
"height": height,
"width": width,
"num_frames": num_frames,
"num_inference_steps": num_inference_steps,
"guidance_scale": guidance_scale,
"flow_shift": flow_shift,
"seed": seed,
"fps": fps,
"image_path": image_path,
"enable_teacache": enable_teacache,
}
# Filter out keys where value is -99999, handling different types properly
args = {}
for k, v in raw_args.items():
try:
if isinstance(v, str):
if v != "-99999":
args[k] = v
elif v != -99999:
# If it's not a string, compare directly
args[k] = v
except (ValueError, TypeError):
# Include any value that causes an error in comparison
args[k] = v
return (args, )
+103
View File
@@ -0,0 +1,103 @@
import hashlib
import os
import folder_paths
import numpy as np
import torch
from PIL import Image, ImageOps, ImageSequence
from .node_helpers import pillow
class LoadImagePath:
@classmethod
def INPUT_TYPES(s):
input_dir = folder_paths.get_input_directory()
files = [
f for f in os.listdir(input_dir)
if os.path.isfile(os.path.join(input_dir, f))
]
files = folder_paths.filter_files_content_types(files, ["image"])
return {
"required": {
"image": (sorted(files), {
"image_upload": True
})
},
}
CATEGORY = "fastvideo"
RETURN_TYPES = ("STRING", "IMAGE", "MASK")
RETURN_NAMES = ("image_path", "IMAGE", "MASK")
FUNCTION = "load_image"
def load_image(self, image):
image_path = folder_paths.get_annotated_filepath(image)
img = pillow(Image.open, image_path)
output_images: list[torch.Tensor] = []
output_masks: list[torch.Tensor] = []
w, h = None, None
excluded_formats = ['MPO']
for i in ImageSequence.Iterator(img):
processed_image = pillow(ImageOps.exif_transpose, i)
if processed_image is None:
continue
if processed_image.mode == 'I':
processed_image = processed_image.point(lambda i: i * (1 / 255))
image = processed_image.convert("RGB")
if len(output_images) == 0:
w = image.size[0]
h = image.size[1]
if image.size[0] != w or image.size[1] != h:
continue
image = np.array(image).astype(np.float32) / 255.0
image = torch.from_numpy(image)[
None,
]
if 'A' in processed_image.getbands():
mask = np.array(processed_image.getchannel('A')).astype(
np.float32) / 255.0
mask = 1. - torch.from_numpy(mask)
elif processed_image.mode == 'P' and 'transparency' in processed_image.info:
mask = np.array(
processed_image.convert('RGBA').getchannel('A')).astype(
np.float32) / 255.0
mask = 1. - torch.from_numpy(mask)
else:
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
output_images.append(image)
output_masks.append(mask.unsqueeze(0))
if len(output_images) > 1 and img.format not in excluded_formats:
output_image = torch.cat(output_images, dim=0)
output_mask = torch.cat(output_masks, dim=0)
else:
output_image = output_images[0]
output_mask = output_masks[0]
return (image_path, output_image, output_mask)
@classmethod
def IS_CHANGED(s, image):
image_path = folder_paths.get_annotated_filepath(image)
m = hashlib.sha256()
with open(image_path, 'rb') as f:
m.update(f.read())
return m.digest().hex()
@classmethod
def VALIDATE_INPUTS(s, image):
if not folder_paths.exists_annotated_filepath(image):
return "Invalid image file: {}".format(image)
return True
+68
View File
@@ -0,0 +1,68 @@
import hashlib
from collections.abc import Callable
from typing import Any, TypeVar
import torch
from comfy.cli_args import args
from PIL import ImageFile, UnidentifiedImageError
T = TypeVar('T')
def conditioning_set_values(conditioning: list[Any],
values: dict[str, Any] | None = None) -> list[Any]:
if values is None:
values = {}
c = []
for t in conditioning:
n = [t[0], t[1].copy()]
for k in values:
n[1][k] = values[k]
c.append(n)
return c
def pillow(fn: Callable[[Any], T], arg: Any) -> T:
prev_value = None
try:
x = fn(arg)
except (OSError, UnidentifiedImageError, ValueError
): #PIL issues #4472 and #2445, also fixes ComfyUI issue #3416
prev_value = ImageFile.LOAD_TRUNCATED_IMAGES
ImageFile.LOAD_TRUNCATED_IMAGES = True
x = fn(arg)
finally:
if prev_value is not None:
ImageFile.LOAD_TRUNCATED_IMAGES = prev_value
return x
def hasher() -> Callable[[], Any]:
hashfuncs = {
"md5": hashlib.md5,
"sha1": hashlib.sha1,
"sha256": hashlib.sha256,
"sha512": hashlib.sha512
}
return hashfuncs[args.default_hashing_function]
def string_to_torch_dtype(string: str) -> torch.dtype | None:
if string == "fp32":
return torch.float32
if string == "fp16":
return torch.float16
if string == "bf16":
return torch.bfloat16
return None
def image_alpha_fix(destination: torch.Tensor,
source: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
if destination.shape[-1] < source.shape[-1]:
source = source[..., :destination.shape[-1]]
elif destination.shape[-1] > source.shape[-1]:
destination = torch.nn.functional.pad(destination, (0, 1))
destination[..., -1] = 1.0
return destination, source
+24
View File
@@ -0,0 +1,24 @@
from .dit_config import DITConfig
from .inference_args import InferenceArgs
from .load_image import LoadImagePath
from .text_encoder_config import TextEncoderConfig
from .vae_config import VAEConfig
from .video_generator import VideoGenerator
NODE_CLASS_MAPPINGS = {
"VideoGenerator": VideoGenerator,
"InferenceArgs": InferenceArgs,
"VAEConfig": VAEConfig,
"TextEncoderConfig": TextEncoderConfig,
"DITConfig": DITConfig,
"LoadImagePath": LoadImagePath
}
NODE_DISPLAY_NAME_MAPPINGS = {
"VideoGenerator": "Video Generator",
"InferenceArgs": "Inference Args",
"VAEConfig": "VAE Config",
"TextEncoderConfig": "Text Encoder Config",
"DITConfig": "DIT Config",
"LoadImagePath": "Load Image Path"
}
@@ -0,0 +1,38 @@
class TextEncoderConfig:
@classmethod
def INPUT_TYPES(cls):
return {
"optional": {
"prefix": ("STRING", {
"default": ""
}),
"quant_config": ("STRING", {
"default": ""
}),
"lora_config": ("STRING", {
"default": ""
}),
}
}
@classmethod
def VALIDATE_INPUTS(cls, **kwargs):
return True
RETURN_TYPES = ("TEXT_ENCODER_CONFIG", )
RETURN_NAMES = ("text_encoder_config", )
FUNCTION = "set_args"
CATEGORY = "fastvideo"
def set_args(self, prefix, quant_config, lora_config):
raw_args = {
"prefix": prefix,
"quant_config": quant_config,
"lora_config": lora_config
}
# Filter out keys where value is -99999
args = {k: v for k, v in raw_args.items() if str(int(v)) != str(-99999)}
return (args, )
+88
View File
@@ -0,0 +1,88 @@
class VAEConfig:
@classmethod
def INPUT_TYPES(cls):
return {
"optional": {
"load_encoder": ([True, False], {
"default": True
}),
"load_decoder": ([True, False], {
"default": True
}),
"tile_sample_min_height": ("INT", {
"default": 256
}),
"tile_sample_min_width": ("INT", {
"default": 256
}),
"tile_sample_min_num_frames": ("INT", {
"default": 16
}),
"tile_sample_stride_height": ("INT", {
"default": 192
}),
"tile_sample_stride_width": ("INT", {
"default": 192
}),
"tile_sample_stride_num_frames": ("INT", {
"default": 12
}),
"blend_num_frames": ("INT", {
"default": 0
}),
"use_tiling": ([True, False], {
"default": True
}),
"use_temporal_tiling": ([True, False], {
"default": True
}),
"use_parallel_tiling": ([True, False], {
"default": True
}),
}
}
@classmethod
def VALIDATE_INPUTS(cls, **kwargs):
return True
RETURN_TYPES = ("VAE_CONFIG", )
RETURN_NAMES = ("vae_config", )
FUNCTION = "set_args"
CATEGORY = "fastvideo"
def set_args(
self,
load_encoder,
load_decoder,
tile_sample_min_height,
tile_sample_min_width,
tile_sample_min_num_frames,
tile_sample_stride_height,
tile_sample_stride_width,
tile_sample_stride_num_frames,
blend_num_frames,
use_tiling,
use_temporal_tiling,
use_parallel_tiling,
):
raw_args = {
"load_encoder": load_encoder,
"load_decoder": load_decoder,
"tile_sample_min_height": tile_sample_min_height,
"tile_sample_min_width": tile_sample_min_width,
"tile_sample_min_num_frames": tile_sample_min_num_frames,
"tile_sample_stride_height": tile_sample_stride_height,
"tile_sample_stride_width": tile_sample_stride_width,
"tile_sample_stride_num_frames": tile_sample_stride_num_frames,
"blend_num_frames": blend_num_frames,
"use_tiling": use_tiling,
"use_temporal_tiling": use_temporal_tiling,
"use_parallel_tiling": use_parallel_tiling,
}
# Filter out any value explicitly set to -99999
args = {k: v for k, v in raw_args.items() if str(int(v)) != str(-99999)}
return (args, )
+315
View File
@@ -0,0 +1,315 @@
from __future__ import annotations
import glob
import os
import signal
import sys
import threading
import time
from typing import Any
from comfy.model_management import processing_interrupted
from fastvideo import PipelineConfig
from fastvideo import VideoGenerator as FastVideoGenerator
sys.path.insert(
0,
os.path.dirname(
os.path.dirname(
os.path.dirname(os.path.dirname(os.path.abspath(__file__))))))
# Custom exception for interruption
class GenerationInterruptedException(Exception):
pass
# Custom exception for interruption that ComfyUI will recognize
class GenerationCancelledException(Exception):
def __init__(self,
message: str = "Generation was cancelled by user") -> None:
self.message = message
super().__init__(self.message)
def update_config_from_args(config: Any, args_dict: dict[str, Any]) -> None:
"""
Update configuration object from arguments dictionary.
Args:
config: The configuration object to update
args_dict: Dictionary containing arguments
"""
for key, value in args_dict.items():
if hasattr(config, key) and value is not None:
if key == "text_encoder_precisions" and isinstance(value, list):
setattr(config, key, tuple(value))
else:
setattr(config, key, value)
class VideoGenerator:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"prompt": ("STRING", {
"multiline":
True,
"default":
"A ripe orange tumbles gently from a tree and lands on the head of a lounging capybara, "
"who blinks slowly in response. The moment is quietly humorous and oddly serene, framed by "
"lush green foliage and dappled sunlight. Mid-shot, warm and whimsical tones."
}),
"output_path": ("STRING", {
"default": "/workspace/ComfyUI/outputs_video/"
}),
"num_gpus": ("INT", {
"default": 2,
"min": 1,
"max": 16
}),
"model_path": ("STRING", {
"default": "FastVideo/FastHunyuan-diffusers"
})
},
"optional": {
"inference_args": ("INFERENCE_ARGS", ),
"embedded_cfg_scale": ("FLOAT", {
"default": 6.0
}),
"sp_size": ("INT", {
"default": 2
}),
"tp_size": ("INT", {
"default": 2
}),
"vae_config": ("VAE_CONFIG", ),
"vae_precision": (["fp16", "bf16"], {
"default": "fp16"
}),
"vae_tiling": ([True, False], {
"default": True
}),
"vae_sp": ([True, False], {
"default": False
}),
"text_encoder_config": ("TEXT_ENCODER_CONFIG", ),
"text_encoder_precision": (["fp16", "bf16"], {
"default": "fp16"
}),
"dit_config": ("DIT_CONFIG", ),
"precision": (["fp16", "bf16"], {
"default": "fp16"
}),
"dit_cpu_offload": ([True, False], {
"default": False
}),
}
}
@classmethod
def VALIDATE_INPUTS(cls, **kwargs):
return True
RETURN_TYPES = ("STRING", )
RETURN_NAMES = ("video_path", )
FUNCTION = "launch_inference"
CATEGORY = "fastvideo"
generator: FastVideoGenerator | None = None
_interrupt_thread: threading.Thread | None = None
_generation_active: bool = False
_generation_interrupted: bool = False
_interrupt_event: threading.Event = threading.Event()
_generation_thread: threading.Thread | None = None
_generation_result: str | None = None
_generation_exception: Exception | None = None
def _monitor_for_interruption(self):
"""Background thread that monitors for interruption requests"""
time.sleep(2) # Give the generation thread time to send execute_forward
while self._generation_active and not self._interrupt_event.is_set():
if processing_interrupted():
print("Video generation interrupted by user")
self._generation_interrupted = True
# Try to send interrupt signal to worker processes
if self.generator is not None and hasattr(
self.generator, 'executor'):
try:
# The MultiprocExecutor has a workers attribute
if hasattr(self.generator.executor, 'workers'):
for worker in self.generator.executor.workers:
if worker.is_alive():
os.kill(worker.pid, signal.SIGINT)
print("Interrupt signal sent to worker processes")
except Exception as e:
print(f"Error sending interrupt signal: {e}")
# Set the interrupt event to notify other threads
self._interrupt_event.set()
break
time.sleep(0.5)
def _run_generation(self, prompt: str, output_path: str,
inference_args: dict[str, Any]) -> None:
"""Thread function to run the generation"""
try:
if self.generator is not None:
self.generator.generate_video(prompt=prompt,
output_path=output_path,
**inference_args)
self._generation_result = os.path.join(output_path,
f"{prompt[:100]}.mp4")
else:
raise RuntimeError("Generator is not initialized")
except Exception as e:
self._generation_exception = e
self._interrupt_event.set()
def load_output_video(self, output_dir):
video_extensions = ["*.mp4", "*.avi", "*.mov", "*.mkv"]
video_files = []
for ext in video_extensions:
video_files.extend(glob.glob(os.path.join(output_dir, ext)))
if not video_files:
print("No video files found in output directory: %s", output_dir)
return ""
video_files.sort()
return video_files[0]
def launch_inference(
self,
prompt,
output_path,
num_gpus,
model_path,
embedded_cfg_scale,
sp_size,
tp_size,
vae_precision,
vae_tiling,
vae_sp,
text_encoder_precision,
precision,
inference_args=None,
vae_config=None,
text_encoder_config=None,
dit_config=None,
dit_cpu_offload=None,
):
print('Running FastVideo inference')
# Reset interruption flag and event
self._generation_interrupted = False
self._interrupt_event.clear()
self._generation_result = None
self._generation_exception = None
# Load pipeline config from model path
pipeline_config = PipelineConfig.from_pretrained(model_path)
print('pipeline_config', pipeline_config)
# Update configs with provided config dictionaries
if dit_config is not None:
update_config_from_args(pipeline_config.dit_config, dit_config)
if vae_config is not None:
update_config_from_args(pipeline_config.vae_config, vae_config)
if text_encoder_config is not None:
update_config_from_args(pipeline_config.text_encoder_configs,
text_encoder_config)
# Update top-level pipeline config with remaining arguments
raw_pipeline_args = {}
if embedded_cfg_scale is not None:
raw_pipeline_args['embedded_cfg_scale'] = embedded_cfg_scale
if precision is not None:
raw_pipeline_args['precision'] = precision
if vae_precision is not None:
raw_pipeline_args['vae_precision'] = vae_precision
if vae_tiling is not None:
raw_pipeline_args['vae_tiling'] = vae_tiling
if vae_sp is not None:
raw_pipeline_args['vae_sp'] = vae_sp
if text_encoder_precision is not None:
raw_pipeline_args['text_encoder_precision'] = text_encoder_precision
# Filter out any value explicitly set to -99999 (auto values)
pipeline_args = {
k: v
for k, v in raw_pipeline_args.items() if str(int(v)) != str(-99999)
}
update_config_from_args(pipeline_config, pipeline_args)
raw_generation_args = {}
if num_gpus is not None:
raw_generation_args['num_gpus'] = num_gpus
if tp_size is not None:
raw_generation_args['tp_size'] = tp_size
if sp_size is not None:
raw_generation_args['sp_size'] = sp_size
if dit_cpu_offload is not None:
raw_generation_args['dit_cpu_offload'] = dit_cpu_offload
generation_args = {
k: v
for k, v in raw_generation_args.items()
if str(int(v)) != str(-99999)
}
if self.generator is None:
print('generation_args', generation_args)
print('pipeline_config', pipeline_config)
self.generator = FastVideoGenerator.from_pretrained(
model_path=model_path,
**generation_args,
pipeline_config=pipeline_config)
print('inference_args', inference_args)
# Start a thread to run the generation
self._generation_thread = threading.Thread(target=self._run_generation,
args=(prompt, output_path,
inference_args),
daemon=True)
self._generation_thread.start()
# Start a background thread to monitor for interruptions
self._generation_active = True
self._interrupt_thread = threading.Thread(
target=self._monitor_for_interruption, daemon=True)
self._interrupt_thread.start()
# Wait for either completion or interruption
while self._generation_thread.is_alive(
) and not self._interrupt_event.is_set():
self._generation_thread.join(timeout=0.5)
self._generation_active = False
if self._interrupt_thread:
self._interrupt_thread.join(timeout=1.0)
self._interrupt_thread = None
if self._generation_interrupted:
print("Video generation was cancelled by user")
raise GenerationCancelledException()
elif self._generation_exception:
# Re-raise the exception from the generation thread
raise self._generation_exception
elif self._generation_result:
return (self._generation_result, )
else:
# This shouldn't happen, but just in case
print("Generation completed but no result was produced")
raise Exception("Generation failed to produce a result")
+593
View File
@@ -0,0 +1,593 @@
import { app } from '../../../scripts/app.js'
function chainCallback(object, property, callback) {
if (object == undefined) {
console.error("Tried to add callback to non-existent object");
return;
}
if (property in object && object[property]) {
const callback_orig = object[property];
object[property] = function () {
const r = callback_orig.apply(this, arguments);
return callback.apply(this, arguments) ?? r;
};
} else {
object[property] = callback;
}
}
function drawAutoAnnotated(ctx, node, widget_width, y, H) {
const litegraph_base = LiteGraph;
const show_text = app.canvas.ds.scale >= 0.5;
const margin = 15;
const autoTextWidth = 30;
const autoTextRightMargin = 5;
ctx.textAlign = 'left';
ctx.strokeStyle = litegraph_base.WIDGET_OUTLINE_COLOR;
ctx.fillStyle = litegraph_base.WIDGET_BGCOLOR;
ctx.beginPath();
if (show_text && ctx.roundRect) {
ctx.roundRect(margin, y, widget_width - margin * 2, H, [H * 0.5]);
} else {
ctx.rect(margin, y, widget_width - margin * 2, H);
}
ctx.fill();
if (show_text) {
if (!this.disabled) ctx.stroke();
const isAuto = this.isAuto === true;
ctx.save();
if (isAuto) {
ctx.fillStyle = litegraph_base.WIDGET_TEXT_COLOR;
ctx.strokeStyle = litegraph_base.WIDGET_TEXT_COLOR;
} else {
ctx.fillStyle = litegraph_base.WIDGET_SECONDARY_TEXT_COLOR;
ctx.strokeStyle = litegraph_base.WIDGET_SECONDARY_TEXT_COLOR;
}
// Position for the cog
const cogX = widget_width - autoTextRightMargin - autoTextWidth - 6;
const cogY = y + H * 0.5;
const cogRadius = 6;
const toothLength = 2;
const numTeeth = 8;
const holeRadius = 2; // Radius of the center hole
// Draw the cog
ctx.beginPath();
ctx.arc(cogX, cogY, cogRadius - toothLength, 0, Math.PI * 2);
ctx.fill();
// Draw the center hole (by clearing it)
ctx.beginPath();
ctx.arc(cogX, cogY, holeRadius, 0, Math.PI * 2);
ctx.fillStyle = litegraph_base.WIDGET_BGCOLOR;
ctx.fill();
// Reset fill style for the teeth
if (isAuto) {
ctx.fillStyle = litegraph_base.WIDGET_TEXT_COLOR;
} else {
ctx.fillStyle = litegraph_base.WIDGET_SECONDARY_TEXT_COLOR;
}
// Draw teeth
ctx.beginPath();
for (let i = 0; i < numTeeth; i++) {
const angle = (i / numTeeth) * Math.PI * 2;
const innerX = cogX + (cogRadius - toothLength) * Math.cos(angle);
const innerY = cogY + (cogRadius - toothLength) * Math.sin(angle);
const outerX = cogX + cogRadius * Math.cos(angle);
const outerY = cogY + cogRadius * Math.sin(angle);
ctx.moveTo(innerX, innerY);
ctx.lineTo(outerX, outerY);
}
ctx.lineWidth = 2;
ctx.stroke();
ctx.restore();
// Draw label
ctx.fillStyle = litegraph_base.WIDGET_SECONDARY_TEXT_COLOR;
const label = this.label || this.name;
if (label != null) {
ctx.fillText(label, margin * 2 + 5, y + H * 0.7);
}
// Draw value
ctx.textAlign = 'right';
const text = isAuto ? "auto" : this.displayValue();
ctx.fillStyle = isAuto ? litegraph_base.WIDGET_SECONDARY_TEXT_COLOR : litegraph_base.WIDGET_TEXT_COLOR;
ctx.fillText(text, widget_width - autoTextRightMargin - autoTextWidth - 15, y + H * 0.7);
// Draw increment/decrement buttons if not in AUTO mode and not a string widget
if (!isAuto && !this.disabled && this.config[0] !== "FVAUTOSTRING") {
// Draw decrement button (left triangle)
ctx.fillStyle = litegraph_base.WIDGET_TEXT_COLOR;
ctx.beginPath();
ctx.moveTo(margin + 16, y + 5);
ctx.lineTo(margin + 6, y + H * 0.5);
ctx.lineTo(margin + 16, y + H - 5);
ctx.fill();
// Draw increment button (right triangle)
ctx.beginPath();
ctx.moveTo(widget_width - margin - 16, y + 5);
ctx.lineTo(widget_width - margin - 6, y + H * 0.5);
ctx.lineTo(widget_width - margin - 16, y + H - 5);
ctx.fill();
}
}
}
function mouseAutoAnnotated(event, [x, y], node) {
const widget_width = node.size[0];
const margin = 15;
const H = 20; // Widget height
const autoTextWidth = 30;
const autoTextRightMargin = 5;
const cogRadius = 6;
if (this.isAuto) {
if (event.type === "pointerup" || event.type === "mouseup") {
const cogX = widget_width - autoTextRightMargin - autoTextWidth - 6;
const cogLeftEdge = cogX - cogRadius;
const cogRightEdge = cogX + cogRadius;
if (x > cogLeftEdge && x < cogRightEdge) {
this.isAuto = false;
this.value = this.cachedValue !== undefined ? this.cachedValue : (this.options.default || 0);
if (this.callback) {
this.callback(this.value);
}
node.graph.setDirtyCanvas(true, false);
}
}
// Block ALL events in auto mode except cog clicks
event.preventDefault?.();
event.stopPropagation?.();
event.stopImmediatePropagation?.();
return true; // Always return true to indicate event was handled
}
// Determine if clicking on increment/decrement buttons
const delta = this.config[0] === "FVAUTOSTRING" ? 0 :
(x < 40 ? -1 : x > widget_width - 48 ? 1 : 0);
if (event.type === "pointerdown" || event.type === "mousedown") {
// ComfyUI appears to intercept pointerdown events, so this code path is never reached
console.log("pointerdown received (unexpected)");
return false;
} else if (event.type === "pointerup" || event.type === "mouseup") {
// Stop event propagation to prevent double handling
event.preventDefault?.();
event.stopPropagation?.();
event.stopImmediatePropagation?.();
const cogX = widget_width - autoTextRightMargin - autoTextWidth - 6;
const cogLeftEdge = widget_width - autoTextRightMargin - autoTextWidth - 6 - cogRadius;
const cogRightEdge = widget_width - autoTextRightMargin - autoTextWidth - 6 + cogRadius;
if (x > cogLeftEdge && x < cogRightEdge) {
this.isAuto = !this.isAuto;
if (this.isAuto) {
this.cachedValue = this.value;
this.value = -99999;
} else {
this.value = this.cachedValue !== undefined ? this.cachedValue : (this.options.default || 0);
}
if (this.callback) {
this.callback(this.value);
}
node.graph.setDirtyCanvas(true, false);
return true;
}
// If in auto mode and NOT clicking the cog, block all other interactions
if (this.isAuto) {
return true;
}
// Handle increment/decrement buttons if not in auto mode
if (delta !== 0 && !this.isAuto) {
if (this.config[0] === "FVAUTOCOMBO") {
const options = this.options.values || [];
if (options.length === 0) return true;
let currentIndex = -1;
for (let i = 0; i < options.length; i++) {
const optValue = typeof options[i] === 'object' ? options[i].value : options[i];
if (optValue == this.value || String(optValue) === String(this.value)) {
currentIndex = i;
break;
}
}
if (currentIndex === -1) {
currentIndex = 0;
}
let newIndex = currentIndex + delta;
if (newIndex < 0) {
newIndex = options.length - 1;
} else if (newIndex >= options.length) {
newIndex = 0;
}
const newOption = options[newIndex];
this.value = typeof newOption === 'object' ? newOption.value : newOption;
if (this.callback) {
this.callback(this.value);
}
node.graph.setDirtyCanvas(true, false);
return true;
} else {
let v = parseFloat(this.value);
const increment = delta * 0.1 * (this.options.step || 1);
v += increment;
// Apply min/max constraints
if (this.options.min != null) {
v = Math.max(this.options.min, v);
}
if (this.options.max != null) {
v = Math.min(this.options.max, v);
}
// Round to precision or to integer
if (this.config[0] === "FVAUTOINT") {
v = Math.round(v);
} else if (this.options.precision !== undefined) {
const precision = Math.pow(10, this.options.precision);
v = Math.round(v * precision) / precision;
}
this.value = v;
if (this.callback) {
this.callback(this.value);
}
node.graph.setDirtyCanvas(true, false);
return true;
}
}
if (delta === 0 && !this.isAuto) {
if (this.config[0] === "FVAUTOCOMBO") {
const options = this.options.values || [];
// Create menu items
const menuItems = options.map(opt => {
const value = typeof opt === 'object' ? opt.value : opt;
const label = typeof opt === 'object' ? opt.label : opt.toString();
return {
content: label,
callback: () => {
this.value = value;
if (this.callback) {
this.callback(this.value);
}
node.graph.setDirtyCanvas(true, false);
}
};
});
new LiteGraph.ContextMenu(menuItems, {
event: event,
title: null,
callback: null,
extra: node
});
return true;
} else if (this.config[0] === "FVAUTOSTRING") {
const d_callback = (v) => {
this.value = v;
if (this.callback) {
this.callback(this.value);
}
node.graph.setDirtyCanvas(true, false);
};
const dialog = app.canvas.prompt(
'Value',
this.value,
d_callback,
event
);
return true;
} else {
// For numeric widgets, show input dialog
const d_callback = (v) => {
this.value = this.parseValue?.(v) ?? Number(v);
// Apply min/max constraints
if (this.options.min != null) {
this.value = Math.max(this.options.min, this.value);
}
if (this.options.max != null) {
this.value = Math.min(this.options.max, this.value);
}
// Round to precision or to integer
if (this.config[0] === "FVAUTOINT") {
this.value = Math.round(this.value);
} else if (this.options.precision !== undefined) {
const precision = Math.pow(10, this.options.precision);
this.value = Math.round(this.value * precision) / precision;
}
if (this.callback) {
this.callback(this.value);
}
node.graph.setDirtyCanvas(true, false);
};
const dialog = app.canvas.prompt(
'Value',
this.value,
d_callback,
event
);
return true;
}
}
return true;
}
return false;
}
function makeAutoAnnotated(widget, inputData) {
const original = {
callback: widget.callback,
type: widget.type,
value: widget.value
};
// Add AUTO properties to the widget
Object.assign(widget, {
type: "BOOLEAN",
draw: drawAutoAnnotated,
mouse: mouseAutoAnnotated,
onMouse: null, // Explicitly disable original onMouse handler
isAuto: true,
cachedValue: widget.value,
config: inputData,
options: Object.assign({}, inputData[1], widget.options),
original: original, // Store original properties for reference
// Disable other potential mouse handlers with no-op functions
onClick: function () {
return false;
},
onPointerUp: function () {
return false;
},
onPointerDown: function () {
return false;
},
onMouseUp: function () {
return false;
},
onMouseDown: function () {
return false;
},
computeSize(width) {
return [width, 20];
},
displayValue: function () {
if (this.config[0] === "FVAUTOINT") {
return Math.round(this.value).toString();
}
if (this.config[0] === "FVAUTOCOMBO") {
return this.value;
}
if (this.config[0] === "FVAUTOSTRING") {
return this.value;
}
// For FLOAT values, check if it's actually an integer
if (Number.isInteger(this.value)) {
return this.value.toString();
}
return this.value.toFixed(this.options.precision || 2);
},
parseValue: function (v) {
if (this.config[0] === "FVAUTOSTRING") {
return v;
}
if (typeof v === "string") {
return parseFloat(v);
}
return v;
},
serializeValue: function () {
// Return special value for AUTO mode
return this.isAuto ? -99999 : this.value;
},
deserializeValue: function (data) {
if (data === -99999) {
this.isAuto = true;
this.value = -99999;
} else {
this.isAuto = false;
this.value = data;
this.cachedValue = data;
}
}
});
// Override callback to handle AUTO mode
widget.callback = function (v) {
if (this.isAuto) {
return; // Don't call the original callback in AUTO mode
}
const result = original.callback?.call(this, v);
return result;
};
// Override any potential click handlers
const originalOnClick = widget.onClick;
if (originalOnClick) {
widget.onClick = function (...args) {
if (this.isAuto) {
return false;
}
return originalOnClick.call(this, ...args);
};
}
return widget;
}
app.registerExtension({
name: "FastVideo.AutoWidgets",
async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeData?.name == "VideoGenerator" || nodeData?.name === "InferenceArgs" || nodeData?.name === "VAEConfig" ||
nodeData?.name === "TextEncoderConfig" || nodeData?.name === "DITConfig") {
// Add serialization support
chainCallback(nodeType.prototype, "onSerialize", function (info) {
if (!this.widgets) {
return;
}
// Ensure widgets_values exists
if (!info.widgets_values) {
info.widgets_values = {};
}
// Store AUTO widget states in a separate property
if (!info.auto_widget_states) {
info.auto_widget_states = {};
}
// Handle AUTO widgets specially
for (const w of this.widgets) {
if (w.type === "BOOLEAN" && w.isAuto !== undefined) {
// Store the serialized value (for Python node)
info.widgets_values[w.name] = w.serializeValue();
// Store the full state (for UI restoration)
info.auto_widget_states[w.name] = {
isAuto: w.isAuto,
value: w.value,
cachedValue: w.cachedValue
};
}
}
});
// Add deserialization support
chainCallback(nodeType.prototype, "onConfigure", function (info) {
if (!this.widgets) {
return;
}
// First, restore from widgets_values (for backward compatibility)
if (info.widgets_values && Array.isArray(info.widgets_values)) {
for (let i = 0; i < this.widgets.length && i < info.widgets_values.length; i++) {
const w = this.widgets[i];
const value = info.widgets_values[i];
if (w.type === "BOOLEAN" && w.isAuto !== undefined) {
w.deserializeValue(value);
}
}
}
// Then, restore full state if available
if (info.auto_widget_states) {
for (const w of this.widgets) {
if (w.type === "BOOLEAN" && w.isAuto !== undefined && w.name in info.auto_widget_states) {
const state = info.auto_widget_states[w.name];
w.isAuto = state.isAuto;
w.cachedValue = state.cachedValue;
w.value = state.isAuto ? -99999 : state.value;
w.callback?.(w.value);
}
}
}
// Force a redraw
this.graph?.setDirtyCanvas(true, true);
});
// Override addInput to handle AUTO widgets
chainCallback(nodeType.prototype, "onNodeCreated", function () {
// Convert any existing widgets to AUTO widgets if needed
let new_widgets = [];
const intWidgetNames = ["sp_size", "tp_size", "height", "width", "num_frames", "num_inference_steps", "flow_shift", "seed", "fps", "scale_factor",
"tile_sample_min_height", "tile_sample_min_width", "tile_sample_min_num_frames", "tile_sample_stride_height", "tile_sample_stride_width",
"tile_sample_stride_num_frames", "blend_num_frames"
]
const floatWidgetNames = ["embedded_cfg_scale", "guidance_scale"]
const comboWidgetNames = ["vae_tiling", "vae_precision", "vae_sp", "text_encoder_precision", "precision",
"load_encoder", "load_decoder", "use_tiling", "use_temporal_tiling", "use_parallel_tiling", "dit_cpu_offload", "enable_teacache"
]
const stringWidgetNames = ["prefix", "quant_config", "lora_config", "image_path"]
if (this.widgets) {
for (let w of this.widgets) {
if (intWidgetNames.includes(w.name)) {
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOINT", { "default": 0 }]));
} else if (floatWidgetNames.includes(w.name)) {
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOFLOAT", { "default": 0 }]));
} else if (comboWidgetNames.includes(w.name)) {
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOCOMBO", { "default": 0 }]));
} else if (stringWidgetNames.includes(w.name)) {
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOSTRING", { "default": "" }]));
} else {
new_widgets.push(w);
}
}
this.widgets = new_widgets;
const autoWidgets = this.widgets.filter(w => w.type === "BOOLEAN" && w.isAuto !== undefined);
}
this.graph?.setDirtyCanvas(true, true);
});
}
},
async init() {
// Force a redraw of all nodes when the extension initializes
if (app.graph) {
setTimeout(() => {
app.graph.setDirtyCanvas(true, true);
}, 1000);
}
}
});
console.log("FastVideo.core.js loaded");
+48 -18
View File
@@ -1,11 +1,21 @@
# Sliding Tile Atteniton Kernel
# Attention Kernel Used in FastVideo
## 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:
## Video Sparse Attention (VSA)
### Installation
We support H100 (via TK) and any other GPU (via triton) for VSA.
```bash
git submodule update --init --recursive
python setup_vsa.py install
```
If you encounter error during installation, try below:
Install C++20 for ThunderKittens:
```bash
sudo apt update
sudo apt install gcc-11 g++-11
@@ -15,27 +25,48 @@ sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave
sudo apt update
sudo apt install clang-11
```
## Environment Setup
First, set up your CUDA environment:
(If you use CUDA12.8)
```bash
export CUDA_HOME=/usr/local/cuda-12.4
export CUDA_HOME=/usr/local/cuda-12.8
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)
### Verify if you have successfully installed
```bash
# test numerical
python tests/test_vsa.py
# (For H100) test speed
python benchmarks/bench_vsa_hopper.py
```
bench_vsa_hopper.py should print something like this:
```bash
Using topk=76 kv blocks per q block (out of 768 total kv blocks)
=== BLOCK SPARSE ATTENTION BENCHMARK ===
Block Sparse Forward - TFLOPS: 5622.26
Block Sparse Backward - TFLOPS: 3865.68
```
## Sliding Tile Attention (STA)
We only support H100 for STA.
```bash
git submodule update --init --recursive
python setup_sta.py install
```
## Install Video Sparse Attention (VSA)
### Usage
End-2-end inference with FastVideo:
```bash
python setup_vsa.py install
bash scripts/inference/v1_inference_wan_STA.sh
```
## Usage
If you want to use sliding tile attention in your custom model:
```python
from st_attn import sliding_tile_attention
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
@@ -47,22 +78,21 @@ from st_attn import sliding_tile_attention
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
### Test
```bash
python tests/test_sta.py # test STA
python tests/test_block_sparse.py # test VSA
python tests/test_vsa.py # test VSA
```
## Benchmark
### Benchmark
```bash
python benchmarks/bench_sta.py
```
## How Does STA Work?
### How Does STA Work?
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
-2
View File
@@ -86,8 +86,6 @@ def benchmark_attention(configurations):
# print(f"Average TFLOPS: {tflops_bwd}")
# print("=" * 60)
torch.cuda.empty_cache()
return results
@@ -1,9 +1,9 @@
import torch
import argparse
from flash_attn.utils.benchmark import benchmark_forward
from vsa import block_sparse_attention_fwd, block_sparse_attention_backward
from triton.testing import do_bench
from vsa import block_sparse_fwd, block_sparse_bwd
from vsa import BLOCK_M, BLOCK_N
import triton
import numpy as np
import random
@@ -23,7 +23,7 @@ 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('--head_dim', type=int, default=128, 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()
@@ -130,19 +130,19 @@ def benchmark_block_sparse_attention(q, k, v, q2k_block_sparse_index, q2k_block_
# Forward pass
# Warm-up run
o, l_vec = block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
variable_block_sizes = torch.ones(q2k_block_sparse_index.shape[2], device=q.device).int() * BLOCK_M
o, l_vec = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
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'
fwd_time = do_bench(
lambda: block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes),
warmup=5,
rep=20,
quantiles=None
)
sparse_tflops = flops / fwd_time.mean * 1e-12
sparse_tflops = flops / fwd_time * 1e-12 * 1e3
print(f"Block Sparse Forward - TFLOPS: {sparse_tflops:.2f}")
# Backward pass
@@ -150,20 +150,19 @@ def benchmark_block_sparse_attention(q, k, v, q2k_block_sparse_index, q2k_block_
# 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)
block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes)
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_time = do_bench(
lambda: block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes),
warmup=5,
rep=20,
quantiles=None
)
bwd_flops = 2.5 * flops # Approximation
sparse_bwd_tflops = bwd_flops / bwd_time.mean * 1e-12
sparse_bwd_tflops = bwd_flops / bwd_time * 1e-12 * 1e3
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd_tflops:.2f}")
return sparse_tflops, sparse_bwd_tflops
+217
View File
@@ -0,0 +1,217 @@
import torch
import argparse
import triton.testing
from vsa import block_sparse_attn
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+backward pass."""
print("\n=== BLOCK SPARSE ATTENTION FORWARD+BACKWARD BENCHMARK ===")
# Combined forward+backward pass
# Warm-up run
q_fwd = q.clone().requires_grad_(True)
k_fwd = k.clone().requires_grad_(True)
v_fwd = v.clone().requires_grad_(True)
o = block_sparse_attn(q_fwd, k_fwd, v_fwd, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
grad_output = torch.randn_like(o)
o.backward(grad_output)
torch.cuda.synchronize()
# Benchmark forward+backward
def forward_backward_fn():
q_fwd = q.clone().requires_grad_(True)
k_fwd = k.clone().requires_grad_(True)
v_fwd = v.clone().requires_grad_(True)
o = block_sparse_attn(q_fwd, k_fwd, v_fwd, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
grad_output = torch.randn_like(o)
o.backward(grad_output)
total_time = triton.testing.do_bench(
forward_backward_fn,
warmup=25,
rep=100,
return_mode='mean'
)
# Total flops for forward + backward (forward + 2.5x backward approximation)
total_flops = flops + 2.5 * flops # 3.5x the forward flops
sparse_tflops = total_flops / total_time * 1e-12 * 1e3
print(f"Block Sparse Forward+Backward - TFLOPS: {sparse_tflops:.2f}")
return sparse_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 = 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+Backward - TFLOPS: {sparse_fwd:.2f}")
if __name__ == "__main__":
main()
@@ -1,2 +1,2 @@
recursive-include tk *
include config.py
include config_sta.py
+87
View File
@@ -0,0 +1,87 @@
# Attention Kernel Used in FastVideo
## Sliding Tile Attention (STA)
We only support H100 for STA.
### Installation
```bash
pip install st_attn
```
Install from source:
```bash
git submodule update --init --recursive
python setup.py install
```
If you encounter error during installation, try below:
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
```
(If you use CUDA12.8)
```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
```
### Usage
End-2-end inference with FastVideo:
```bash
bash scripts/inference/v1_inference_wan_STA.sh
```
If you want to use sliding tile attention in your custom model:
```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_vsa.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.
@@ -1,7 +1,7 @@
import os
import subprocess
from csrc.attn.config_sta import kernels, sources, target
from config_sta import kernels, sources, target
from setuptools import find_packages, setup
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
@@ -9,7 +9,7 @@ target = target.lower()
# Package metadata
PACKAGE_NAME = "st_attn"
VERSION = "0.0.4"
VERSION = "0.0.6"
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"
-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.")
+156
View File
@@ -0,0 +1,156 @@
import torch
import sys
import os
import numpy as np
from tqdm import tqdm
# Add the parent directory to the path to import block_sparse_attn
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from tests.utils import generate_block_sparse_mask_for_function, create_full_mask_from_block_mask
from vsa import block_sparse_attn
BLOCK_M = 64
BLOCK_N = 64
def pytorch_test(Q, K, V, block_sparse_mask, dO):
q_ = Q.clone().float().requires_grad_()
k_ = K.clone().float().requires_grad_()
v_ = V.clone().float().requires_grad_()
QK = torch.matmul(q_, k_.transpose(-2, -1))
QK /= (q_.size(-1) ** 0.5)
QK = QK.masked_fill(~block_sparse_mask.unsqueeze(0), float('-inf'))
QK = torch.nn.functional.softmax(QK, dim=-1)
output = torch.matmul(QK, v_)
dO_ = dO
output.backward(dO_)
return (
output.to(torch.bfloat16),
q_.grad.to(torch.bfloat16),
k_.grad.to(torch.bfloat16),
v_.grad.to(torch.bfloat16),
)
def block_sparse_kernel_test(Q, K, V, block_sparse_mask, variable_block_sizes, non_pad_index, dO):
Q = Q.detach().requires_grad_()
K = K.detach().requires_grad_()
V = V.detach().requires_grad_()
q_padded = vsa_pad(Q, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
k_padded = vsa_pad(K, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
v_padded = vsa_pad(V, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
output, _= block_sparse_attn(q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes)
output = output[:, :, non_pad_index, :]
output.backward(dO)
return output, Q.grad, K.grad, V.grad
def get_non_pad_index(
vid_len: torch.LongTensor,
n_win: int,
win_size: int,
):
device = vid_len.device
starts_pad = torch.arange(n_win, device=device) * win_size
index_pad = starts_pad[:, None] + torch.arange(win_size, device=device)[None, :]
index_mask = torch.arange(win_size, device=device)[None, :] < vid_len[:, None]
return index_pad[index_mask]
def generate_tensor(shape, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
return tensor
def generate_variable_block_sizes(num_blocks, min_size=32, max_size=64, device="cuda"):
return torch.randint(min_size, max_size + 1, (num_blocks,), device=device, dtype=torch.int32)
def vsa_pad(x, non_pad_index, num_blocks, block_size):
padded_x = torch.zeros((1, x.shape[1], num_blocks * BLOCK_M, x.shape[3]), device=x.device, dtype=x.dtype)
padded_x[:, :, non_pad_index, :] = x
return padded_x
def check_correctness(h, d, num_blocks, k, num_iterations=20, error_mode='all'):
results = {
'gO': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
'gQ': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
'gK': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
'gV': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
}
device = "cuda" if torch.cuda.is_available() else "cpu"
variable_block_sizes = generate_variable_block_sizes(num_blocks, device=device)
S = int(variable_block_sizes.sum().item())
padded_S = num_blocks * BLOCK_M
non_pad_index = get_non_pad_index(variable_block_sizes, num_blocks, BLOCK_M)
block_mask = generate_block_sparse_mask_for_function(h, num_blocks, k, device)
full_mask = create_full_mask_from_block_mask(block_mask, variable_block_sizes, device)
for _ in range(num_iterations):
Q = generate_tensor((1, h, S, d), torch.bfloat16, device)
K = generate_tensor((1, h, S, d), torch.bfloat16, device)
V = generate_tensor((1, h, S, d), torch.bfloat16, device)
dO = generate_tensor((1, h, S, d), torch.bfloat16, device)
# dO_padded = torch.zeros_like(dO_padded)
# dO_padded[:, :, non_pad_index, :] = dO
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, full_mask, dO)
bs_o, bs_qg, bs_kg, bs_vg = block_sparse_kernel_test(Q, K, V, block_mask.unsqueeze(0), variable_block_sizes,non_pad_index, dO)
for name, (pt, bs) in zip(['gQ', 'gK', 'gV', 'gO'], [(pt_qg, bs_qg), (pt_kg, bs_kg), (pt_vg, bs_vg), (pt_o, bs_o)]):
if bs is not None:
diff = pt - bs
abs_diff = torch.abs(diff)
results[name]['sum_diff'] += torch.sum(abs_diff).item()
results[name]['sum_abs'] += torch.sum(torch.abs(pt)).item()
rel_max_diff = torch.max(abs_diff) / torch.mean(torch.abs(pt))
results[name]['max_diff'] = max(results[name]['max_diff'], rel_max_diff.item())
if torch.cuda.is_available():
torch.cuda.empty_cache()
total_elements = h * S * d * num_iterations
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_graphs(h, d, error_mode='all'):
test_configs = [
{"num_blocks": 16, "k": 2, "description": "Small sequence"},
{"num_blocks": 32, "k": 4, "description": "Medium sequence"},
{"num_blocks": 53, "k": 6, "description": "Large sequence"},
]
print(f"\nError Analysis for h={h}, d={d}, mode={error_mode}")
print("=" * 150)
print(f"{'Config':<20} {'Blocks':<8} {'K':<4} "
f"{'gQ Avg':<12} {'Rel gQ Max':<12} "
f"{'gK Avg':<12} {'Rel gK Max':<12} "
f"{'gV Avg':<12} {'Rel gV Max':<12} "
f"{'gO Avg':<12} {'Rel gO Max':<12}")
print("-" * 150)
for config in test_configs:
num_blocks = config["num_blocks"]
k = config["k"]
description = config["description"]
results = check_correctness(h, d, num_blocks, k, error_mode=error_mode)
print(f"{description:<20} {num_blocks:<8} {k:<4} "
f"{results['gQ']['avg_diff']:<12.6e} {results['gQ']['max_diff']:<12.6e} "
f"{results['gK']['avg_diff']:<12.6e} {results['gK']['max_diff']:<12.6e} "
f"{results['gV']['avg_diff']:<12.6e} {results['gV']['max_diff']:<12.6e} "
f"{results['gO']['avg_diff']:<12.6e} {results['gO']['max_diff']:<12.6e}")
print("-" * 150)
if __name__ == "__main__":
h, d = 16, 128
print("Block Sparse Attention with Variable Block Sizes Analysis")
print("=" * 60)
for mode in ['backward']:
generate_error_graphs(h, d, error_mode=mode)
print("\nAnalysis completed for all modes.")
+54
View File
@@ -0,0 +1,54 @@
import torch
def generate_block_sparse_mask_for_function(h, num_blocks, k, device="cuda"):
"""
Generate block sparse mask of shape [h, num_blocks, num_blocks].
Args:
h: number of heads
num_blocks: number of blocks
k: number of kv blocks each q block attends to
device: device to create tensors on
Returns:
block_sparse_mask: [h, num_blocks, num_blocks] bool tensor
"""
k = min(k, num_blocks)
scores = torch.rand(h, num_blocks, num_blocks, device=device)
_, indices = torch.topk(scores, k, dim=-1)
block_sparse_mask = torch.zeros(h, num_blocks, num_blocks, dtype=torch.bool, device=device)
block_sparse_mask = block_sparse_mask.scatter_(2, indices, 1).bool()
return block_sparse_mask
def create_full_mask_from_block_mask(block_sparse_mask, variable_block_sizes, device="cuda"):
"""
Convert block-level sparse mask to full attention mask.
Args:
block_sparse_mask: [h, num_blocks, num_blocks] bool tensor
variable_block_sizes: [num_blocks] tensor
device: device to create tensors on
Returns:
full_mask: [h, S, S] bool tensor where S = total sequence length
"""
h, num_blocks, _ = block_sparse_mask.shape
total_seq_len = variable_block_sizes.sum().item()
cumsum = torch.cat([torch.tensor([0], device=device), variable_block_sizes.cumsum(dim=0)[:-1]])
full_mask = torch.zeros(h, total_seq_len, total_seq_len, dtype=torch.bool, device=device)
for head in range(h):
for q_block in range(num_blocks):
q_start = cumsum[q_block]
q_end = q_start + variable_block_sizes[q_block]
for kv_block in range(num_blocks):
if block_sparse_mask[head, q_block, kv_block]:
kv_start = cumsum[kv_block]
kv_end = kv_start + variable_block_sizes[kv_block]
full_mask[head, q_start:q_end, kv_start:kv_end] = True
return full_mask
Submodule csrc/attn/tk deleted from 1719fb7264
+2
View File
@@ -0,0 +1,2 @@
recursive-include tk *
include config_vsa.py
+61
View File
@@ -0,0 +1,61 @@
# Attention Kernel Used in FastVideo
## Video Sparse Attention (VSA)
### Installation
We support H100 (via TK) and any other GPU (via triton) for VSA.
```bash
pip install vsa
```
Install from source:
```bash
git submodule update --init --recursive
python setup.py install
```
If you encounter error during installation, try below:
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
```
(If you use CUDA12.8)
```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
```
### Verify if you have successfully installed
```bash
# test numerical
python ../tests/test_vsa.py
# (For H100) test speed
python ../benchmarks/bench_vsa_hopper.py
```
bench_vsa_hopper.py should print something like this:
```bash
Using topk=76 kv blocks per q block (out of 768 total kv blocks)
=== BLOCK SPARSE ATTENTION BENCHMARK ===
Block Sparse Forward - TFLOPS: 5622.26
Block Sparse Backward - TFLOPS: 3865.68
```
## Acknowledgement
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
@@ -1,7 +1,7 @@
import os
import subprocess
from csrc.attn.config_vsa import kernels, sources, target
from config_vsa import kernels, sources, target
from setuptools import find_packages, setup
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
@@ -9,10 +9,10 @@ target = target.lower()
# Package metadata
PACKAGE_NAME = "vsa"
VERSION = "0.0.1"
VERSION = "0.0.3"
AUTHOR = "Hao AI Lab"
DESCRIPTION = "Video Sparse Attention Kernel Used in FastVideo"
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn"
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn/video_sparse_attn"
# Set environment variables
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
@@ -51,21 +51,26 @@ for k in kernels:
source_files.append(sources[k]['source_files'][target])
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
ext_modules = [
CUDAExtension('vsa_cuda',
sources=source_files,
extra_compile_args={
'cxx': cpp_flags,
'nvcc': cuda_flags
},
libraries=['cuda'])
]
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'])
],
ext_modules=ext_modules,
cmdclass={'build_ext': BuildExtension},
classifiers=[
"Programming Language :: Python :: 3",
@@ -9,10 +9,10 @@
#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
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q2k_block_sparse_index, torch::Tensor q2k_block_sparse_num, torch::Tensor block_size
);
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
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, torch::Tensor block_size
);
#endif
@@ -0,0 +1,80 @@
import torch
from typing import Tuple
block_sparse_attn=None
import torch
major, minor = torch.cuda.get_device_capability(0)
if major == 9 and minor == 0:# check if H100
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
from vsa.block_sparse_wrapper import block_sparse_attn_SM90
block_sparse_attn = block_sparse_attn_SM90
else:
from vsa.block_sparse_wrapper import block_sparse_attn_triton
block_sparse_fwd = None
block_sparse_bwd = None
block_sparse_attn = block_sparse_attn_triton
BLOCK_M = 64
BLOCK_N = 64
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 video_sparse_attn(q, k, v, variable_block_sizes, 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]
NOTE: We assume q, k, v is zero padded!!
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
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).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(q.dtype)
k_compress = (k.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(k.dtype)
v_compress = (v.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(v.dtype)
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)
topK_indices = torch.topk(block_attn_score, topk, dim=-1).indices
block_mask = torch.zeros_like(block_attn_score, dtype=torch.bool).scatter_(-1, topK_indices, True)
output_select, _ = block_sparse_attn(q, k, v, block_mask, variable_block_sizes)
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
@@ -0,0 +1,449 @@
"""
Fused Attention
===============
This is a Triton implementation of the Flash Attention v2 algorithm from Tri Dao
(https://tridao.me/publications/flash2/flash2.pdf)
Credits: OpenAI kernel team
"""
import pytest
import torch
import triton
import triton.language as tl
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
import math # small utility needed by the sparse wrapper
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
# We don't run auto-tuning every time to keep the tutorial fast. Keeping
# the code below and commenting out the equivalent parameters is convenient for
# re-tuning.
configs = [
triton.Config({'BLOCK_M': BM, 'BLOCK_N': BN}, num_stages=s, num_warps=w) \
for BM in [64]\
for BN in [64]\
for s in [3, 4, 7]\
for w in [4, 8]\
]
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
@triton.autotune(configs, key=["N_CTX", "HEAD_DIM"])
@triton.jit
def _attn_fwd_sparse(Q, K, V, sm_scale, #
q2k_index, q2k_num, max_kv_blks, #
variable_block_sizes,
M, Out, #
stride_qz, stride_qh, stride_qm, stride_qk,
stride_kz, stride_kh, stride_kn, stride_kk,
stride_vz, stride_vh, stride_vk, stride_vn,
stride_oz, stride_oh, stride_om, stride_on,
Z, H, N_CTX, #
HEAD_DIM: tl.constexpr, #
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
STAGE: tl.constexpr):
"""
64×64 **block-sparse** forward kernel. Back-prop kernels remain dense
(32×64 and 64×32) – memory footprint unchanged.
"""
# ----- program-id mapping -----
q_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(1) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX // BLOCK_M
meta_base = ((b * H + h) * q_tiles + q_blk)
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
# ----- base pointers -----
qvk_off = (b.to(tl.int64) * stride_qz +
h.to(tl.int64) * stride_qh)
Q_ptr = tl.make_block_ptr(
base=Q + qvk_off, shape=(N_CTX, HEAD_DIM),
strides=(stride_qm, stride_qk),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
K_base = tl.make_block_ptr(
base=K + qvk_off, shape=(HEAD_DIM, N_CTX),
strides=(stride_kk, stride_kn),
offsets=(0, 0),
block_shape=(HEAD_DIM, BLOCK_N), order=(0, 1))
v_order: tl.constexpr = (0, 1) if V.dtype.element_ty == tl.float8e5 else (1, 0)
V_base = tl.make_block_ptr(
base=V + qvk_off, shape=(N_CTX, HEAD_DIM),
strides=(stride_vk, stride_vn),
offsets=(0, 0),
block_shape=(BLOCK_N, HEAD_DIM), order=v_order)
O_ptr = tl.make_block_ptr(
base=Out + qvk_off, shape=(N_CTX, HEAD_DIM),
strides=(stride_om, stride_on),
offsets=(q_blk * BLOCK_M, 0),
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
# ----- accumulators -----
offs_m = q_blk * BLOCK_M + tl.arange(0, BLOCK_M)
m_i = tl.full([BLOCK_M], -float("inf"), tl.float32)
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
qk_scale = sm_scale * 1.44269504 # 1/ln2
q = tl.load(Q_ptr)
# ----- sparse loop over valid K/V tiles -----
for i in range(0, kv_blocks):
kv_idx = tl.load(kv_ptr + i).to(tl.int32)
block_size = tl.load(variable_block_sizes + kv_idx)
K_ptr = tl.advance(K_base, (0, kv_idx * BLOCK_N))
V_ptr = tl.advance(V_base, (kv_idx * BLOCK_N, 0))
k = tl.load(K_ptr)
qk = tl.dot(q, k)
# mask out invalid columns
mask = tl.arange(0, BLOCK_N) < block_size
qk = tl.where(mask[None, :], qk, -float("inf"))
m_ij = tl.maximum(m_i, tl.max(qk, 1) * qk_scale)
p = tl.math.exp2(qk * qk_scale - m_ij[:, None])
l_ij = tl.sum(p, 1)
alpha = tl.math.exp2(m_i - m_ij)
l_i = l_i * alpha + l_ij
acc = acc * alpha[:, None]
v = tl.load(V_ptr)
acc = tl.dot(p.to(tl.bfloat16), v, acc)
m_i = m_ij
# ----- epilogue -----
m_i += tl.math.log2(l_i)
acc = acc / l_i[:, None]
tl.store(M + off_hz * N_CTX + offs_m, m_i)
tl.store(O_ptr, acc.to(Out.type.element_ty))
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
@triton.jit
def _attn_bwd_preprocess(O, DO, #
Delta, #
Z, H, N_CTX, #
BLOCK_M: tl.constexpr, HEAD_DIM: tl.constexpr #
):
off_m = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
off_hz = tl.program_id(1)
off_n = tl.arange(0, HEAD_DIM)
# load
o = tl.load(O + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :])
do = tl.load(DO + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :]).to(tl.float32)
delta = tl.sum(o * do, axis=1)
# write-back
tl.store(Delta + off_hz * N_CTX + off_m, delta)
# The main inner-loop logic for computing dK and dV.
@triton.jit
def _attn_bwd_dkdv(dk, dv, #
Q, k, v, sm_scale, #
DO, #
M, D, #
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
# shared by Q/K/V/DO.
stride_tok, stride_d, #
H, N_CTX, BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
HEAD_DIM: tl.constexpr, #
# Filled in by the wrapper.
start_n, start_m, num_steps):
offs_m = start_m + tl.arange(0, BLOCK_M1)
offs_n = start_n + tl.arange(0, BLOCK_N1)
offs_k = tl.arange(0, HEAD_DIM)
qT_ptrs = Q + offs_m[None, :] * stride_tok + offs_k[:, None] * stride_d
do_ptrs = DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
# BLOCK_N1 must be a multiple of BLOCK_M1, otherwise the code wouldn't work.
tl.static_assert(BLOCK_N1 % BLOCK_M1 == 0)
step_m = BLOCK_M1
kv_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(2) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX // BLOCK_N1
meta_base = ((b * H + h) * q_tiles + kv_blk)
q_blocks = tl.load(k2q_num + meta_base) # int32
q_ptr = k2q_index + meta_base * max_q_blks # ptr to list
block_size = tl.load(variable_block_sizes + kv_blk)
for blk_idx in range(q_blocks*2):
block_sparse_offset = (tl.load(q_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_m
qT = tl.load(qT_ptrs + block_sparse_offset * stride_tok)
# Load m before computing qk to reduce pipeline stall.
offs_m = start_m + block_sparse_offset + tl.arange(0, BLOCK_M1)
m = tl.load(M + offs_m)
qkT = tl.dot(k, qT)
pT = tl.math.exp2(qkT - m[None, :])
mask = tl.arange(0, BLOCK_N1) < block_size
pT = tl.where(mask[:, None], pT, 0.0)
do = tl.load(do_ptrs + block_sparse_offset * stride_tok)
# Compute dV.
ppT = pT
ppT = ppT.to(tl.bfloat16)
dv += tl.dot(ppT, do)
# D (= delta) is pre-divided by ds_scale.
Di = tl.load(D + offs_m)
# Compute dP and dS.
dpT = tl.dot(v, tl.trans(do)).to(tl.float32)
dsT = pT * (dpT - Di[None, :])
dsT = dsT.to(tl.bfloat16)
dk += tl.dot(dsT, tl.trans(qT))
# Increment pointers.
return dk, dv
# the main inner-loop logic for computing dQ
@triton.jit
def _attn_bwd_dq(dq, q, K, V, #
do, m, D,
# shared by Q/K/V/DO.
q2k_index, q2k_num, max_kv_blks,
variable_block_sizes,
stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M2: tl.constexpr, #
BLOCK_N2: tl.constexpr, #
HEAD_DIM: tl.constexpr,
# Filled in by the wrapper.
start_m, start_n, num_steps):
offs_m = start_m + tl.arange(0, BLOCK_M2)
offs_n = start_n + tl.arange(0, BLOCK_N2)
offs_k = tl.arange(0, HEAD_DIM)
kT_ptrs = K + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
vT_ptrs = V + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
# D (= delta) is pre-divided by ds_scale.
Di = tl.load(D + offs_m)
# BLOCK_M2 must be a multiple of BLOCK_N2, otherwise the code wouldn't work.
tl.static_assert(BLOCK_M2 % BLOCK_N2 == 0)
step_n = BLOCK_N2
q_blk = tl.program_id(0) # Q-tile index
off_hz = tl.program_id(2) # fused (batch, head)
b = off_hz // H
h = off_hz % H
q_tiles = N_CTX // BLOCK_M2
meta_base = ((b * H + h) * q_tiles + q_blk)
kv_blocks = tl.load(q2k_num + meta_base) # int32
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
for blk_idx in range(kv_blocks*2):
block_sparse_offset = (tl.load(kv_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_n * stride_tok
block_size = tl.load(variable_block_sizes + blk_idx//2) - (blk_idx%2) * step_n
kT = tl.load(kT_ptrs + block_sparse_offset)
vT = tl.load(vT_ptrs + block_sparse_offset)
qk = tl.dot(q, kT)
p = tl.math.exp2(qk - m)
mask = tl.arange(0, BLOCK_N2) < block_size.to(tl.int32)
p = tl.where(mask[None, :], p , 0.0)
# Compute dP and dS.
dp = tl.dot(do, vT).to(tl.float32)
ds = p * (dp - Di[:, None])
ds = ds.to(tl.bfloat16)
# Compute dQ.
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
dq += tl.dot(ds, tl.trans(kT))
# Increment pointers.
return dq
@triton.jit
def _attn_bwd(Q, K, V, sm_scale, #
DO, #
DQ, DK, DV, #
M, D,
q2k_index, q2k_num, max_kv_blks,
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
# shared by Q/K/V/DO.
stride_z, stride_h, stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M1: tl.constexpr, #
BLOCK_N1: tl.constexpr, #
BLOCK_M2: tl.constexpr, #
BLOCK_N2: tl.constexpr, #
HEAD_DIM: tl.constexpr):
LN2 = 0.6931471824645996 # = ln(2)
bhid = tl.program_id(2)
off_chz = (bhid * N_CTX).to(tl.int64)
adj = (stride_h * (bhid % H) + stride_z * (bhid // H)).to(tl.int64)
pid = tl.program_id(0)
# offset pointers for batch/head
Q += adj
K += adj
V += adj
DO += adj
DQ += adj
DK += adj
DV += adj
M += off_chz
D += off_chz
# load scales
offs_k = tl.arange(0, HEAD_DIM)
start_n = pid * BLOCK_N1
start_m = 0
offs_n = start_n + tl.arange(0, BLOCK_N1)
dv = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
dk = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
# load K and V: they stay in SRAM throughout the inner loop.
k = tl.load(K + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
v = tl.load(V + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
num_steps = N_CTX // BLOCK_M1
dk, dv = _attn_bwd_dkdv( #
dk, dv, #
Q, k, v, sm_scale, #
DO, #
M, D, #
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M1, BLOCK_N1, HEAD_DIM, #
start_n, start_m, num_steps #
)
dv_ptrs = DV + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
tl.store(dv_ptrs, dv)
# Write back dK.
dk *= sm_scale
dk_ptrs = DK + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
tl.store(dk_ptrs, dk)
# THIS BLOCK DOES DQ:
start_m = pid * BLOCK_M2
end_n = 0
offs_m = start_m + tl.arange(0, BLOCK_M2)
q = tl.load(Q + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
dq = tl.zeros([BLOCK_M2, HEAD_DIM], dtype=tl.float32)
do = tl.load(DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
m = tl.load(M + offs_m)
m = m[:, None]
num_steps = N_CTX // BLOCK_N2
dq = _attn_bwd_dq(dq, q, K, V, #
do, m, D, #
q2k_index, q2k_num, max_kv_blks,
variable_block_sizes,
stride_tok, stride_d, #
H, N_CTX, #
BLOCK_M2, BLOCK_N2, HEAD_DIM, #
start_m, end_n, num_steps #
)
# Write back dQ.
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
dq *= LN2
tl.store(dq_ptrs, dq)
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num, variable_block_sizes):
B, H, T, D = q.shape
sm_scale = 1.0 / math.sqrt(D)
max_kv_blks = q2k_index.shape[-1]
assert T % 64 == 0, f"T must be a multiple of 64, but got {T}"
assert T // 64 == q2k_num.shape[-1], f"shape mismatch, T // 64 = {T // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
o = torch.empty_like(q)
M = torch.empty((B, H, T), dtype=torch.float32, device=q.device)
grid = lambda _: (triton.cdiv(T, 64), B * H, 1)
_attn_fwd_sparse[grid](
q, k, v, sm_scale,
q2k_index, q2k_num, max_kv_blks,
variable_block_sizes,
M, o,
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
k.stride(0), k.stride(1), k.stride(2), k.stride(3),
v.stride(0), v.stride(1), v.stride(2), v.stride(3),
o.stride(0), o.stride(1), o.stride(2), o.stride(3),
B, H, T,
HEAD_DIM=D, STAGE=3
)
return o, M
def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q_index, k2q_num, variable_block_sizes):
assert do.is_contiguous()
assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
B, H, T, D = q.shape
sm_scale = 1.0 / math.sqrt(D)
dq = torch.empty_like(q)
dk = torch.empty_like(k)
dv = torch.empty_like(v)
BATCH, N_HEAD, N_CTX = q.shape[:3]
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
arg_k = k
arg_k = arg_k * (sm_scale * RCP_LN2)
PRE_BLOCK = 64
assert N_CTX % PRE_BLOCK == 0
pre_grid = (N_CTX // PRE_BLOCK, BATCH * N_HEAD)
delta = torch.empty_like(M)
_attn_bwd_preprocess[pre_grid](
o, do, #
delta, #
BATCH, N_HEAD, N_CTX, #
BLOCK_M=PRE_BLOCK, HEAD_DIM=D #
)
max_q_blks = k2q_index.shape[-1]
max_kv_blks = q2k_index.shape[-1]
grid = (N_CTX // BLOCK_N1, 1, BATCH * N_HEAD)
_attn_bwd[grid](
q, arg_k, v, sm_scale, do, dq, dk, dv, #
M, delta, #
q2k_index, q2k_num, max_kv_blks,
k2q_index, k2q_num, max_q_blks,
variable_block_sizes,
q.stride(0), q.stride(1), q.stride(2), q.stride(3), #
N_HEAD, N_CTX, #
BLOCK_M1=BLOCK_M1, BLOCK_N1=BLOCK_N1, #
BLOCK_M2=BLOCK_M2, BLOCK_N2=BLOCK_N2, #
HEAD_DIM=D #
)
return dq, dk, dv
@@ -46,225 +46,9 @@ template<int D> struct fwd_globals {
int32_t *__restrict__ q2k_block_sparse_index;
int32_t *__restrict__ q2k_block_sparse_num;
int32_t *__restrict__ block_size;
};
template<int D>
__global__ __launch_bounds__(128, 3) // encourage compiler to reduce register usage so that an SM can hold 3 CTAs. Performance will drop from 391T to 353T if not specified explicitly.
void fwd_attend_ker_even(const __grid_constant__ fwd_globals<D> g) {
extern __shared__ int __shm[];
tma_swizzle_allocator al((int*)&__shm[0]);
using K = fwd_attend_ker_tile_dims<D>;
using q_tile = st_bf<64, K::tile_width>;
using k_tile = st_bf<128, K::tile_width>;
using v_tile = st_bf<128, K::tile_width>;
using k_tile_half = st_bf<64, K::tile_width>;
using v_tile_half = st_bf<64, K::tile_width>;
using l_col_vec = col_vec<st_fl<64, K::tile_width>>;
using o_tile = st_bf<64, K::tile_width>;
q_tile (&q_smem)[1] = al.allocate<q_tile, 1>();
k_tile_half (&k_smem_0)[1] = al.allocate<k_tile_half, 1 >();
k_tile_half (&k_smem_1)[1] = al.allocate<k_tile_half, 1 >();
k_tile (*k_smem) = reinterpret_cast<k_tile(*)>(k_smem_0);
v_tile_half (&v_smem_0)[1] = al.allocate<v_tile_half, 1 >();
v_tile_half (&v_smem_1)[1] = al.allocate<v_tile_half, 1 >();
v_tile (*v_smem) = reinterpret_cast<v_tile(*)>(v_smem_0);
l_col_vec (&l_smem)[1] = al.allocate<l_col_vec, 1>();
auto (*o_smem) = reinterpret_cast<o_tile(*)>(q_smem);
int kv_head_idx = blockIdx.y / g.hr;
int seq_idx = blockIdx.x;
int32_t* q2k_block_sparse_index_ptr = g.q2k_block_sparse_index + blockIdx.z * gridDim.y * gridDim.x * g.max_kv_blocks_per_q + blockIdx.y * gridDim.x * g.max_kv_blocks_per_q + blockIdx.x * g.max_kv_blocks_per_q;
int32_t* q2k_block_sparse_num_ptr = g.q2k_block_sparse_num + blockIdx.z * gridDim.y * gridDim.x + blockIdx.y * gridDim.x + blockIdx.x;
int32_t kv_blocks = q2k_block_sparse_num_ptr[0] / 2; // each iter load 2 kv blocks
__shared__ kittens::semaphore qsmem_semaphore, k_smem_arrived, v_smem_arrived;
if (threadIdx.x == 0) {
int32_t kv_block_index[2];
reinterpret_cast<float2*>(kv_block_index)[0] = reinterpret_cast<float2*>(q2k_block_sparse_index_ptr)[0];
init_semaphore(qsmem_semaphore, 0, 1);
init_semaphore(k_smem_arrived, 0, 1);
init_semaphore(v_smem_arrived, 0, 1);
// preload q block
coord<q_tile> q_tile_idx = {blockIdx.z, blockIdx.y, seq_idx, 0};
tma::expect_bytes(qsmem_semaphore, sizeof(q_smem));
tma::load_async(q_smem[0], g.q, q_tile_idx, qsmem_semaphore);
// preload the zeroth block of kv
tma::expect_bytes(k_smem_arrived, sizeof(k_tile));
coord<k_tile_half> k_tile_idx_0 = {blockIdx.z, kv_head_idx, kv_block_index[0], 0};
coord<k_tile_half> k_tile_idx_1 = {blockIdx.z, kv_head_idx, kv_block_index[1], 0};
tma::load_async(k_smem_0[0], g.k, k_tile_idx_0, k_smem_arrived);
tma::load_async(k_smem_1[0], g.k, k_tile_idx_1, k_smem_arrived);
tma::expect_bytes(v_smem_arrived, sizeof(v_tile));
coord<v_tile_half> v_tile_idx_0 = {blockIdx.z, kv_head_idx, kv_block_index[0], 0};
coord<v_tile_half> v_tile_idx_1 = {blockIdx.z, kv_head_idx, kv_block_index[1], 0};
tma::load_async(v_smem_0[0], g.v, v_tile_idx_0, v_smem_arrived);
tma::load_async(v_smem_1[0], g.v, v_tile_idx_1, v_smem_arrived);
}
__syncthreads();
rt_fl<16, 128> att_block;
rt_bf<16, 128> att_block_mma;
rt_fl<16, K::tile_width> o_reg;
col_vec<rt_fl<16, 128>> max_vec, norm_vec, max_vec_last_scaled, max_vec_scaled;
neg_infty(max_vec);
zero(norm_vec);
zero(o_reg);
// wait for q block
wait(qsmem_semaphore, 0);
for (int kv_idx = 0; kv_idx < kv_blocks - 1; kv_idx++) {
// preload kv index
int32_t kv_block_index[2];
reinterpret_cast<float2*>(kv_block_index)[0] = reinterpret_cast<float2*>(q2k_block_sparse_index_ptr)[kv_idx + 1];
// wait k
wait(k_smem_arrived, kv_idx % 2);
// compute QK^T
warpgroup::mm_ABt(att_block, q_smem[0], k_smem[0]);
copy(max_vec_last_scaled, max_vec);
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
warpgroup::mma_async_wait();
// load K
if (threadIdx.x == 0) {
tma::expect_bytes(k_smem_arrived, sizeof(k_tile));
coord<k_tile_half> k_tile_idx_0 = {blockIdx.z, kv_head_idx, kv_block_index[0], 0};
coord<k_tile_half> k_tile_idx_1 = {blockIdx.z, kv_head_idx, kv_block_index[1], 0};
tma::load_async(k_smem_0[0], g.k, k_tile_idx_0, k_smem_arrived);
tma::load_async(k_smem_1[0], g.k, k_tile_idx_1, k_smem_arrived);
}
// exp
row_max(max_vec, att_block, max_vec);
if constexpr (D == 64) {
mul(att_block, att_block, 1.44269504089f*0.125f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
}
else {
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
}
sub_row(att_block, att_block, max_vec_scaled);
exp2(att_block, att_block);
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
exp2(max_vec_last_scaled, max_vec_last_scaled);
mul(norm_vec, norm_vec, max_vec_last_scaled);
row_sum(norm_vec, att_block, norm_vec);
add(att_block, att_block, 0.f);
copy(att_block_mma, att_block);
mul_row(o_reg, o_reg, max_vec_last_scaled);
// wait v
wait(v_smem_arrived, kv_idx % 2);
// compute SV
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[0]);
warpgroup::mma_async_wait();
// load V
if (threadIdx.x == 0) {
tma::expect_bytes(v_smem_arrived, sizeof(v_tile));
// coord<v_tile> v_tile_idx = {blockIdx.z, kv_head_idx, kv_idx + 1, 0};
// tma::load_async(v_smem[0], g.v, v_tile_idx, v_smem_arrived);
coord<v_tile_half> v_tile_idx_0 = {blockIdx.z, kv_head_idx, kv_block_index[0], 0};
coord<v_tile_half> v_tile_idx_1 = {blockIdx.z, kv_head_idx, kv_block_index[1], 0};
tma::load_async(v_smem_0[0], g.v, v_tile_idx_0, v_smem_arrived);
tma::load_async(v_smem_1[0], g.v, v_tile_idx_1, v_smem_arrived);
}
}
// last iter
{
int kv_idx = kv_blocks - 1;
// wait k
wait(k_smem_arrived, kv_idx % 2);
// compute QK^T
warpgroup::mm_ABt(att_block, q_smem[0], k_smem[0]);
copy(max_vec_last_scaled, max_vec);
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
warpgroup::mma_async_wait();
// exp
row_max(max_vec, att_block, max_vec);
if constexpr (D == 64) {
mul(att_block, att_block, 1.44269504089f*0.125f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
}
else {
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
}
sub_row(att_block, att_block, max_vec_scaled);
exp2(att_block, att_block);
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
exp2(max_vec_last_scaled, max_vec_last_scaled);
mul(norm_vec, norm_vec, max_vec_last_scaled);
row_sum(norm_vec, att_block, norm_vec);
add(att_block, att_block, 0.f);
copy(att_block_mma, att_block);
mul_row(o_reg, o_reg, max_vec_last_scaled);
// wait v
wait(v_smem_arrived, kv_idx % 2);
// compute SV
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[0]);
warpgroup::mma_async_wait();
}
div_row(o_reg, o_reg, norm_vec);
warpgroup::store(o_smem[0], o_reg);
__syncthreads();
// TK store_async internally calls syncwarp so we need to route on warp level
if (threadIdx.x / 32 == 0) {
coord<o_tile> o_tile_idx = {blockIdx.z, blockIdx.y, seq_idx, 0};
tma::store_async(g.o, o_smem[0], o_tile_idx);
}
mul(max_vec_scaled, max_vec_scaled, 0.69314718056f);
log(norm_vec, norm_vec);
add(norm_vec, norm_vec, max_vec_scaled);
if constexpr (D == 64) { mul(norm_vec, norm_vec, -8.0f); }
else { mul(norm_vec, norm_vec, -11.313708499f); }
warpgroup::store(l_smem[0], norm_vec);
__syncthreads();
if (threadIdx.x / 32 == 0) {
coord<l_col_vec> tile_idx = {blockIdx.z, blockIdx.y, 0, seq_idx};
tma::store_async(g.l, l_smem[0], tile_idx);
}
tma::store_async_wait();
}
template<int D>
__global__ __launch_bounds__(128, 4)
@@ -357,6 +141,7 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) { // use block siz
}
// exp
right_fill(att_block, att_block, g.block_size[q2k_block_sparse_index_ptr[kv_idx]], base_types::constants<float>::neg_infty());
row_max(max_vec, att_block, max_vec);
if constexpr (D == 64) {
@@ -409,6 +194,8 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) { // use block siz
warpgroup::mma_async_wait();
// exp
right_fill(att_block, att_block, g.block_size[q2k_block_sparse_index_ptr[kv_idx]], base_types::constants<float>::neg_infty());
row_max(max_vec, att_block, max_vec);
if constexpr (D == 64) {
@@ -484,8 +271,9 @@ struct bwd_prep_globals {
d_gl d;
};
constexpr int PREP_NUM_WARPS = (1);
template<int D>
__global__ __launch_bounds__(4*kittens::WARP_THREADS, (D == 64) ? 2 : 1)
__global__ __launch_bounds__(PREP_NUM_WARPS*kittens::WARP_THREADS, (D == 64) ? 6 / PREP_NUM_WARPS : 3 / PREP_NUM_WARPS)
void bwd_attend_prep_ker(const __grid_constant__ bwd_prep_globals<D> g) {
extern __shared__ int __shm[];
tma_swizzle_allocator al((int*)&__shm[0]);
@@ -496,9 +284,9 @@ void bwd_attend_prep_ker(const __grid_constant__ bwd_prep_globals<D> g) {
using o_tile = st_bf<4*16, D>;
using d_tile = col_vec<st_fl<4*16, D>>;
og_tile (&og_smem)[4] = al.allocate<og_tile, 4>();
o_tile (&o_smem) [4] = al.allocate<o_tile , 4>();
d_tile (&d_smem) [4] = al.allocate<d_tile , 4>();
og_tile (&og_smem)[PREP_NUM_WARPS] = al.allocate<og_tile, PREP_NUM_WARPS>();
o_tile (&o_smem) [PREP_NUM_WARPS] = al.allocate<o_tile , PREP_NUM_WARPS>();
d_tile (&d_smem) [PREP_NUM_WARPS] = al.allocate<d_tile , PREP_NUM_WARPS>();
rt_fl<4*16, D> og_reg, o_reg;
col_vec<rt_fl<4*16, D>> d_reg;
@@ -507,13 +295,13 @@ void bwd_attend_prep_ker(const __grid_constant__ bwd_prep_globals<D> g) {
if (threadIdx.x == 0) {
init_semaphore(smem_semaphore, 0, 1);
tma::expect_bytes(smem_semaphore, sizeof(og_smem[0]) * 4 * 2);
tma::expect_bytes(smem_semaphore, sizeof(og_smem[0]) * PREP_NUM_WARPS * 2);
}
__syncthreads();
if (warpid == 0) {
for (int w = 0; w < 4; w++) {
coord<o_tile> tile_idx = {blockIdx.z, blockIdx.y, (blockIdx.x * 4) + w, 0};
for (int w = 0; w < PREP_NUM_WARPS; w++) {
coord<o_tile> tile_idx = {blockIdx.z, blockIdx.y, (blockIdx.x * PREP_NUM_WARPS) + w, 0};
tma::load_async(o_smem[w], g.o, tile_idx, smem_semaphore);
tma::load_async(og_smem[w], g.og, tile_idx, smem_semaphore);
}
@@ -528,8 +316,8 @@ void bwd_attend_prep_ker(const __grid_constant__ bwd_prep_globals<D> g) {
__syncthreads();
if (warpid == 0) {
for (int w = 0; w < 4; w++) {
coord<d_tile> tile_idx = {blockIdx.z, blockIdx.y, 0, (blockIdx.x * 4) + w};
for (int w = 0; w < PREP_NUM_WARPS; w++) {
coord<d_tile> tile_idx = {blockIdx.z, blockIdx.y, 0, (blockIdx.x * PREP_NUM_WARPS) + w};
tma::store_async(g.d, d_smem[w], tile_idx);
}
}
@@ -591,6 +379,7 @@ struct bwd_globals {
int32_t *__restrict__ k2q_block_sparse_index;
int32_t *__restrict__ k2q_block_sparse_num;
int32_t *__restrict__ block_size;
};
__device__ static inline void
@@ -713,7 +502,7 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
// wait for kv
wait(kv_b, 0);
int fill_start = g.block_size[blockIdx.x] - 16 * kittens::warpid();
for (int qo_idx = 0; qo_idx < qo_blocks - 1; qo_idx++) {
// preload q index
store_qg_block_index = load_q_block_index;
@@ -732,7 +521,8 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
if constexpr (D == 64) { mul(s_block_t, s_block_t, 1.44269504089f*0.125f); }
else { mul(s_block_t, s_block_t, 1.44269504089f*0.08838834764f); }
lower_fill(s_block_t, s_block_t, fill_start, base_types::constants<float>::neg_infty());
exp2(s_block_t, s_block_t); // P_i
copy(p_block_t, s_block_t);
copy(p_block_t_mma, s_block_t);
@@ -778,7 +568,6 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
__syncthreads(); // wait for sd_smem shared memory write
warpgroup::mm_AtB(qg_reg, ds_smem_t[0], k_smem[0]); //delat dQ = dSK
warpgroup::mma_commit_group();
tma::store_async_wait();
warpgroup::mma_async_wait();
// store qg to shared memory
warpgroup::store(qg_smem, qg_reg);
@@ -788,6 +577,7 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
if (threadIdx.x / 32 == 0) {
coord<qg_tile> tile_idx = {blockIdx.z, blockIdx.y, store_qg_block_index, 0};
tma::store_add_async(g.qg, qg_smem, tile_idx);
tma::store_async_wait();
}
}
@@ -810,7 +600,7 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
if constexpr (D == 64) { mul(s_block_t, s_block_t, 1.44269504089f*0.125f); }
else { mul(s_block_t, s_block_t, 1.44269504089f*0.08838834764f); }
lower_fill(s_block_t, s_block_t, fill_start, base_types::constants<float>::neg_infty());
exp2(s_block_t, s_block_t); // P_i
copy(p_block_t, s_block_t);
copy(p_block_t_mma, s_block_t);
@@ -834,7 +624,6 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
__syncthreads(); // wait for sd_smem shared memory write
warpgroup::mm_AtB(qg_reg, ds_smem_t[0], k_smem[0]); //delat dQ = dSK
warpgroup::mma_commit_group();
tma::store_async_wait();
warpgroup::mma_async_wait();
// store qg to shared memory
warpgroup::store(qg_smem, qg_reg);
@@ -844,13 +633,14 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
if (threadIdx.x / 32 == 0) {
coord<qg_tile> tile_idx = {blockIdx.z, blockIdx.y, store_qg_block_index, 0};
tma::store_add_async(g.qg, qg_smem, tile_idx);
tma::store_async_wait();
}
}
// store kq and vq
// ! the following two line seems unnecessary.
tma::store_async_wait(); // ensure qg is finished
// tma::store_async_wait(); // ensure qg is finished
__syncthreads();
warpgroup::store(kg_smem[0], kg_reg);
@@ -876,7 +666,14 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
#include <iostream>
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)
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,
torch::Tensor block_size
)
{
CHECK_INPUT(q);
CHECK_INPUT(k);
@@ -888,6 +685,10 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
auto qo_heads = q.size(1);
auto kv_heads = k.size(1);
auto max_kv_blocks_per_q = q2k_block_sparse_index.size(3);
auto num_q_blocks = block_size.size(0);
TORCH_CHECK(batch==1, "Batch size dim will be removed in the future, please set batch to 1");
TORCH_CHECK(num_q_blocks * 64 == seq_len, "This kernel supports variable block size, but it assumes the input sequence is properly padded.");
TORCH_CHECK(num_q_blocks == q2k_block_sparse_index.size(2), "Number of Q blocks does not match between q2k_block_sparse_index and block_size");
// check to see that these dimensions match for all inputs
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
@@ -967,7 +768,19 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(hr), static_cast<int>(max_kv_blocks_per_q), reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()), reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr())};
globals g{
qg_arg,
kg_arg,
vg_arg,
lg_arg,
og_arg,
static_cast<int>(seq_len),
static_cast<int>(hr),
static_cast<int>(max_kv_blocks_per_q),
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
reinterpret_cast<int32_t*>(block_size.data_ptr())
};
constexpr int mem_size = 54000;
@@ -1006,7 +819,19 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(hr), static_cast<int>(max_kv_blocks_per_q), reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()), reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr())};
globals g{
qg_arg,
kg_arg,
vg_arg,
lg_arg,
og_arg,
static_cast<int>(seq_len),
static_cast<int>(hr),
static_cast<int>(max_kv_blocks_per_q),
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
reinterpret_cast<int32_t*>(block_size.data_ptr())
};
constexpr int mem_size = 54000;
@@ -1036,7 +861,8 @@ block_sparse_attention_backward(torch::Tensor q,
torch::Tensor l_vec,
torch::Tensor og,
torch::Tensor k2q_block_sparse_index,
torch::Tensor k2q_block_sparse_num)
torch::Tensor k2q_block_sparse_num,
torch::Tensor block_size)
{
CHECK_INPUT(q);
CHECK_INPUT(k);
@@ -1049,7 +875,7 @@ block_sparse_attention_backward(torch::Tensor q,
auto seq_len = q.size(2);
auto head_dim = q.size(3);
auto max_q_blocks_per_kv = k2q_block_sparse_index.size(3);
TORCH_CHECK(k2q_block_sparse_index.size(2) == block_size.size(0), "k2q_block_sparse_index.size(2) must match block_size.size(0)");
// check to see that these dimensions match for all inputs
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
@@ -1136,7 +962,7 @@ block_sparse_attention_backward(torch::Tensor q,
float* d_vg = reinterpret_cast<float*>(vg_ptr);
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
int threads = 4 * kittens::WARP_THREADS;
int threads = PREP_NUM_WARPS * kittens::WARP_THREADS;
//cudadevicesynchronize();
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
@@ -1145,7 +971,7 @@ block_sparse_attention_backward(torch::Tensor q,
// cudaStreamSynchronize(stream);
// TORCH_CHECK(seq_len % (4*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 256");
dim3 grid_bwd(seq_len/(4*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
dim3 grid_bwd(seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
if (head_dim == 64) {
using og_tile = st_bf<4*16, 64>;
@@ -1220,8 +1046,8 @@ block_sparse_attention_backward(torch::Tensor q,
static_cast<int>(hr),
static_cast<int>(max_q_blocks_per_kv),
reinterpret_cast<int32_t*>(k2q_block_sparse_index.data_ptr()),
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr())
};
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr()),
reinterpret_cast<int32_t*>(block_size.data_ptr())};
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
threads = 128;
@@ -1324,8 +1150,8 @@ block_sparse_attention_backward(torch::Tensor q,
static_cast<int>(hr),
static_cast<int>(max_q_blocks_per_kv),
reinterpret_cast<int32_t*>(k2q_block_sparse_index.data_ptr()),
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr())
};
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr()),
reinterpret_cast<int32_t*>(block_size.data_ptr())};
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
threads = 128;
@@ -1348,4 +1174,4 @@ block_sparse_attention_backward(torch::Tensor q,
return {qg, kg, vg};
//cudadevicesynchronize();
}
}
@@ -0,0 +1,185 @@
import torch
try:
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
except ImportError:
block_sparse_fwd = None
block_sparse_bwd = None
from vsa.block_sparse_attn_triton import triton_block_sparse_attn_forward, triton_block_sparse_attn_backward
assert torch.__version__ >= "2.4.0", "VSA requires PyTorch 2.4.0 or higher"
from vsa.index import map_to_index
from typing import Tuple, Optional
@torch.library.custom_op("vsa::block_sparse_attn_triton", mutates_args=(), device_types="cuda")
def block_sparse_attn_triton(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
block_map = block_map.int()
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
o, M = triton_block_sparse_attn_forward(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
return o, M
@torch.library.register_fake("vsa::block_sparse_attn_triton")
def _block_sparse_attn_triton_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
o = torch.empty_like(q)
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
return o, M
@torch.library.custom_op("vsa::block_sparse_attn_backward_triton", mutates_args=(), device_types="cuda")
def block_sparse_attn_backward_triton(
grad_output_padded: torch.Tensor,
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
o_padded: torch.Tensor,
M: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
grad_output_padded = grad_output_padded.contiguous()
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(block_map.transpose(-1, -2))
dq, dk, dv = triton_block_sparse_attn_backward(grad_output_padded, q_padded, k_padded, v_padded, o_padded, M, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes)
return dq, dk, dv
@torch.library.register_fake("vsa::block_sparse_attn_backward_triton")
def _block_sparse_attn_backward_triton_fake(
grad_output_padded: torch.Tensor,
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
o_padded: torch.Tensor,
M: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
grad_output_padded = grad_output_padded.contiguous()
dq = torch.empty_like(grad_output_padded)
dk = torch.empty_like(grad_output_padded)
dv = torch.empty_like(grad_output_padded)
return dq, dk, dv
def backward_triton(ctx, grad_output1, grad_output2):
q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_triton(grad_output1, q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes)
return dq, dk, dv, None, None
def setup_context_triton(ctx, inputs, output):
q_padded, k_padded, v_padded, block_map, variable_block_sizes = inputs
o_padded, M = output
ctx.save_for_backward(q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes)
block_sparse_attn_triton.register_autograd(backward_triton, setup_context=setup_context_triton)
major, minor = torch.cuda.get_device_capability(0)
if major == 9 and minor == 0:# check if H100
@torch.library.custom_op("vsa::block_sparse_attn_SM90", mutates_args=(), device_types="cuda")
def block_sparse_attn_SM90(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
)-> Tuple[torch.Tensor, torch.Tensor]:
q_padded = q_padded.contiguous()
k_padded = k_padded.contiguous()
v_padded = v_padded.contiguous()
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
variable_block_sizes = variable_block_sizes.int()
o_padded, lse_padded = block_sparse_fwd(q_padded, k_padded, v_padded, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
return o_padded, lse_padded
@torch.library.register_fake("vsa::block_sparse_attn_SM90")
def _block_sparse_attn_SM90_fake(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
q_padded, k_padded, v_padded = [x.contiguous() for x in (q_padded, k_padded, v_padded)]
B, H, S, D = q_padded.shape
o_padded = torch.empty_like(q_padded)
lse_padded = torch.empty((B, H, S, 1), device=q_padded.device, dtype=torch.float32)
return o_padded, lse_padded
@torch.library.custom_op("vsa::block_sparse_attn_backward_SM90", mutates_args=(), device_types="cuda")
def block_sparse_attn_backward_SM90(
grad_output_padded: torch.Tensor,
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
)-> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
grad_output_padded = grad_output_padded.contiguous()
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(block_map.transpose(-1, -2))
grad_q_padded, grad_k_padded, grad_v_padded = block_sparse_bwd(
q_padded, k_padded, v_padded, o_padded, lse_padded, grad_output_padded, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes
)
grad_q_padded = grad_q_padded.to(grad_output_padded.dtype)
grad_k_padded = grad_k_padded.to(grad_output_padded.dtype)
grad_v_padded = grad_v_padded.to(grad_output_padded.dtype)
return grad_q_padded, grad_k_padded, grad_v_padded
@torch.library.register_fake("vsa::block_sparse_attn_backward_SM90")
def _block_sparse_attn_backward_SM90_fake(
grad_output_padded: torch.Tensor,
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
torch._check(grad_output_padded.dtype == torch.bfloat16)
torch._check(lse_padded.dtype == torch.float32)
grad_output_padded = grad_output_padded.contiguous()
dq = torch.empty_like(grad_output_padded)
dk = torch.empty_like(grad_output_padded)
dv = torch.empty_like(grad_output_padded)
return dq, dk, dv
def backward_SM90(ctx, grad_output1, grad_output2):
q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes= ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_SM90(grad_output1, q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
return dq, dk, dv, None, None
def setup_context_SM90(ctx, inputs, output):
q_padded, k_padded, v_padded, block_map, variable_block_sizes = inputs
o_padded, lse_padded = output
ctx.save_for_backward(q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
block_sparse_attn_SM90.register_autograd(backward_SM90, setup_context=setup_context_SM90)
+152
View File
@@ -0,0 +1,152 @@
## pytorch sdpa version of block sparse ##
import triton
import triton.language as tl
import torch
@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,
):
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,
):
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.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
+32
View File
@@ -0,0 +1,32 @@
# Attention Kernel Used in FastVideo
## VMoBA: Mixture-of-Block Attention for Video Diffusion Models (VMoBA)
### Installation
Please ensure that you have installed FlashAttention version **2.7.1 or higher**, as some interfaces have changed in recent releases.
### Usage
You can use `moba_attn_varlen` in the following ways:
**Install from source:**
```bash
python setup.py install
```
**Import after installation:**
```python
from vmoba import moba_attn_varlen
```
**Or import directly from the project root:**
```python
from csrc.attn.vmoba_attn.vmoba import moba_attn_varlen
```
### Verify if you have successfully installed
```bash
python csrc/attn/vmoba_attn/vmoba/vmoba.py
```
+24
View File
@@ -0,0 +1,24 @@
# SPDX-License-Identifier: Apache-2.0
from setuptools import find_packages, setup
PACKAGE_NAME = "vmoba"
VERSION = "0.0.0"
AUTHOR = "JianzongWu"
DESCRIPTION = "VMoBA: Mixture-of-Block Attention for Video Diffusion Models"
URL = "https://github.com/KwaiVGI/VMoBA"
setup(
name=PACKAGE_NAME,
version=VERSION,
author=AUTHOR,
description=DESCRIPTION,
url=URL,
packages=find_packages(),
classifiers=[
"Programming Language :: Python :: 3",
"License :: OSI Approved :: Apache Software License",
],
python_requires='>=3.12',
install_requires=[]
)
@@ -0,0 +1,97 @@
# SPDX-License-Identifier: Apache-2.0
import torch
import pytest
import random
from csrc.attn.vmoba_attn.vmoba import moba_attn_varlen
def generate_test_data(batch_size, total_seqlen, num_heads, head_dim, dtype, device="cuda"):
"""
Generates random data for testing the variable-length attention function.
"""
torch.manual_seed(42)
random.seed(42)
torch.cuda.manual_seed_all(42)
# Generate sequence lengths for each item in the batch
if batch_size > 1:
# Ensure sequence lengths are reasonably distributed
avg_seqlen = total_seqlen // batch_size
seqlens = [random.randint(avg_seqlen // 2, avg_seqlen + avg_seqlen // 2) for _ in range(batch_size - 1)]
remaining_len = total_seqlen - sum(seqlens)
if remaining_len > 0:
seqlens.append(remaining_len)
else: # Adjust if sum exceeds total_seqlen
seqlens.append(avg_seqlen)
current_sum = sum(seqlens)
seqlens[-1] -= (current_sum - total_seqlen)
# Ensure all lengths are positive
seqlens = [max(1, s) for s in seqlens]
# Final adjustment to match total_seqlen
seqlens[-1] += total_seqlen - sum(seqlens)
else:
seqlens = [total_seqlen]
cu_seqlens = torch.tensor([0] + list(torch.cumsum(torch.tensor(seqlens), 0)), device=device, dtype=torch.int32)
max_seqlen = max(seqlens) if seqlens else 0
q = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
k = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
v = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
return q, k, v, cu_seqlens, max_seqlen
@pytest.mark.parametrize("batch_size", [1, 2])
@pytest.mark.parametrize("total_seqlen", [512, 1024])
@pytest.mark.parametrize("num_heads", [8])
@pytest.mark.parametrize("head_dim", [64])
@pytest.mark.parametrize("moba_chunk_size", [64])
@pytest.mark.parametrize("moba_topk", [2, 4])
@pytest.mark.parametrize("select_mode", ["topk", "threshold"])
@pytest.mark.parametrize("threshold_type", ["query_head", "head_global", "overall"])
@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16])
def test_moba_attn_varlen_forward(
batch_size, total_seqlen, num_heads, head_dim, moba_chunk_size, moba_topk, select_mode, threshold_type, dtype
):
"""
Tests the forward pass of moba_attn_varlen for basic correctness.
It checks output shape, dtype, and for the presence of NaNs/Infs.
"""
if dtype == torch.float32:
pytest.skip("float32 is not supported in flash attention")
q, k, v, cu_seqlens, max_seqlen = generate_test_data(
batch_size, total_seqlen, num_heads, head_dim, dtype
)
# Ensure chunk size is not larger than the smallest sequence length
min_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).min().item()
if moba_chunk_size > min_seqlen:
pytest.skip("moba_chunk_size is larger than the minimum sequence length in the batch")
try:
output = moba_attn_varlen(
q=q,
k=k,
v=v,
cu_seqlens=cu_seqlens,
max_seqlen=max_seqlen,
moba_chunk_size=moba_chunk_size,
moba_topk=moba_topk,
select_mode=select_mode,
threshold_type=threshold_type,
simsum_threshold=0.5, # A reasonable default for threshold mode
)
except Exception as e:
pytest.fail(f"moba_attn_varlen forward pass failed with exception: {e}")
# 1. Check output shape
assert output.shape == q.shape, f"Expected output shape {q.shape}, but got {output.shape}"
# 2. Check output dtype
assert output.dtype == q.dtype, f"Expected output dtype {q.dtype}, but got {output.dtype}"
# 3. Check for NaNs or Infs in the output
assert torch.all(torch.isfinite(output)), "Output contains NaN or Inf values"
+2
View File
@@ -0,0 +1,2 @@
# SPDX-License-Identifier: Apache-2.0
from .vmoba import moba_attn_varlen, process_moba_input, process_moba_output
+860
View File
@@ -0,0 +1,860 @@
# SPDX-License-Identifier: Apache-2.0
# Adapt from https://github.com/KwaiVGI/VMoBA/blob/main/src/vmoba.py
import random
import time
import os
import torch
from typing import Tuple
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
from functools import lru_cache
from einops import rearrange
@lru_cache(maxsize=16)
def calc_chunks(cu_seqlen, moba_chunk_size):
"""
Calculate chunk boundaries.
For vision tasks we include all chunks (even the last one which might be shorter)
so that every chunk can be selected.
"""
batch_sizes = cu_seqlen[1:] - cu_seqlen[:-1]
batch_num_chunk = (batch_sizes + (moba_chunk_size - 1)) // moba_chunk_size
cu_num_chunk = torch.ones(
batch_num_chunk.numel() + 1,
device=cu_seqlen.device,
dtype=batch_num_chunk.dtype,
)
cu_num_chunk[1:] = batch_num_chunk.cumsum(dim=0)
num_chunk = cu_num_chunk[-1]
chunk_sizes = torch.full(
(num_chunk + 1,), moba_chunk_size, dtype=torch.int32, device=cu_seqlen.device
)
chunk_sizes[0] = 0
batch_last_chunk_size = batch_sizes - (batch_num_chunk - 1) * moba_chunk_size
chunk_sizes[cu_num_chunk[1:]] = batch_last_chunk_size
cu_chunk = chunk_sizes.cumsum(dim=-1, dtype=torch.int32)
chunk_to_batch = torch.zeros(
(num_chunk,), dtype=torch.int32, device=cu_seqlen.device
)
chunk_to_batch[cu_num_chunk[1:-1]] = 1
chunk_to_batch = chunk_to_batch.cumsum(dim=0, dtype=torch.int32)
# Do not filter out any chunk
filtered_chunk_indices = torch.arange(
num_chunk, device=cu_seqlen.device, dtype=torch.int32
)
num_filtered_chunk = num_chunk
return cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch
# --- Threshold Selection Helper Functions ---
def _select_threshold_query_head(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects chunks for each <query, head> pair based on threshold.
Normalization and sorting happen along the chunk dimension (dim=0).
"""
C, H, S = gate.shape
eps = 1e-6
# LSE‐style normalization per <head, query> (across chunks)
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
row_min = gate_min_val.amin(dim=0) # (H, S)
row_max = gate_masked.amax(dim=0) # (H, S)
denom = row_max - row_min
denom = torch.where(denom <= eps, torch.ones_like(denom), denom) # avoid divide‑by‑zero
gate_norm = (gate - row_min.unsqueeze(0)) / denom.unsqueeze(0)
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 1) pull out the self‐chunk’s normalized weight for each <head,seq>
self_norm = (gate_norm * gate_self_chunk_mask).sum(dim=0) # (H, S)
# 2) compute how much more normalized weight we need beyond self
total_norm_sum = gate_norm.sum(dim=0) # (H, S)
remain_ratio = simsum_threshold - self_norm / (total_norm_sum + eps) # (H, S)
remain_ratio = torch.clamp(remain_ratio, min=0.0) # if already ≥ thresh, no extra needed
# 3) zero out the self‐chunk in a copy, so we only sort “others”
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0
# 4) sort the other chunks by descending norm, per <head,seq>
sorted_norm, sorted_idx = torch.sort(others_norm, descending=True, dim=0) # (C, H, S)
# 5) cumulative‑sum the sorted norms per <head,seq>
cumsum_others = sorted_norm.cumsum(dim=0) # (C, H, S)
# 6) for each <head,seq>, find the smallest k where cumsum_ratio ≥ remain_ratio
ratio = cumsum_others / (total_norm_sum.unsqueeze(0) + eps) # (C, H, S)
cond = ratio >= remain_ratio.unsqueeze(0) # (C, H, S) boolean mask
any_cond = cond.any(dim=0) # (H, S)
# Find the index of the first True value along dim 0. If none, use C-1.
cutoff = torch.where(any_cond, cond.float().argmax(dim=0), torch.full_like(any_cond, fill_value=C - 1)) # (H, S)
# 7) build a mask in sorted order up to that cutoff
idx_range = torch.arange(C, device=gate.device).view(-1, 1, 1) # (C, 1, 1)
sorted_mask = idx_range <= cutoff.unsqueeze(0) # (C, H, S)
# 8) scatter it back to original chunk order
others_mask = torch.zeros_like(gate, dtype=torch.bool)
others_mask.scatter_(0, sorted_idx, sorted_mask)
# 9) finally, include every self‐chunk plus all selected others
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
return final_gate_mask
def _select_threshold_block(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects <query, head> pairs for each block based on threshold.
Normalization and sorting happen across the head and sequence dimensions (dim=1, 2).
"""
C, H, S = gate.shape
HS = H * S
eps = 1e-6
# LSE‐style normalization per block (across heads and queries)
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
block_max = gate_masked.amax(dim=(1, 2), keepdim=True) # (C, 1, 1)
block_min = gate_min_val.amin(dim=(1, 2), keepdim=True) # (C, 1, 1)
block_denom = block_max - block_min
block_denom = torch.where(block_denom <= eps, torch.ones_like(block_denom), block_denom) # (C, 1, 1)
gate_norm = (gate - block_min) / block_denom # (C, H, S)
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 1) identify normalized weights of entries that *are* self-chunks (from query perspective)
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
# Sum these weights *per block*
self_norm_sum_per_block = self_norm_entries.sum(dim=(1, 2)) # (C,)
# 2) compute how much more normalized weight each block needs beyond its self-chunk contributions
total_norm_sum_per_block = gate_norm.sum(dim=(1, 2)) # (C,)
remain_ratio = simsum_threshold - self_norm_sum_per_block / (total_norm_sum_per_block + eps) # (C,)
remain_ratio = torch.clamp(remain_ratio, min=0.0) # (C,)
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
# 4) sort the other <head, seq> pairs by descending norm, per block
others_flat = others_norm.contiguous().view(C, HS) # (C, H*S)
sorted_others_flat, sorted_indices_flat = torch.sort(others_flat, dim=1, descending=True) # (C, H*S)
# 5) cumulative‑sum the sorted norms per block
cumsum_others_flat = sorted_others_flat.cumsum(dim=1) # (C, H*S)
# 6) for each block, find the smallest k where cumsum_ratio ≥ remain_ratio
ratio_flat = cumsum_others_flat / (total_norm_sum_per_block.unsqueeze(1) + eps) # (C, H*S)
cond_flat = ratio_flat >= remain_ratio.unsqueeze(1) # (C, H*S) boolean mask
any_cond = cond_flat.any(dim=1) # (C,)
# Find the index of the first True value along dim 1. If none, use HS-1.
cutoff_flat = torch.where(any_cond, cond_flat.float().argmax(dim=1), torch.full_like(any_cond, fill_value=HS - 1)) # (C,)
# 7) build a mask in sorted order up to that cutoff per block
idx_range_flat = torch.arange(HS, device=gate.device).unsqueeze(0) # (1, H*S)
sorted_mask_flat = idx_range_flat <= cutoff_flat.unsqueeze(1) # (C, H*S)
# 8) scatter it back to original <head, seq> order per block
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C, H*S)
others_mask_flat.scatter_(1, sorted_indices_flat, sorted_mask_flat)
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
# 9) finally, include every self‐chunk entry plus all selected others
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
return final_gate_mask
def _select_threshold_overall(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects <chunk, query, head> triplets globally based on threshold.
Normalization and sorting happen across all valid entries.
"""
C, H, S = gate.shape
CHS = C * H * S
eps = 1e-6
# LSE‐style normalization globally across all valid entries
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
overall_max = gate_masked.max() # scalar
overall_min = gate_min_val.min() # scalar
overall_denom = overall_max - overall_min
overall_denom = torch.where(overall_denom <= eps, torch.tensor(1.0, device=gate.device, dtype=gate.dtype), overall_denom)
gate_norm = (gate - overall_min) / overall_denom # (C, H, S)
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 1) identify normalized weights of entries that *are* self-chunks
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
# Sum these weights globally
self_norm_sum_overall = self_norm_entries.sum() # scalar
# 2) compute how much more normalized weight is needed globally beyond self-chunk contributions
total_norm_sum_overall = gate_norm.sum() # scalar
remain_ratio = simsum_threshold - self_norm_sum_overall / (total_norm_sum_overall + eps) # scalar
remain_ratio = torch.clamp(remain_ratio, min=0.0) # scalar
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
# 4) sort all other entries by descending norm, globally
others_flat = others_norm.flatten() # (C*H*S,)
valid_others_mask_flat = valid_gate_mask.flatten() & ~gate_self_chunk_mask.flatten() # Mask for valid, non-self entries
# Only sort the valid 'other' entries
valid_others_indices = torch.where(valid_others_mask_flat)[0]
valid_others_values = others_flat[valid_others_indices]
sorted_others_values, sort_perm = torch.sort(valid_others_values, descending=True) # (N_valid_others,)
sorted_original_indices = valid_others_indices[sort_perm] # Original indices in C*H*S space, sorted by value
# 5) cumulative‑sum the sorted valid 'other' norms globally
cumsum_others_values = sorted_others_values.cumsum(dim=0) # (N_valid_others,)
# 6) find the smallest k where cumsum_ratio ≥ remain_ratio globally
ratio_values = cumsum_others_values / (total_norm_sum_overall + eps) # (N_valid_others,)
cond_values = ratio_values >= remain_ratio # (N_valid_others,) boolean mask
any_cond = cond_values.any() # scalar
# Find the index of the first True value in the *sorted* list. If none, use all valid others.
cutoff_idx_in_sorted = torch.where(
any_cond,
cond_values.float().argmax(dim=0),
torch.tensor(len(sorted_others_values) - 1, device=gate.device, dtype=torch.long)
)
# 7) build a mask selecting the top-k others based on the cutoff
# Select the original indices corresponding to the top entries in the sorted list
selected_other_indices = sorted_original_indices[:cutoff_idx_in_sorted + 1]
# 8) create the mask in the original flat shape
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C*H*S,)
if selected_other_indices.numel() > 0: # Check if any 'other' indices were selected
others_mask_flat[selected_other_indices] = True
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
# 9) finally, include every self‐chunk entry plus all selected others
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
return final_gate_mask
def _select_threshold_head_global(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects <chunk, query> globally for each head based on threshold.
"""
C, H, S = gate.shape
eps = 1e-6
# 1) LSE‐style normalization per head (across chunks and sequence dims)
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf)
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf)
max_per_head = gate_masked.amax(dim=(0, 2), keepdim=True) # (1, H, 1)
min_per_head = gate_min_val.amin(dim=(0, 2), keepdim=True) # (1, H, 1)
denom = max_per_head - min_per_head
denom = torch.where(denom <= eps, torch.ones_like(denom), denom)
gate_norm = (gate - min_per_head) / denom
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 2) sum normalized self‐chunk contributions per head
self_norm_sum = (gate_norm * gate_self_chunk_mask).sum(dim=(0, 2)) # (H,)
# 3) total normalized sum per head
total_norm_sum = gate_norm.sum(dim=(0, 2)) # (H,)
# 4) how much more normalized weight needed per head
remain_ratio = simsum_threshold - self_norm_sum / (total_norm_sum + eps) # (H,)
remain_ratio = torch.clamp(remain_ratio, min=0.0)
# 5) zero out self‐chunk entries to focus on "others"
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0 # (C, H, S)
# 6) flatten chunk and sequence dims, per head
CS = C * S
others_flat = others_norm.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
valid_flat = (valid_gate_mask & ~gate_self_chunk_mask) \
.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
# 7) vectorized selection of “others” per head
masked_flat = torch.where(valid_flat, others_flat, torch.zeros_like(others_flat))
sorted_vals, sorted_idx = torch.sort(masked_flat, dim=1, descending=True) # (H, C*S)
cumsum_vals = sorted_vals.cumsum(dim=1) # (H, C*S)
ratio_vals = cumsum_vals / (total_norm_sum.unsqueeze(1) + eps) # (H, C*S)
cond = ratio_vals >= remain_ratio.unsqueeze(1) # (H, C*S)
has_cutoff = cond.any(dim=1) # (H,)
default = torch.full((H,), CS - 1, device=gate.device, dtype=torch.long)
cutoff = torch.where(has_cutoff, cond.float().argmax(dim=1), default) # (H,)
idx_range = torch.arange(CS, device=gate.device).unsqueeze(0) # (1, C*S)
sorted_mask = idx_range <= cutoff.unsqueeze(1) # (H, C*S)
selected_flat = torch.zeros_like(valid_flat) # (H, C*S)
selected_flat.scatter_(1, sorted_idx, sorted_mask) # (H, C*S)
# 8) reshape selection mask back to (C, H, S)
others_mask = selected_flat.reshape(H, C, S).permute(1, 0, 2) # (C, H, S)
# 9) include self‐chunks plus selected others, and obey valid mask
final_gate_mask = valid_gate_mask & (gate_self_chunk_mask | others_mask)
return final_gate_mask
class MixedAttention(torch.autograd.Function):
@staticmethod
def forward(
ctx,
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
max_seqlen,
moba_chunk_size,
moba_q_sh_indices,
):
ctx.max_seqlen = max_seqlen
ctx.moba_chunk_size = moba_chunk_size
ctx.softmax_scale = softmax_scale = q.shape[-1] ** (-0.5)
# Non-causal self-attention branch
# return out, softmax_lse, S_dmask, rng_state
self_attn_out_sh, self_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
q=q,
k=k,
v=v,
cu_seqlens_q=self_attn_cu_seqlen,
cu_seqlens_k=self_attn_cu_seqlen,
max_seqlen_q=max_seqlen,
max_seqlen_k=max_seqlen,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
)
# MOBA attention branch (non-causal)
moba_attn_out, moba_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
q=moba_q,
k=moba_kv[:, 0],
v=moba_kv[:, 1],
cu_seqlens_q=moba_cu_seqlen_q,
cu_seqlens_k=moba_cu_seqlen_kv,
max_seqlen_q=max_seqlen,
max_seqlen_k=moba_chunk_size,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
)
self_attn_lse_sh = self_attn_lse_hs.t().contiguous()
moba_attn_lse = moba_attn_lse_hs.t().contiguous()
output = torch.zeros((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
output_2d = output.view(-1, q.shape[2])
max_lse_1d = self_attn_lse_sh.view(-1)
max_lse_1d = max_lse_1d.index_reduce(
0, moba_q_sh_indices, moba_attn_lse.view(-1), "amax"
)
self_attn_lse_sh = self_attn_lse_sh - max_lse_1d.view_as(self_attn_lse_sh)
moba_attn_lse = (
moba_attn_lse.view(-1)
.sub(max_lse_1d.index_select(0, moba_q_sh_indices))
.reshape_as(moba_attn_lse)
)
mixed_attn_se_sh = self_attn_lse_sh.exp()
moba_attn_se = moba_attn_lse.exp()
mixed_attn_se_sh.view(-1).index_add_(
0, moba_q_sh_indices, moba_attn_se.view(-1)
)
mixed_attn_lse_sh = mixed_attn_se_sh.log()
# Combine self-attention output
factor = (self_attn_lse_sh - mixed_attn_lse_sh).exp() # [S, H]
self_attn_out_sh = self_attn_out_sh * factor.unsqueeze(-1)
output_2d += self_attn_out_sh.reshape_as(output_2d)
# Combine MOBA attention output
mixed_attn_lse = (
mixed_attn_lse_sh.view(-1)
.index_select(0, moba_q_sh_indices)
.view_as(moba_attn_lse)
)
factor = (moba_attn_lse - mixed_attn_lse).exp() # [S, H]
moba_attn_out = moba_attn_out * factor.unsqueeze(-1)
raw_attn_out = moba_attn_out.view(-1, moba_attn_out.shape[-1])
output_2d.index_add_(0, moba_q_sh_indices, raw_attn_out)
output = output.to(q.dtype)
mixed_attn_lse_sh = mixed_attn_lse_sh + max_lse_1d.view_as(mixed_attn_se_sh)
ctx.save_for_backward(
output,
mixed_attn_lse_sh,
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
moba_q_sh_indices,
)
return output
@staticmethod
def backward(ctx, d_output):
max_seqlen = ctx.max_seqlen
moba_chunk_size = ctx.moba_chunk_size
softmax_scale = ctx.softmax_scale
(
output,
mixed_attn_vlse_sh,
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
moba_q_sh_indices,
) = ctx.saved_tensors
d_output = d_output.contiguous()
dq = torch.empty_like(q)
dk = torch.empty_like(k)
dv = torch.empty_like(v)
_ = _flash_attn_varlen_backward(
dout=d_output,
q=q,
k=k,
v=v,
out=output,
softmax_lse=mixed_attn_vlse_sh.t().contiguous(),
dq=dq,
dk=dk,
dv=dv,
cu_seqlens_q=self_attn_cu_seqlen,
cu_seqlens_k=self_attn_cu_seqlen,
max_seqlen_q=max_seqlen,
max_seqlen_k=max_seqlen,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
softcap=0.0,
alibi_slopes=None,
deterministic=True,
window_size_left=-1,
window_size_right=-1
)
headdim = q.shape[-1]
d_moba_output = (
d_output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
)
moba_output = (
output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
)
mixed_attn_vlse = (
mixed_attn_vlse_sh.view(-1).index_select(0, moba_q_sh_indices).view(1, -1)
)
dmq = torch.empty_like(moba_q)
dmkv = torch.empty_like(moba_kv)
_ = _flash_attn_varlen_backward(
dout=d_moba_output,
q=moba_q,
k=moba_kv[:, 0],
v=moba_kv[:, 1],
out=moba_output,
softmax_lse=mixed_attn_vlse,
dq=dmq,
dk=dmkv[:,0],
dv=dmkv[:,1],
cu_seqlens_q=moba_cu_seqlen_q,
cu_seqlens_k=moba_cu_seqlen_kv,
max_seqlen_q=max_seqlen,
max_seqlen_k=moba_chunk_size,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
softcap=0.0,
alibi_slopes=None,
deterministic=True,
window_size_left=-1,
window_size_right=-1
)
return dq, dk, dv, None, dmq, dmkv, None, None, None, None, None
def moba_attn_varlen(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
cu_seqlens: torch.Tensor,
max_seqlen: int,
moba_chunk_size: int,
moba_topk: int,
select_mode: str = 'threshold', # "topk" or "threshold"
simsum_threshold: float = 0.25,
threshold_type: str = 'query_head',
) -> torch.Tensor:
"""
Accelerated MOBA attention for vision tasks with proper LSE normalization.
This version:
- Splits KV into chunks.
- For each query head, selects the top-k relevant KV chunks (including the self chunk)
by amplifying the diagonal (self-chunk) logits.
- Aggregates the attention outputs from the selected chunks using a log-sum-exp
reduction so that attending to each query over the selected chunks is equivalent
to the original algorithm.
"""
# Stack keys and values.
kv = torch.stack((k, v), dim=1)
seqlen, num_head, head_dim = q.shape
# Compute chunk boundaries.
cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch = calc_chunks(
cu_seqlens, moba_chunk_size
)
self_attn_cu_seqlen = cu_chunk
# Update top-k selection to include the self chunk.
moba_topk = min(moba_topk, num_filtered_chunk)
# --- Build filtered KV from chunks ---
chunk_starts = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
chunk_ends = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
chunk_lengths = chunk_ends - chunk_starts # [num_filtered_chunk]
max_chunk_len = int(chunk_lengths.max().item())
range_tensor = torch.arange(max_chunk_len, device=kv.device, dtype=chunk_starts.dtype).unsqueeze(0)
indices = chunk_starts.unsqueeze(1) + range_tensor
indices = torch.clamp(indices, max=kv.shape[0] - 1)
valid_mask = range_tensor < chunk_lengths.unsqueeze(1)
gathered = kv[indices.view(-1)].view(num_filtered_chunk, max_chunk_len, *kv.shape[1:])
gathered = gathered * valid_mask.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1).type_as(gathered)
# Compute key_gate_weight over valid tokens.
key_values = gathered[:, :, 0].float() # [num_filtered_chunk, max_chunk_len, num_head, head_dim]
valid_mask_exp = valid_mask.unsqueeze(-1).unsqueeze(-1)
key_sum = (key_values * valid_mask_exp).sum(dim=1)
divisor = valid_mask.sum(dim=1).unsqueeze(-1).unsqueeze(-1)
key_gate_weight = key_sum / divisor # [num_filtered_chunk, num_head, head_dim]
# Compute gate logits between key_gate_weight and queries.
q_float = q.float()
# gate = torch.einsum("nhd,shd->nhs", key_gate_weight, q_float) # [num_filtered_chunk, num_head, seqlen]
gate = torch.bmm(key_gate_weight.permute(1, 0, 2), q_float.permute(1, 0, 2).transpose(1, 2)).permute(1, 0, 2)
# Amplify the diagonal (self chunk) contributions.
gate_seq_idx = torch.arange(seqlen, device=q.device, dtype=torch.int32).unsqueeze(0).expand(num_filtered_chunk, seqlen)
chunk_start = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
chunk_end = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
gate_self_chunk_mask = ((gate_seq_idx >= chunk_start.unsqueeze(1)) &
(gate_seq_idx < chunk_end.unsqueeze(1))).unsqueeze(1).expand(-1, num_head, -1)
amplification_factor = 1e9 # Example factor; adjust as needed.
origin_gate = gate.clone()
gate = gate.clone()
if select_mode == "topk":
gate[gate_self_chunk_mask] += amplification_factor
# Exclude positions that are outside the valid batch boundaries.
batch_starts = cu_seqlens[chunk_to_batch[filtered_chunk_indices]]
batch_ends = cu_seqlens[chunk_to_batch[filtered_chunk_indices] + 1]
gate_batch_start_mask = gate_seq_idx < batch_starts.unsqueeze(1)
gate_batch_end_mask = gate_seq_idx >= batch_ends.unsqueeze(1)
gate_inf_mask = gate_batch_start_mask | gate_batch_end_mask
gate.masked_fill_(gate_inf_mask.unsqueeze(1), -float("inf"))
if select_mode == 'topk':
# We amplify self‐chunk in gate already, so self entries will rank highest.
valid_gate_mask = gate != -float("inf")
if threshold_type == 'query_head':
# === per‐<head,seq> top-k across chunks (original behavior) ===
# gate: (C, H, S)
_, gate_topk_idx = torch.topk(gate, k=moba_topk, dim=0, largest=True, sorted=False)
gate_idx_mask = torch.zeros_like(gate, dtype=torch.bool)
gate_idx_mask.scatter_(0, gate_topk_idx, True)
gate_mask = valid_gate_mask & gate_idx_mask
elif threshold_type == 'overall':
# === global top-k across all (chunk, head, seq) entries ===
C, H, S = gate.shape
flat_gate = gate.flatten()
flat_mask = valid_gate_mask.flatten()
flat_gate_masked = torch.where(flat_mask, flat_gate, -float("inf"))
# pick topk global entries
vals, idx = torch.topk(flat_gate_masked, k=moba_topk * H * S, largest=True, sorted=False)
others_mask_flat = torch.zeros_like(flat_mask, dtype=torch.bool)
others_mask_flat[idx] = True
gate_mask = (valid_gate_mask.flatten() & others_mask_flat).view(gate.shape)
elif threshold_type == 'head_global':
# per-head top-k across all chunks and sequence positions
C, H, S = gate.shape
CS = C * S
flat_gate = gate.permute(1, 0, 2).reshape(H, CS)
flat_valid = valid_gate_mask.permute(1, 0, 2).reshape(H, CS)
flat_gate_masked = torch.where(flat_valid, flat_gate, torch.full_like(flat_gate, -float('inf')))
# pick top-k indices per head
_, topk_idx = torch.topk(flat_gate_masked, k=moba_topk * S, dim=1, largest=True, sorted=False)
gate_idx_flat = torch.zeros_like(flat_valid, dtype=torch.bool)
gate_idx_flat.scatter_(1, topk_idx, True)
gate_mask = gate_idx_flat.reshape(H, C, S).permute(1, 0, 2)
else:
raise ValueError(
f"Invalid threshold_type for topk: {threshold_type}. "
"Choose 'query_head', 'block', or 'overall'."
)
elif select_mode == 'threshold':
# Delegate to the specific thresholding function
valid_gate_mask = gate != -float("inf") # (num_chunk, num_head, seqlen)
if threshold_type == 'query_head':
gate_mask = _select_threshold_query_head(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
elif threshold_type == 'block':
gate_mask = _select_threshold_block(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
elif threshold_type == 'overall':
gate_mask = _select_threshold_overall(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
elif threshold_type == 'head_global':
gate_mask = _select_threshold_head_global(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
else:
raise ValueError(f"Invalid threshold_type: {threshold_type}. Choose 'query_head', 'block', or 'overall'.")
else:
raise ValueError(f"Invalid select_mode: {select_mode}. Choose 'topk' or 'threshold'.")
# eliminate self_chunk in MoBA branch
gate_mask = gate_mask & ~gate_self_chunk_mask
# if gate_mask is all false, perform flash_attn instead
if gate_mask.sum() == 0:
return flash_attn_varlen_func(
q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen, causal=False
)
# Determine which query positions are selected.
# nonzero_indices has shape [N, 3] where each row is [chunk_index, head_index, seq_index].
moba_q_indices = gate_mask.reshape(gate_mask.shape[0], -1).nonzero(as_tuple=True)[-1] # [(h s k)]
moba_q_sh_indices = (moba_q_indices % seqlen) * num_head + (moba_q_indices // seqlen)
moba_q = rearrange(q, "s h d -> (h s) d").index_select(0, moba_q_indices).unsqueeze(1)
# Build cumulative sequence lengths for the selected queries.
moba_seqlen_q = gate_mask.sum(dim=-1).flatten()
q_zero_mask = moba_seqlen_q == 0
valid_expert_mask = ~q_zero_mask
if q_zero_mask.sum() > 0:
moba_seqlen_q = moba_seqlen_q[valid_expert_mask]
moba_cu_seqlen_q = torch.cat(
(
torch.tensor([0], device=q.device, dtype=moba_seqlen_q.dtype),
moba_seqlen_q.cumsum(dim=0),
),
dim=0,
).to(torch.int32)
# Rearrange gathered KV for the MOBA branch.
experts_tensor = rearrange(gathered, "nc cl two h d -> (nc h) cl two d")
valid_expert_lengths = chunk_lengths.unsqueeze(1).expand(num_filtered_chunk, num_head).reshape(-1).to(torch.int32)
if q_zero_mask.sum() > 0:
experts_tensor = experts_tensor[valid_expert_mask]
valid_expert_lengths = valid_expert_lengths[valid_expert_mask]
seq_range = torch.arange(experts_tensor.shape[1], device=experts_tensor.device).unsqueeze(0)
mask = seq_range < valid_expert_lengths.unsqueeze(1)
moba_kv = experts_tensor[mask] # Shape: ((nc h cl_valid) two d)
moba_kv = moba_kv.unsqueeze(2) # Shape: ((nc h cl_valid) two 1 d)
moba_cu_seqlen_kv = torch.cat(
[torch.zeros(1, device=experts_tensor.device, dtype=torch.int32),
valid_expert_lengths.cumsum(dim=0)],
dim=0,
).to(torch.int32)
assert (
moba_cu_seqlen_kv.shape == moba_cu_seqlen_q.shape
), f"Mismatch between moba_cu_seqlen_kv.shape and moba_cu_seqlen_q.shape: {moba_cu_seqlen_kv.shape} vs {moba_cu_seqlen_q.shape}"
return MixedAttention.apply(
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
max_seqlen,
moba_chunk_size,
moba_q_sh_indices,
)
def process_moba_input(
x,
patch_resolution,
chunk_size,
):
"""
Process inputs for the attention function.
Args:
x (torch.Tensor): Input tensor with shape [batch_size, num_patches, num_heads, head_dim].
patch_resolution (tuple): Tuple containing the patch resolution (t, h, w).
chunk_size (int): Size of the chunk. (maybe tuple or int, according to chunk type)
Returns:
torch.Tensor: Processed input tensor.
"""
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
moba_chunk_size = int(chunk_size * patch_resolution[1] * patch_resolution[2])
else:
assert isinstance(chunk_size, (Tuple, list)), f"chunk_size should be a tuple, list, or int, now it is: {type(chunk_size)}"
if len(chunk_size) == 2:
assert patch_resolution[1] % chunk_size[0] == 0 and patch_resolution[2] % chunk_size[1] == 0, f"spatial patch_resolution {patch_resolution[1:]} should be divisible by 2d chunk_size {chunk_size}"
nch, ncw = patch_resolution[1] // chunk_size[0], patch_resolution[2] // chunk_size[1]
x = rearrange(x, "b (t nch ch ncw cw) n d -> b (nch ncw t ch cw) n d", t=patch_resolution[0], nch=nch, ncw=ncw, ch=chunk_size[0], cw=chunk_size[1])
moba_chunk_size = patch_resolution[0] * chunk_size[0] * chunk_size[1]
elif len(chunk_size) == 3:
assert patch_resolution[0] % chunk_size[0] == 0 and patch_resolution[1] % chunk_size[1] == 0 and patch_resolution[2] % chunk_size[2] == 0, f"patch_resolution {patch_resolution} should be divisible by 3d chunk_size {chunk_size}"
nct, nch, ncw = patch_resolution[0] // chunk_size[0], patch_resolution[1] // chunk_size[1], patch_resolution[2] // chunk_size[2]
x = rearrange(x, "b (nct ct nch ch ncw cw) n d -> b (nct nch ncw ct ch cw) n d", nct=nct, nch=nch, ncw=ncw, ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
moba_chunk_size = chunk_size[0] * chunk_size[1] * chunk_size[2]
else:
raise ValueError(f"chunk_size should be a int, or a tuple of length 2 or 3, now it is: {len(chunk_size)}")
return x, moba_chunk_size
def process_moba_output(
x,
patch_resolution,
chunk_size,
):
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
pass
elif len(chunk_size) == 2:
x = rearrange(x, "b (nch ncw t ch cw) n d -> b (t nch ch ncw cw) n d", nch=patch_resolution[1] // chunk_size[0], ncw=patch_resolution[2] // chunk_size[1], t=patch_resolution[0], ch=chunk_size[0], cw=chunk_size[1])
elif len(chunk_size) == 3:
x = rearrange(x, "b (nct nch ncw ct ch cw) n d -> b (nct ct nch ch ncw cw) n d", nct=patch_resolution[0] // chunk_size[0], nch=patch_resolution[1] // chunk_size[1], ncw=patch_resolution[2] // chunk_size[2], ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
return x
# TEST
def generate_data(batch_size, seqlen, num_head, head_dim, dtype):
random.seed(0)
torch.manual_seed(0)
torch.cuda.manual_seed(0)
device = torch.cuda.current_device()
q = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
k = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
v = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
print(f"q.shape: {q.shape}, k.shape: {k.shape}, v.shape: {v.shape}")
cu_seqlens = torch.arange(0, q.shape[0] * q.shape[1] + 1, q.shape[1], dtype=torch.int32, device='cuda')
max_seqlen = q.shape[1]
q = rearrange(q, "b s ... -> (b s) ...")
k = rearrange(k, "b s ... -> (b s) ...")
v = rearrange(v, "b s ... -> (b s) ...")
return q, k, v, cu_seqlens, max_seqlen
def test_attn_varlen_moba_speed(batch, head, seqlen, head_dim, moba_chunk_size, moba_topk, dtype=torch.bfloat16, select_mode='threshold', simsum_threshold=0.25, threshold_type='query_head'):
"""Speed test comparing flash_attn vs moba_attention"""
# Get data
q, k, v, cu_seqlen, max_seqlen = generate_data(batch, seqlen, head, head_dim, dtype)
print(f"batch:{batch} head:{head} seqlen:{seqlen} chunk:{moba_chunk_size} topk:{moba_topk} select_mode: {select_mode} simsum_threshold:{simsum_threshold}")
vo_grad = torch.randn_like(q)
# Warmup
warmup_iters = 3
perf_test_iters = 10
# Warmup
for _ in range(warmup_iters):
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
torch.autograd.backward(o, vo_grad)
torch.cuda.synchronize()
start_flash = time.perf_counter()
for _ in range(perf_test_iters):
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
torch.autograd.backward(o, vo_grad)
torch.cuda.synchronize()
time_flash = (time.perf_counter() - start_flash) / perf_test_iters * 1000
# Warmup
for _ in range(warmup_iters):
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
torch.autograd.backward(om, vo_grad)
torch.cuda.synchronize()
start_moba = time.perf_counter()
for _ in range(perf_test_iters):
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
torch.autograd.backward(om, vo_grad)
torch.cuda.synchronize()
time_moba = (time.perf_counter() - start_moba) / perf_test_iters * 1000
print(f"Flash: {time_flash:.2f}ms, MoBA: {time_moba:.2f}ms")
print(f"Speedup: {time_flash / time_moba:.2f}x")
if __name__ == "__main__":
"""
CUDA_VISIBLE_DEVICES=1 \
python -u csrc/attn/vmoba_attn/vmoba/vmoba.py
"""
test_attn_varlen_moba_speed(batch=1, head=12, seqlen=32760, head_dim=128, moba_chunk_size=32760 // 3 // 6 // 4, moba_topk=3, select_mode='threshold', simsum_threshold=0.3, threshold_type='query_head')
-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
+45 -21
View File
@@ -1,7 +1,9 @@
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
FROM nvidia/cuda:12.8.0-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,36 @@ 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 flash-attn==2.8.3 --no-build-isolation
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 STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn/sliding_tile_attn && \
git submodule update --init --recursive && \
python setup.py install
# 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
# Install VSA
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn/video_sparse_attn && \
git submodule update --init --recursive && \
python setup.py install
EXPOSE 22
EXPOSE 22
+45 -21
View File
@@ -1,7 +1,9 @@
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
FROM nvidia/cuda:12.8.0-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,36 @@ 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 flash-attn==2.8.3 --no-build-isolation
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 STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn/sliding_tile_attn && \
git submodule update --init --recursive && \
python setup.py install
# 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
# Install VSA
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn/video_sparse_attn && \
git submodule update --init --recursive && \
python setup.py install
EXPOSE 22
EXPOSE 22
+6 -6
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 flash-attn==2.8.3 --no-build-isolation
COPY . .
@@ -58,15 +58,15 @@ RUN source $HOME/.local/bin/env && \
# Install STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn && \
cd csrc/attn/sliding_tile_attn && \
git submodule update --init --recursive && \
python setup_sta.py install
python setup.py install
# Install VSA
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn && \
cd csrc/attn/video_sparse_attn && \
git submodule update --init --recursive && \
python setup_vsa.py install
python setup.py install
EXPOSE 22
EXPOSE 22
+72
View File
@@ -0,0 +1,72 @@
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 STA (Sliding Tile Attention)
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn/sliding_tile_attn && \
git submodule update --init --recursive && \
python setup.py install
# Install VSA
RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
cd csrc/attn/video_sparse_attn && \
git submodule update --init --recursive && \
python setup.py install
EXPOSE 22
+1
View File
@@ -23,3 +23,4 @@ clean:
@$(SPHINXBUILD) -M clean "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
rm -rf "$(SOURCEDIR)/getting_started/examples"
rm -rf "$(SOURCEDIR)/inference/examples"
rm -rf "$(SOURCEDIR)/training/examples"
+1 -1
View File
@@ -11,5 +11,5 @@ commonmark # Required by sphinx-argparse when using :markdownhelp:
# packages to install to build the documentation
cachetools
-f https://download.pytorch.org/whl/cpu
# -f https://download.pytorch.org/whl/cpu
torch
Binary file not shown.

After

Width:  |  Height:  |  Size: 194 KiB

+2 -2
View File
@@ -9,11 +9,11 @@
## Initialization Configuration
```{autodoc2-summary}
fastvideo.v1.configs.pipelines.PipelineConfig
fastvideo.configs.pipelines.PipelineConfig
```
## Sampling Configuration
```{autodoc2-summary}
fastvideo.v1.configs.sample.SamplingParam
fastvideo.configs.sample.SamplingParam
```
+2 -5
View File
@@ -18,7 +18,6 @@ import os
import re
import sys
from pathlib import Path
from typing import Optional
import requests
@@ -97,8 +96,7 @@ copybutton_prompt_is_regexp = True
#
html_title = project
html_theme = 'sphinx_book_theme'
html_logo = '../../assets/logo.jpg'
#html_favicon = 'assets/logos/vllm-logo-only-light.ico'
html_logo = '../../assets/logos/icon_simple.svg'
html_theme_options = {
'path_to_docs': 'docs/source',
'repository_url': 'https://github.com/hao-ai-lab/FastVideo/',
@@ -168,8 +166,7 @@ _cached_base: str = ""
_cached_branch: str = ""
def get_repo_base_and_branch(
pr_number: str) -> tuple[Optional[str], Optional[str]]:
def get_repo_base_and_branch(pr_number: str) -> tuple[str | None, str | None]:
global _cached_base, _cached_branch
if _cached_base and _cached_branch:
return _cached_base, _cached_branch
@@ -3,12 +3,12 @@
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:
@@ -6,14 +6,16 @@ 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)
When creating your pod template, use this image:
```
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest
```
Paste Container Start Command to support SSH ([RunPod Docs](https://docs.runpod.io/pods/configuration/use-ssh)):
+24 -3
View File
@@ -22,10 +22,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 +46,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 +60,14 @@ 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
```
+30 -30
View File
@@ -1,24 +1,24 @@
# 🔍 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
- [`fastvideo/pipelines/`](#design-pipeline-system) - Core diffusion pipeline components
- [`fastvideo/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/attention/`](#design-optimized-attention) - Optimized attention implementations
- [`fastvideo/distributed/`](#design-distributed-processing) - Distributed computing utilities
- [`fastvideo/layers/`](#design-tensor-parallelism) - Custom neural network layers
- [`fastvideo/platforms/`](#design-platforms) - Hardware platform abstractions
- [`fastvideo/worker/`](#design-executor-and-worker-abstractions) - Multi-GPU process management
- [`fastvideo/fastvideo_args.py`](#design-fastvideo-args) - Argument handling
- [`fastvideo/forward_context.py`](#design-forwardcontext) - Forward pass context management
- `fastvideo/utils.py` - Utility functions
- [`fastvideo/logger.py`](#design-logger) - Logging infrastructure
## Core Architecture
@@ -32,7 +32,7 @@ FastVideo separates model components from execution logic with these principles:
(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
@@ -111,7 +111,7 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward
(design-forwardbatch)=
### 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
@@ -123,14 +123,14 @@ 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`
@@ -157,7 +157,7 @@ def forward(
VAEs handle conversion between pixel space and latent space:
- **Location**: `fastvideo/v1/models/vaes/`
- **Location**: `fastvideo/models/vaes/`
- **Examples**:
- `AutoencoderKLWan`
- `AutoencoderKLHunyuanVideo`
@@ -175,7 +175,7 @@ FastVideo's VAE implementations include:
Encoders process conditioning inputs into embeddings:
- **Location**: `fastvideo/v1/models/encoders/`
- **Location**: `fastvideo/models/encoders/`
- **Text Encoders**:
- `CLIPTextModel`
- `LlamaModel`
@@ -193,7 +193,7 @@ FastVideo implements optimizations such as:
Schedulers manage the diffusion sampling process:
- **Location**: `fastvideo/v1/models/schedulers/`
- **Location**: `fastvideo/models/schedulers/`
- **Examples**:
- `UniPCMultistepScheduler`
- `FlowMatchEulerDiscreteScheduler`
@@ -219,7 +219,7 @@ def step(
(design-optimized-attention)=
## 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:
@@ -248,7 +248,7 @@ Supports various patterns with memory optimization techniques:
(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
@@ -260,7 +260,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 +288,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,
@@ -312,7 +312,7 @@ Efficient communication primitives minimize distributed overhead:
### 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
@@ -333,7 +333,7 @@ with set_forward_context(current_timestep, attn_metadata, fastvideo_args):
(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
@@ -360,7 +360,7 @@ This design allows FastVideo to efficiently utilize multiple GPUs while providin
(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
@@ -377,7 +377,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):
@@ -398,9 +398,9 @@ 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:
@@ -0,0 +1,43 @@
(v0-data-preprocess)=
# 🧱 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 computations, 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](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation.html). 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
+240 -41
View File
@@ -5,7 +5,6 @@ import itertools
import re
from dataclasses import dataclass, field
from pathlib import Path
from typing import Optional
ROOT_DIR = Path(__file__).parent.parent.parent.resolve()
ROOT_DIR_RELATIVE = '../../../..'
@@ -28,6 +27,15 @@ def fix_case(text: str) -> str:
"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
}
@@ -89,7 +97,7 @@ class Example:
generate() -> str: Generates the documentation content.
""" # noqa: E501
path: Path
category: Optional[str] = None
category: str | None = None
main_file: Path = field(init=False)
other_files: list[Path] = field(init=False)
title: str = field(init=False)
@@ -162,31 +170,35 @@ class Example:
return content
def generate_examples(generate_main_index=False):
"""
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 empty indices with dynamic paths
@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/source/examples"
if not main_index_dir.exists():
main_index_dir.mkdir(parents=True)
# Create the main examples index only if requested
examples_index = None
if generate_main_index:
examples_index = Index(
path=main_index_dir / "examples_index.md",
title="💡 Examples",
description=
"A collection of examples demonstrating usage of FastVideo.\nAll documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.", # noqa: E501
caption="Examples",
maxdepth=2)
# Category indices with dynamic paths based on category names
category_indices = {
"inference":
Index(
@@ -194,28 +206,54 @@ def generate_examples(generate_main_index=False):
"docs/source/inference/examples/examples_inference_index.md",
title="🚀 Examples",
description=
"Inference examples demonstrate how to use FastVideo in an offline setting, where the model is queried for predictions in batches. We recommend starting with <project:basic.md>.", # noqa: E501
"Inference examples demonstrate how to use FastVideo inference. We recommend starting with <project:basic.md>.",
caption="Examples",
maxdepth=1,
),
"training":
Index(
path=ROOT_DIR /
"docs/source/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/source/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 category, index in category_indices.items():
category_dir = index.path.parent
if not category_dir.exists():
category_dir.mkdir(parents=True)
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"]
# Find categorised examples
for category in category_indices:
print(category)
category_dir = EXAMPLE_DIR / category
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
for path in category_dir.glob("*/*.md"):
# 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
@@ -230,36 +268,197 @@ def generate_examples(generate_main_index=False):
continue
examples.append(Example(path.parent))
# Create document directories for each category based on category name and generate files
for example in sorted(examples, key=lambda e: e.path.stem):
print(example)
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]]]] = {}
for example in examples:
if example.category not in ["training", "distillation"]:
continue
category_dir = EXAMPLE_DIR / example.category
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 # Default to main index if available
index = examples_index
else:
# Skip examples without a category if no main index
print(f"Skipping {example.path} (no category and no main index)")
continue
# Place generated example markdown in the same directory as its index
# 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())
# Add the example to the index
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=False):
"""
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/source/examples"
examples_index = Index(
path=main_index_dir / "examples_index.md",
title="💡 Examples",
description=
"A collection of examples demonstrating usage of FastVideo.\nAll documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.",
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:
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)
assert examples_index is not None
examples_index.documents.insert(
0,
str(rel_path).replace(".md", ""))
+11 -113
View File
@@ -1,120 +1,18 @@
(fastvideo-installation)=
(installation-index)=
# 🔧 Installation
FastVideo currently only supports Linux and NVIDIA CUDA GPUs.
FastVideo supports the following hardware platforms:
## Requirements
:::{toctree}
:maxdepth: 1
:hidden:
- **OS: Linux**
- **Python: 3.10-3.12**
- **CUDA 12.4**
- **At least 1 NVIDIA GPU**
## 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-Linux-x86_64.sh
bash Miniconda3-latest-Linux-x86_64.sh
source ~/.bashrc
```
##### 2. Create and activate a Conda environment for FastVideo
```bash
# (Recommended) Create a new conda environment.
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.
installation/gpu
installation/mps
:::
#### 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
# (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
```
### Installation
```bash
pip install fastvideo
# or if you are using uv
uv pip install fastvideo
```
Also optionally install flash-attn:
```bash
pip install flash-attn==2.7.4.post1 --no-build-isolation
```
### 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 .
```
### Optional Dependencies
#### Flash Attention
```bash
pip install flash-attn==2.7.4.post1 --no-build-isolation
```
## Set up using Docker
We also have prebuilt docker images with FastVideo dependencies pre-installed:
[Docker Images](#docker)
## Development Environment Setup
If you're planning to contribute to FastVideo please see the following page:
[Contributor Guide](#developer-overview)
## Hardware Requirements
### For Basic Inference
- NVIDIA GPU with CUDA 12.4 support
### For Lora Finetuning
- 40GB GPU memory each for 2 GPUs with lora
- 30GB GPU memory each for 2 GPUs with CPU offload and lora
### For Full Finetuning/Distillation
- Multiple high-memory GPUs recommended (e.g., H100)
## 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-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg) for additional support.
- <project:installation/gpu.md>
- NVIDIA CUDA
- <project:installation/mps.md>
- Apple silicon
@@ -0,0 +1,119 @@
# NVIDIA GPU
Instructions to install FastVideo for NVIDIA CUDA GPUs.
## Requirements
- **OS: Linux or Windows WSL**
- **Python: 3.10-3.12**
- **CUDA 12.8**
- **At least 1 NVIDIA GPU**
## 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-Linux-x86_64.sh
bash Miniconda3-latest-Linux-x86_64.sh
source ~/.bashrc
```
##### 2. Create and activate a Conda environment for FastVideo
```bash
# (Recommended) Create a new conda environment.
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.
Note that you can also use `uv` to install FastVideo in a Conda environment.
:::
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
```
### Installation
```bash
pip install fastvideo
# or if you are using uv
uv pip install fastvideo
```
Also optionally install flash-attn:
```bash
pip install flash-attn --no-build-isolation
```
### 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 .
```
### Optional Dependencies
#### Flash Attention
```bash
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)
## Development Environment Setup
If you're planning to contribute to FastVideo please see the following page:
[Contributor Guide](#developer-overview)
## Hardware Requirements
### For Basic Inference
- NVIDIA GPU with CUDA 12.8 support
### For Lora Finetuning
- 40GB GPU memory each for 2 GPUs with lora
- 30GB GPU memory each for 2 GPUs with CPU offload and lora
### For Full Finetuning/Distillation
- Multiple high-memory GPUs recommended (e.g., H100)
## 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.
@@ -0,0 +1,102 @@
# 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
```
:::{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.
Note that you can also use `uv` to install FastVideo in a Conda environment.
:::
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](#developer-overview)
## 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.
+18 -18
View File
@@ -10,20 +10,20 @@ This class will be the primary Python API for generating videos and images.
fastvideo.VideoGenerator
```
`````{py:class} VideoGenerator(fastvideo_args: fastvideo.v1.fastvideo_args.FastVideoArgs, executor_class: type[fastvideo.v1.worker.executor.Executor], log_stats: bool)
:canonical: fastvideo.v1.entrypoints.video_generator.VideoGenerator
`````{py:class} VideoGenerator(fastvideo_args: fastvideo.fastvideo_args.FastVideoArgs, executor_class: type[fastvideo.worker.executor.Executor], log_stats: bool)
:canonical: fastvideo.entrypoints.video_generator.VideoGenerator
```{autodoc2-docstring} fastvideo.v1.entrypoints.video_generator.VideoGenerator
```{autodoc2-docstring} fastvideo.entrypoints.video_generator.VideoGenerator
:parser: docs.source.autodoc2_docstring_parser
```
`VideoGenerator.from_pretrained()` should be the primary way of creating a new video generator.
````{py:method} from_pretrained(model_path: str, device: typing.Optional[str] = None, torch_dtype: typing.Optional[torch.dtype] = None, pipeline_config: typing.Optional[typing.Union[str | fastvideo.v1.configs.pipelines.PipelineConfig]] = None, **kwargs) -> fastvideo.v1.entrypoints.video_generator.VideoGenerator
:canonical: fastvideo.v1.entrypoints.video_generator.VideoGenerator.from_pretrained
````{py:method} from_pretrained(model_path: str, device: typing.Optional[str] = None, torch_dtype: typing.Optional[torch.dtype] = None, pipeline_config: typing.Optional[typing.Union[str | fastvideo.configs.pipelines.PipelineConfig]] = None, **kwargs) -> fastvideo.entrypoints.video_generator.VideoGenerator
:canonical: fastvideo.entrypoints.video_generator.VideoGenerator.from_pretrained
:classmethod:
```{autodoc2-docstring} fastvideo.v1.entrypoints.video_generator.VideoGenerator.from_pretrained
```{autodoc2-docstring} fastvideo.entrypoints.video_generator.VideoGenerator.from_pretrained
:parser: docs.source.autodoc2_docstring_parser
```
@@ -38,25 +38,25 @@ The follow two classes `PipelineConfig` and `SamplingParam` are used to configur
```
`````{py:class} PipelineConfig
:canonical: fastvideo.v1.configs.pipelines.base.PipelineConfig
:canonical: fastvideo.configs.pipelines.base.PipelineConfig
```{autodoc2-docstring} fastvideo.v1.configs.pipelines.base.PipelineConfig
```{autodoc2-docstring} fastvideo.configs.pipelines.base.PipelineConfig
:parser: docs.source.autodoc2_docstring_parser
```
````{py:method} from_pretrained(model_path: str) -> fastvideo.v1.configs.pipelines.base.PipelineConfig
:canonical: fastvideo.v1.configs.pipelines.base.PipelineConfig.from_pretrained
````{py:method} from_pretrained(model_path: str) -> fastvideo.configs.pipelines.base.PipelineConfig
:canonical: fastvideo.configs.pipelines.base.PipelineConfig.from_pretrained
:classmethod:
```{autodoc2-docstring} fastvideo.v1.configs.pipelines.base.PipelineConfig.from_pretrained
```{autodoc2-docstring} fastvideo.configs.pipelines.base.PipelineConfig.from_pretrained
:parser: docs.source.autodoc2_docstring_parser
```
````{py:method} dump_to_json(file_path: str)
:canonical: fastvideo.v1.configs.pipelines.base.PipelineConfig.dump_to_json
:canonical: fastvideo.configs.pipelines.base.PipelineConfig.dump_to_json
```{autodoc2-docstring} fastvideo.v1.configs.pipelines.base.PipelineConfig.dump_to_json
```{autodoc2-docstring} fastvideo.configs.pipelines.base.PipelineConfig.dump_to_json
:parser: docs.source.autodoc2_docstring_parser
```
@@ -68,16 +68,16 @@ The follow two classes `PipelineConfig` and `SamplingParam` are used to configur
```
`````{py:class} SamplingParam
:canonical: fastvideo.v1.configs.sample.base.SamplingParam
:canonical: fastvideo.configs.sample.base.SamplingParam
```{autodoc2-docstring} fastvideo.v1.configs.sample.base.SamplingParam
```{autodoc2-docstring} fastvideo.configs.sample.base.SamplingParam
:parser: docs.source.autodoc2_docstring_parser
```
````{py:method} from_pretrained(model_path: str) -> fastvideo.v1.configs.sample.base.SamplingParam
:canonical: fastvideo.v1.configs.sample.base.SamplingParam.from_pretrained
````{py:method} from_pretrained(model_path: str) -> fastvideo.configs.sample.base.SamplingParam
:canonical: fastvideo.configs.sample.base.SamplingParam.from_pretrained
:classmethod:
```{autodoc2-docstring} fastvideo.v1.configs.sample.base.SamplingParam.from_pretrained
```{autodoc2-docstring} fastvideo.configs.sample.base.SamplingParam.from_pretrained
:parser: docs.source.autodoc2_docstring_parser
```
+31 -19
View File
@@ -1,6 +1,6 @@
# Welcome to FastVideo
:::{figure} ../../assets/logo.jpg
:::{figure} ../../assets/logos/logo.svg
:align: center
:alt: FastVideo
:class: no-scaled-link
@@ -9,7 +9,7 @@
:::{raw} html
<p style="text-align:center">
<strong>FastVideo is a unified framework for accelerated video generation.
<strong>FastVideo is a unified inference and post-training framework for accelerated video generation.
</strong>
</p>
@@ -21,11 +21,10 @@
</p>
:::
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.
FastVideo is an inference and post-training framework for diffusion models. It features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
<div style="text-align: center;">
<img src=_static/images/perf.png width="100%"/>
<img src=_static/images/fastwan.png width="100%"/>
</div>
## Key Features
@@ -35,16 +34,11 @@ FastVideo has the following features:
- [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.
- E2E post-training support
- Data preprocessing pipeline for video data.
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 using [Video Sparse Attention](https://arxiv.org/pdf/2505.13389) and [Distribution Matching Distillation](https://tianweiy.github.io/dmd2/)
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs.
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
## Documentation
@@ -63,22 +57,31 @@ getting_started/installation
:maxdepth: 1
inference/inference_quick_start
inference/examples/examples_inference_index
inference/configuration
inference/optimizations
inference/comfyui
inference/support_matrix
inference/examples/examples_inference_index
inference/cli
inference/add_pipeline
inference/v0_inference
:::
:::{toctree}
:caption: Training
:maxdepth: 1
training/examples/examples_training_index
training/data_preprocess
training/distillation
training/finetune
<!-- training/finetune -->
:::
:::{toctree}
:caption: Distillation
:maxdepth: 1
distillation/examples/examples_distillation_index
distillation/data_preprocess
distillation/dmd
:::
% What is STA Kernel?
@@ -91,6 +94,15 @@ sliding_tile_attention/installation
sliding_tile_attention/demo
:::
% What is VSA Kernel?
:::{toctree}
:caption: Video Sparse Attention
:maxdepth: 1
video_sparse_attention/installation
:::
:::{toctree}
:caption: Design
:maxdepth: 1
+22 -22
View File
@@ -12,7 +12,7 @@ This guide explains how to implement a custom diffusion pipeline in FastVideo, l
4. **Register Your Pipeline** - Make it discoverable by the framework
5. **Configure Your Pipeline** - (Coming soon)
Need help? Join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg).
Need help? Join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ).
## Step 1: Pipeline Modules
@@ -46,25 +46,25 @@ FastVideo uses the Hugging Face Diffusers format for model organization:
### Implementing Modules
Place new modules in the appropriate directories:
- Encoders: `fastvideo/v1/models/encoders/`
- VAEs: `fastvideo/v1/models/vaes/`
- Transformer models: `fastvideo/v1/models/dits/`
- Schedulers: `fastvideo/v1/models/schedulers/`
- Encoders: `fastvideo/models/encoders/`
- VAEs: `fastvideo/models/vaes/`
- Transformer models: `fastvideo/models/dits/`
- Schedulers: `fastvideo/models/schedulers/`
### Adapting Model Layers
#### Layer Replacements
Replace standard PyTorch layers with FastVideo optimized versions:
- nn.LayerNorm → fastvideo.v1.layers.layernorm.RMSNorm
- Embedding layers → fastvideo.v1.layers.vocab_parallel_embedding modules
- Activation functions → versions from fastvideo.v1.layers.activation
- nn.LayerNorm → fastvideo.layers.layernorm.RMSNorm
- Embedding layers → fastvideo.layers.vocab_parallel_embedding modules
- Activation functions → versions from fastvideo.layers.activation
#### Distributed Linear Layers
Use appropriate parallel layers for distribution:
```python
# Output dimension parallelism
from fastvideo.v1.layers.linear import ColumnParallelLinear
from fastvideo.layers.linear import ColumnParallelLinear
self.q_proj = ColumnParallelLinear(
input_size=hidden_size,
output_size=head_size * num_heads,
@@ -73,7 +73,7 @@ self.q_proj = ColumnParallelLinear(
)
# Fused QKV projection
from fastvideo.v1.layers.linear import QKVParallelLinear
from fastvideo.layers.linear import QKVParallelLinear
self.qkv_proj = QKVParallelLinear(
hidden_size=hidden_size,
head_size=attention_head_dim,
@@ -82,7 +82,7 @@ self.qkv_proj = QKVParallelLinear(
)
# Input dimension parallelism
from fastvideo.v1.layers.linear import RowParallelLinear
from fastvideo.layers.linear import RowParallelLinear
self.out_proj = RowParallelLinear(
input_size=head_size * num_heads,
output_size=hidden_size,
@@ -96,8 +96,8 @@ Replace standard attention with FastVideo's optimized attention:
```python
# Local attention patterns
from fastvideo.v1.attention import LocalAttention
from fastvideo.v1.attention.backends.abstract import _Backend
from fastvideo.attention import LocalAttention
from fastvideo.attention.backends.abstract import _Backend
self.attn = LocalAttention(
num_heads=num_heads,
head_size=head_dim,
@@ -108,7 +108,7 @@ self.attn = LocalAttention(
)
# Distributed attention for long sequences
from fastvideo.v1.attention import DistributedAttention
from fastvideo.attention import DistributedAttention
self.attn = DistributedAttention(
num_heads=num_heads,
head_size=head_dim,
@@ -130,7 +130,7 @@ self.attn = DistributedAttention(
Register implemented modules in the model registry:
```python
# In fastvideo/v1/models/registry.py
# In fastvideo/models/registry.py
_TEXT_TO_VIDEO_DIT_MODELS = {
"YourTransformerModel": ("dits", "yourmodule", "YourTransformerClass"),
}
@@ -145,7 +145,7 @@ _VAE_MODELS = {
Create a new directory for your pipeline:
```
fastvideo/v1/pipelines/
fastvideo/pipelines/
├── your_pipeline/
│ ├── __init__.py
│ └── your_pipeline.py
@@ -167,13 +167,13 @@ Pipelines are composed of stages, each handling a specific part of the diffusion
### Creating Your Pipeline
```python
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.v1.pipelines.stages import (
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (
InputValidationStage, CLIPTextEncodingStage, TimestepPreparationStage,
LatentPreparationStage, DenoisingStage, DecodingStage
)
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
import torch
class MyCustomPipeline(ComposedPipelineBase):
@@ -246,7 +246,7 @@ EntryClass = MyCustomPipeline
If existing stages don't meet your needs, create custom ones:
```python
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.base import PipelineStage
class MyCustomStage(PipelineStage):
"""Custom processing stage for the pipeline."""
@@ -305,7 +305,7 @@ EntryClass = [MyCustomPipeline, MyOtherPipeline]
```
The registry will automatically:
1. Scan all packages under `fastvideo/v1/pipelines/`
1. Scan all packages under `fastvideo/pipelines/`
2. Look for `EntryClass` variables
3. Register pipelines using their class names as identifiers
+3 -3
View File
@@ -27,7 +27,7 @@ fastvideo generate --help
### Hardware Configuration
- `--num-gpus {NUM_GPUS}`: Number of GPUs to use
- `--tp-size {TP_SIZE}`: Tensor parallelism size (Typically should match the number of GPUs)
- `--tp-size {TP_SIZE}`: Tensor parallelism size (only for the encoder, should not be larger than 1 if text encoder offload is enabled, as layerwise offload + prefetch is faster)
- `--sp-size {SP_SIZE}`: Sequence parallelism size (Typically should match the number of GPUs)
#### Video Configuration
@@ -68,7 +68,7 @@ Example configuration file (config.json):
"output_path": "outputs/",
"num_gpus": 2,
"sp_size": 2,
"tp_size": 2,
"tp_size": 1,
"num_frames": 45,
"height": 720,
"width": 1280,
@@ -102,7 +102,7 @@ prompt: "A beautiful woman in a red dress walking down a street"
output_path: "outputs/"
num_gpus: 2
sp_size: 2
tp_size: 2
tp_size: 1
num_frames: 45
height: 720
width: 1280
+5
View File
@@ -0,0 +1,5 @@
# FastVideo + ComfyUI
FastVideo provides a custom node suite for ComfyUI.
See this [README](https://github.com/hao-ai-lab/FastVideo/tree/main/comfyui) for instructions.
+1 -1
View File
@@ -27,7 +27,7 @@ def main():
model_name = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
config = PipelineConfig.from_pretrained(model_name)
config.vae_precision = "fp16"
config.use_cpu_offload = True
config.dit_cpu_offload = True
# Create the generator
generator = VideoGenerator.from_pretrained(
@@ -5,7 +5,7 @@ This page contains step-by-step instructions to get you quickly started with vid
## Requirements
- **OS**: Linux (Tested on Ubuntu 22.04+)
- **Python**: 3.10-3.12
- **CUDA**: 12.4
- **CUDA**: 12.8
- **GPU**: At least one NVIDIA GPU
## Installation
@@ -121,4 +121,4 @@ If the generated video doesn't match your prompt:
- Learn about using [Optimizations](#inference-optimizations)
- See [Examples](../examples/examples_inference_index.md) for more usage scenarios
- Join our [Community Discord](https://discord.gg/JA7cksDz86).
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg).
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ).

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