178 Commits
Author SHA1 Message Date
kijai 088128b224 Don't count .disabled as duplicate 2026-05-24 16:07:12 +03:00
Jukka Seppänen 126819826a Merge pull request #1948 from haosenwang1018/fix/bare-excepts
fix: replace 47 bare excepts with except Exception
2026-05-24 16:02:36 +03:00
kijai 8f1804bf72 Fix multitalk_audio_stride init 2026-05-24 15:59:40 +03:00
kijai 5437b016e3 Initial LongCatAvatar 1.5 support 2026-05-23 23:13:28 +03:00
kijai d18cdb1859 Fix offload on interrupt 2026-05-05 14:00:56 +03:00
haosenwang1018 0d78230336 fix: replace 47 bare excepts with except Exception
Bare except catches KeyboardInterrupt and SystemExit, masking real errors.
2026-02-25 06:04:52 +00:00
Jukka Seppänen df8f3e49da Delete .github/FUNDING.yml 2026-02-22 15:18:05 +02:00
Jukka Seppänen 86ad93d616 Merge pull request #1840 from little6neko/main
feat(wanvideo): add manual start reference support for WanAnimate loop
2026-02-16 15:44:15 +02:00
Jukka Seppänen 309491b269 Merge pull request #1908 from jjdejong/patch-1
Implement conditional offloading for diffusion model
2026-02-16 15:35:45 +02:00
kijai 06122f1e9d Rope offset should be disabled by default 2026-02-16 15:32:54 +02:00
kijai 5d36631795 I2V cross_attn fixes 2026-02-16 15:31:20 +02:00
kijai 3d7b49e2df Move string_to_seed import 2026-02-02 17:02:05 +02:00
kijai e091c4a774 version 1.4.7 2026-02-02 01:43:38 +02:00
kijai 60a387579e Update wanvideo_2_1_14B_I2V_SkyReelsV3_TalkingAvatar_example_01.json 2026-02-01 15:16:33 +02:00
kijai 9ae1a4f2c9 Update multitalk_loop.py 2026-02-01 15:15:37 +02:00
kijai f55b7b3d89 Update wanvideo_2_1_14B_I2V_SkyReelsV3_TalkingAvatar_example_01.json 2026-02-01 15:15:33 +02:00
kijai e4e7f413f7 Make compatible with latest ComfyUI version 2026-02-01 15:15:27 +02:00
kijai e21fe20a4d Support SkyReels TalkingAvatar (A2V) 2026-01-31 23:24:29 +02:00
kijai 2c5a04cc63 Init image_cond_mask 2026-01-30 18:59:41 +02:00
kijai d00abe52d7 version 1.4.6 2026-01-28 02:09:50 +02:00
kijai 2952f5d9dd Show T5 loading progress bar 2026-01-26 15:53:48 +02:00
kijai 339e0fec81 Make NAG inplace application optional 2026-01-23 18:35:30 +02:00
kijai 2c2a6e1889 RoPE frequency offset option for storymem
I'm not 100% sure on this, I initially tested this when I noticed the original code doesn't, but it's described in the paper... now I see the original code has added it too so it seems to be the intended way to use it.
2026-01-23 14:31:04 +02:00
kijai 8640bfad52 VRAM optimizations
Minor for almost everything, major for multitalk when using masks
2026-01-22 19:46:52 +02:00
kijai 3b4a711a40 Reduce NAG memory usage 2026-01-22 17:43:34 +02:00
Jean J. de Jong b99eac73da Implement conditional offloading for diffusion model
Add condition to skip offloading if load_device is main_device.
2026-01-21 17:33:58 +01:00
kijai 9f9a8e71c2 Revert " Fix: Fix torch.arange bounds error in context window processing"
This reverts commit 707bbcd72b.
2026-01-17 20:13:43 +02:00
小六妞儿 0fbcbed06a Merge branch 'kijai:main' into main 2026-01-16 15:52:00 +08:00
Jukka Seppänen f2e2e2550a Merge pull request #1877 from YFJack/patch-1
Fix: Fix torch.arange bounds error in context window processing
2026-01-15 17:06:31 +02:00
小六妞儿 6f9832ed47 Merge branch 'kijai:main' into main 2026-01-15 14:05:52 +08:00
Yifu Wang 707bbcd72b Fix: Fix torch.arange bounds error in context window processing
Fixed "upper bound and lower bound inconsistent with step sign" error
       when using WanVideoContextOptions with pose latents. Changed end index
       from c[-1] to c[-1] + 1 to properly include the last frame in the range.
2026-01-09 11:36:53 +08:00
Jukka Seppänen 855d103ee6 Merge pull request #1868 from vantagewithai/main
longcat avatar GGUF support.
2026-01-08 13:36:43 +02:00
kijai 64191921d4 Squashed commit of the following:
commit fdb23dec7d
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jan 5 22:11:04 2026 +0200

    Update model.py

commit 07d7d8ca8e
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jan 5 22:10:02 2026 +0200

    remove prints

commit 01869d4bf5
Merge: 55c6720 bf1d77f
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jan 5 18:47:48 2026 +0200

    Merge branch 'main' into longvie2

commit 55c672028b
Merge: b551ec9 be41f67
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 29 15:39:43 2025 +0200

    Merge branch 'main' into longvie2

commit b551ec9e31
Merge: 9f019d7 19bcee6
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 29 15:03:53 2025 +0200

    Merge branch 'main' into longvie2

commit 9f019d7dfb
Merge: fc5322f c5d3fb4
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 23 23:40:25 2025 +0200

    Merge branch 'main' into longvie2

commit fc5322fae4
Merge: 222fc70 e75f814
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 23 22:04:15 2025 +0200

    Merge branch 'main' into longvie2

commit 222fc70eb7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 23 17:18:55 2025 +0200

    Update nodes.py

commit 8509236da1
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 23 14:20:18 2025 +0200

    init
2026-01-05 22:11:20 +02:00
Vantage with AI 576f065073 Update tensor name condition for multitalk audio proj, to support longcat video avatar GGUF 2026-01-04 21:13:54 +05:30
kijai bf1d77fe15 Update wan_video_vae.py 2025-12-31 15:56:10 +02:00
kijai cced7fefbe Make VAE memory reporting optional to reduce log spam, and other logging updates 2025-12-31 15:46:40 +02:00
kijai 36bb0c73ee Correct infinite talk loop log entry
Just visual bug, no effect on anything else
2025-12-31 14:50:27 +02:00
kijai b982b4ef0c Update nodes_model_loading.py 2025-12-30 02:55:51 +02:00
kijai 3a7100bc39 Add node for UltraVico params 2025-12-30 02:53:38 +02:00
kijai be41f67fae Fix res_multistep steps 2025-12-29 15:39:27 +02:00
kijai 351829a2b6 Fix er_sde steps 2025-12-29 15:36:43 +02:00
kijai 19bcee67ed remove print 2025-12-29 02:50:04 +02:00
kijai 2fe4834178 Adjust ultravico frame_tokens 2025-12-29 02:48:59 +02:00
kijai 486564060f Add node to set attention mode per step and/or blocks 2025-12-29 01:35:32 +02:00
kijai 3730ccf603 Fix s2v 2025-12-28 10:34:13 +02:00
小六妞儿 58c1bcb7ce Merge branch 'kijai:main' into main 2025-12-28 12:59:47 +08:00
kijai 1eab022bb0 Update nodes.py 2025-12-28 02:35:42 +02:00
kijai 6fd4c6640c Adjust scheduler graph drawing 2025-12-28 01:27:18 +02:00
kijai 7a6efc1456 Add node for SVI 2.0 Pro 2025-12-28 00:35:14 +02:00
小六妞儿 64cbd28e00 feat(wanvideo): add manual start reference support for WanAnimate loop 2025-12-27 16:04:00 +08:00
kijai e855726f10 Fix enhance-a-video 2025-12-26 23:12:18 +02:00
kijai a896101ec8 I don't know why this suddenly errors 2025-12-26 22:28:04 +02:00
kijai fd818faa08 Fix context window ref latent device 2025-12-26 16:09:01 +02:00
kijai b132a82f7a Fix LongCat-Avatar audio padding when not enough audio provided for given window 2025-12-26 15:57:30 +02:00
kijai 4b709a7a04 Cleanup example_workflows folder some 2025-12-26 15:44:08 +02:00
kijai 220aac2771 Update LongCatAvatar_audio_image_to_video_example_01.json 2025-12-26 13:52:27 +02:00
kijai 027bed8c3d Cleanup 2025-12-26 13:51:05 +02:00
kijai b2c520ca44 Fix s2v 2025-12-26 12:55:13 +02:00
kijai 20942b8fd9 Cleanup: Remove flowedit code, restructure some other code
Due to lack of use and code maintainability
2025-12-26 12:51:22 +02:00
kijai 74f337e06c Adjust StoryMem lora scaling
This was probably too high afterall
2025-12-26 01:41:32 +02:00
kijai 264212dddb Fix uncond variable name when using zero star or fresca 2025-12-26 01:27:48 +02:00
kijai f988d19fdb StoryMem latents are supposed to be encoded one by one 2025-12-25 20:44:37 +02:00
kijai 95255c7ffa Add node to add story memory latents 2025-12-24 18:44:05 +02:00
kijai c42bf94b07 Automatically adjust LoRA alpha for Peft rs_lora weights
At least the original StoryMem -LoRAs need this
2025-12-24 17:58:30 +02:00
kijai dcae850b96 Allow higher lora scale 2025-12-24 17:07:59 +02:00
kijai c5d3fb450c Allow loading StoryMem -LoRAs
https://huggingface.co/Kevin-thu/StoryMem
2025-12-23 22:26:37 +02:00
kijai f28e7da442 version 1.4.5 2025-12-23 22:16:03 +02:00
kijai ac7d8cab98 Update LongCatAvatar_audio_image_to_video_example_01.json 2025-12-23 22:05:02 +02:00
kijai e75f814312 LongCat-Avatar example 2025-12-23 20:43:10 +02:00
kijai 1c24ef50f8 Remove print 2025-12-23 12:25:55 +02:00
kijai 5360eeb345 Squashed commit of the following:
commit fd32b14fdc
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 23 02:02:31 2025 +0200

    Clean prints

commit 1776695e26
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 23 01:48:12 2025 +0200

    Update nodes_model_loading.py

commit ef36204fa8
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 23 01:35:08 2025 +0200

    Reduce peak VRAM use

commit c6f32c1424
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 22 23:53:41 2025 +0200

    Norm dtype

commit 6d4a0f6e53
Merge: e7e0006 3e45021
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 22 22:11:38 2025 +0200

    Merge branch 'main' into longcat_avatar

commit e7e00061e5
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 22 00:43:01 2025 +0200

    Update nodes_sampler.py

commit eb5ec262a0
Merge: 7c0ba84 fed3b22
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 22 00:42:53 2025 +0200

    Merge branch 'main' into longcat_avatar

commit 7c0ba84a26
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Dec 21 23:00:43 2025 +0200

    remove prints

commit 06a86923e7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Dec 21 22:53:25 2025 +0200

    Fix ref latent

    oops

commit dca3106f10
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 20 18:46:32 2025 +0200

    Expose more options, make vid2vid easier

commit 175418b8d2
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 20 03:15:24 2025 +0200

    Create LongCatAvatar_testing_wip.json

commit 4a6e2d3c6c
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 20 03:14:49 2025 +0200

    Init
2025-12-23 12:17:07 +02:00
kijai 3e45021422 Support mixed precision fp8 scaled models 2025-12-22 21:40:14 +02:00
kijai 41683a0423 Possibly fix uni3c + multitalk 2025-12-22 19:44:53 +02:00
kijai 0ac7401271 Stand-in RoPE adjustments
The offset was a bit wrong
2025-12-22 18:56:04 +02:00
kijai fed3b22bc9 Update nodes_sampler.py 2025-12-21 15:06:46 +02:00
kijai 392d0305fc Add portrait_cfg 2025-12-21 15:06:14 +02:00
kijai be411bdef6 Support loading FlashPortrait model
Can just re-use FantasyPortrait for basic usage
2025-12-20 19:56:26 +02:00
Jukka Seppänen 2ee9e2f356 version 1.4.4 2025-12-17 18:35:11 +02:00
kijai ef82826161 Add LoadNLFModel -node
For local loading
2025-12-17 18:10:18 +02:00
kijai 93f7af6dc8 Make uni3c offloading optional 2025-12-17 16:22:37 +02:00
kijai ae6fe0853e Update nodes.py 2025-12-16 00:06:21 +02:00
kijai c80fed01fe Don't error on empty 2025-12-15 20:27:06 +02:00
kijai 95097fefc2 NLF predit: Add batching option, output bboxes as well 2025-12-15 20:11:40 +02:00
kijai c49fe98e55 version 1.4.3 2025-12-15 18:04:01 +02:00
kijai ebceb165cc SCAIL example 2025-12-15 18:03:37 +02:00
kijai e6bd1b413a Don't trigger context windows if not enough frames 2025-12-15 18:02:41 +02:00
kijai 4dbba4d06d Update readme.md 2025-12-15 17:35:48 +02:00
kijai a9e21f164c Squashed commit of the following:
commit 916fc0b1bcfd37b6bd9ece0daeb5b3cbaa53d0a9
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 15 17:30:37 2025 +0200

    Update nodes.py

commit 63818324f5dbb0b300064bea0402c4cd1bd57b2b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 15 17:30:26 2025 +0200

    Refactor RoPE caching

commit bb0c55da4d8f8bca4968704e877fd057a90a1eeb
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 15 01:59:16 2025 +0200

    Update nodes_sampler.py

commit a0447d55534857051606ee4201bc7f4e25aa73ae
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 15 01:28:09 2025 +0200

    Fix non scale wfs

commit fa761cc2f2a426faa9c391aeede62cf6f0fd7266
Merge: ea1677b 3aae54f
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 15 01:26:23 2025 +0200

    Merge branch 'main' into SCAIL

commit ea1677bd4ad42f19e369551590a9d4f17a36fa29
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Dec 14 19:41:43 2025 +0200

    Handle torchscript issue better

    Some other custom nodes globally set torch._C._jit_set_profiling_executor(False) which breaks the NLF model

commit e3cfa64bd3712ac153ce84a75215842c884a8ba4
Merge: ad7a0b9 3611341
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Dec 14 16:49:04 2025 +0200

    Merge branch 'main' into SCAIL

commit ad7a0b925de61ff705b928cd802e752e46089b42
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Dec 14 16:10:34 2025 +0200

    Fix possible uni3c issue

commit 74d97fa4bb7c58a0edf8516cc9fad4468da5c57e
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Dec 14 15:58:42 2025 +0200

    Match Uni3C temporal dim

commit 056d8ad96ffa5a223a8cd88c900a573a8d450e22
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Dec 14 14:47:58 2025 +0200

    Add warning for potential other overrides on torch.jit.script

commit f6dff002ffdcd880451955db298872ea90a4e3f8
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Dec 14 14:19:33 2025 +0200

    Add option to warmup the NLF model on load and fix it's offloading

commit a19107501dff23804e7db984d7da304a9955adc9
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Dec 14 13:45:20 2025 +0200

    Add error to indicate ComfyUI-RMBG currently breaks the NLF model

commit e2cfa486e48ead50195884167d9794c7caf0a69f
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 13 23:29:49 2025 +0200

    Cleanup unnecessary code

commit 462b61855fb96b0cb18cbccd48593256992808d7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 13 18:05:10 2025 +0200

    context windows

commit e57d4baeebf12c43e851c6c2467d698d7dbb4d03
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 13 16:55:23 2025 +0200

    Start/end percentages and strength

commit 3e507ae32256ed3e41cea69d9e26c30b5272968e
Merge: 1e5c7cb 0fa5383
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 13 16:09:16 2025 +0200

    Merge branch 'main' into SCAIL

commit 1e5c7cb2113138bdeae562d266f911c1e3edee91
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 13 15:45:39 2025 +0200

    Update nodes.py

commit 98f8e56bcacfc07e12cbb4b26555b2b28d9db92f
Merge: 9652146 78e3e18
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 13 15:42:44 2025 +0200

    Merge branch 'main' into SCAIL

commit 9652146763fb27e916a6853a8125efd0a67cd601
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 13 02:41:06 2025 +0200

    Add imitation of SCAIL pose drawing to the existing NLF node

    This only draws the pose with same colors, it's not meant as final solution, just for testing.

commit 1f86cebdaa97570ed88da0c9986b85c6664d62dc
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 13 01:11:56 2025 +0200

    test pose inputs

commit b348b21dbef0dcb92c0961df85959648e78da6aa
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Dec 12 20:10:48 2025 +0200

    Init
2025-12-15 17:31:01 +02:00
kijai a19d6bf12d Update nodes_sampler.py 2025-12-15 01:59:39 +02:00
kijai 3aae54f220 Allow WanMove to work with context windows 2025-12-15 01:22:21 +02:00
kijai 3611341339 Fix uni3c reloading 2025-12-14 16:48:55 +02:00
kijai 0fa5383106 Fix default rope for v2 sampler 2025-12-13 16:09:08 +02:00
kijai 78e3e1857c Add v2 sampling nodes for cleaner workflows
No functional changes, just cleaner nodes
2025-12-13 15:42:35 +02:00
kijai 164a6bbebd Restore LoRA step scheduling functionality 2025-12-13 00:54:22 +02:00
kijai 4a3ab6958a Update nodes.py 2025-12-12 15:43:48 +02:00
kijai 91413b33f5 One-to-all pose cfg 2025-12-12 15:23:56 +02:00
kijai 60b3ec57dd Support 1.3B OneToAll 2025-12-12 12:32:04 +02:00
kijai f6e2d6dd47 version 1.4.2 2025-12-11 17:43:31 +02:00
kijai 68d65979d6 Fix crash on latest frontend version 2025-12-11 17:30:08 +02:00
kijai 66a21e544a WanMove: Make compatible with TRACK -input type and make strength param do something 2025-12-11 01:45:05 +02:00
kijai 7273470faa Optimize trajectory drawing 2025-12-10 20:08:54 +02:00
kijai c36e56875e Update nodes.py 2025-12-10 20:05:49 +02:00
kijai f9287e3ecd Fix cfg when using WanMove 2025-12-10 20:05:41 +02:00
kijai 2790c532cc Handle OneToAll extra layer devices 2025-12-10 19:46:02 +02:00
kijai 369043c5d8 Add comfy core attention as option
For edge cases that may benefit from the older attention types
2025-12-10 19:39:05 +02:00
Jukka Seppänen eb6dd96049 Merge pull request #1749 from wong00/fix-flashattn3-attention-compute
[bugfix] Fixed the logic of computing attention in flashAttention3.
2025-12-10 11:52:18 +02:00
wangxin68 38a48c670a [bugfix] Fixed the logic for processing return values after computing attention in flashAttention3 2025-12-10 17:27:37 +08:00
kijai dd58511d4e Expose some WanMove track preview drawing options
Don't need to be so huge
2025-12-10 02:26:51 +02:00
kijai 4c97d27583 Fix LongCat error 2025-12-09 21:09:44 +02:00
kijai bcdc0c0661 Add node to use WanMove with native workflows 2025-12-09 21:04:56 +02:00
kijai 83c25644ef Create wanvideo_WanMove_I2V_example_01.json 2025-12-09 20:28:45 +02:00
kijai 113b6df04d Fix multiple tracks 2025-12-09 20:11:45 +02:00
kijai d000fbc645 Support WanMove
https://github.com/ali-vilab/Wan-Move
2025-12-09 19:55:56 +02:00
kijai cc58964027 Refactor init.py to something more sensible 2025-12-09 19:53:26 +02:00
kijai fa57681424 Reduce recompiles with unmerged loras 2025-12-09 14:40:52 +02:00
kijai 0a65354247 Allow larger wananimate window size for disabling windowing 2025-12-09 14:15:00 +02:00
kijai a27a4892b6 version 1.4.1 2025-12-09 12:56:40 +02:00
kijai 8b037bce2e Squashed commit of the following:
commit c3eb0f49faf68ab953f1b08b7e00225e041e5d0b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 9 12:55:49 2025 +0200

    move workflow

commit e129e25c26f9b55b527dd3e9f15c6e3f215af11f
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 9 11:17:17 2025 +0200

    Fix padding

commit f252f34eff5cc15ec6fc475f929cafa3e5b7f46c
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 9 01:38:17 2025 +0200

    Add long video example

commit 09ceab808b67a3b2fb7d1ee5fc0a1ad667739e2a
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 9 01:31:48 2025 +0200

    Support extension

commit 7ca221874e8a2cabfc766c51bb63774fde3c851b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 12:28:29 2025 +0200

    Might as well not even do control pass on uncond...

commit b55caf299e4d89148f5885e8e56bf8e411472dc3
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 12:15:59 2025 +0200

    Cfg fixes

commit fd54ba23e6746acb33a8bf124e5bc7de9d947ff1
Merge: 2f97b1b e867e64
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 10:39:55 2025 +0200

    Merge branch 'main' into onetoall

commit 2f97b1bd887367962542b9a6058f9f6e3c4ad4d7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 09:32:09 2025 +0200

    Add ref_mask input

commit 74cad232fd35347c50f2ed7465ff13e179ef8402
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 03:44:42 2025 +0200

    Update nodes_model_loading.py

commit 01a038eb4a30f29d868fbaef190e6e90da1a058d
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 03:11:08 2025 +0200

    Fix indentation

commit a95f4d6eaa4468e818910fec7ba11e1f92423d9b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 02:54:47 2025 +0200

    Update model.py

commit ad006985a1bafdf5941c0fa85a47852eb20a818a
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 02:54:19 2025 +0200

    Fix token replace

commit b5f0f44f1720586950756ad142a538e04814270f
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 02:50:52 2025 +0200

    Don't use token replace by default

commit 874174ec2921c528a4373097fd0bebbbb5257606
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 02:24:47 2025 +0200

    Create WanToAllAnimation_test.json

commit 9e6175855618c94c1bcb89c4b89879219410ce53
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 02:23:15 2025 +0200

    Add token replacement

commit 41fd76dfcbf0e70a3a7308a6fa0652fb492ed1f6
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 00:45:33 2025 +0200

    Use correct norm for reference attn

commit 705f5dcc8b6cd5fa6fe453f9bd01ffdf43a23078
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 00:11:17 2025 +0200

    cleanup

commit 4f095d97f80da807417d49d9aa7e9ee47145c85f
Merge: 3e4e4db 2369cdb
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Dec 7 18:44:01 2025 +0200

    Merge branch 'main' into onetoall

commit 3e4e4db35d3e266c39d48cd683f60384a737eca5
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Dec 7 00:27:23 2025 +0200

    handle controlnet better

commit c5742552a9af4a3ae208f9c2ead6e1105cc2c348
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 6 17:24:45 2025 +0200

    cleanup

commit c06ff9c06651c32953236802bd7fb385b9cf93ab
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 6 03:41:02 2025 +0200

    3D rope for controlnet

commit 948ea6b783f54892515cbc9cfe66484913904ee7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 6 03:08:04 2025 +0200

    pose input scaling

commit 90c2eff3b2d30d3a92ff5c27e4327a0ac80b642c
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 6 02:37:48 2025 +0200

    Cleanup

commit 9f7683422c1aa8ebe4d3380a86be98d6c589b270
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Dec 5 23:29:05 2025 +0200

    pose control

commit 0f217be4d8742741b0f89db50138214302a58dc3
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Dec 5 20:55:10 2025 +0200

    Support reference input
2025-12-09 12:56:11 +02:00
kijai e867e642d4 Fix longcat scheduler 2025-12-08 10:26:38 +02:00
kijai 00971fc121 Don't remove lora in loop sampling offload 2025-12-08 09:43:36 +02:00
Jukka Seppänen 0e81adb843 Fix for HuMo with I2V 2025-12-08 03:17:29 +02:00
kijai 7f9560eb33 Fix mistake in sageattn_ultravico 2025-12-07 19:56:03 +02:00
kijai 2369cdbbe9 Fix prompt splitting 2025-12-07 18:34:21 +02:00
kijai 123c9ca312 Allow HuMo to work with start_image 2025-12-07 18:18:16 +02:00
kijai 2b62866945 Fix SteadyDancer GGUF dtypes 2025-12-06 17:25:57 +02:00
kijai b5415573c5 Update nodes_model_loading.py 2025-12-05 20:55:33 +02:00
kijai 9e005fab90 Fix scaling for light TAE 2025-12-05 14:03:33 +02:00
kijai f85ea5a48a add "empty_frame_pad_image" input to I2V encode node for easier SVI lora use
This pads the empty frames with the given image, as should be done with SVI-shot and SVI 2.0 LoRAs
2025-12-05 13:05:35 +02:00
kijai ee1bb5c5c7 Update nodes_utility.py 2025-12-05 12:07:48 +02:00
kijai 041c31fe2e cleanup 2025-12-05 12:01:34 +02:00
kijai 4e7b8dd92c Clearer error on attention mode imports 2025-12-05 02:00:31 +02:00
kijai e2333d0f04 Update nodes_sampler.py 2025-12-04 16:51:30 +02:00
kijai 2dd4c6bf15 Merge branch 'pr/1705' 2025-12-04 16:48:50 +02:00
kijai 9d0a21228b Update nodes_sampler.py 2025-12-04 16:48:23 +02:00
kijai 001bc0e24a Merge branch 'main' of https://github.com/kijai/ComfyUI-WanVideoWrapper 2025-12-04 16:44:29 +02:00
kijai 7189922f36 Update __init__.py 2025-12-04 16:44:18 +02:00
kijai faf9c5927d Merge branch 'pr/1510' 2025-12-04 16:44:02 +02:00
Jukka Seppänen 6c5dc0ba2a Merge pull request #1546 from wzxysf/main
Enable vae tiling with end frame
2025-12-04 16:40:29 +02:00
Jukka Seppänen 995804166d Merge pull request #1651 from jamesjjcondon/patch-1
Refactor data_mean and data_std initialization
2025-12-04 16:40:05 +02:00
mossmatrix 4e2bf0f1fa Fix device mismatch error with I2V + Lynx embeds
Fixed RuntimeError when using I2V embeds with Lynx embeds where tensors
were on different devices (cuda:0 and cpu).

The issue occurred because lynx_ref_text_embed["prompt_embeds"] were not
explicitly moved to the GPU device before being passed to the transformer
during Lynx reference buffer extraction.

Changes:
- Move lynx text embeddings to device in both conditional and unconditional
  buffer extraction calls (lines 1114 and 1128)
- Ensures all tensors are on the same device during cross-attention operations
2025-12-04 00:34:54 -05:00
kijai b06c7d2d6d Add UltraViCo -sage attention mode and refactor some attention code
https://github.com/thu-ml/DiT-Extrapolation/
2025-12-03 13:32:45 +02:00
kijai 014e711972 cleanup 2025-12-03 11:09:09 +02:00
kijai c47b1ded69 Apply conv3d workaround to UniAnimate 2025-12-02 15:00:06 +02:00
kijai 7bc45daaf2 Cleanup whitespaces 2025-12-02 14:50:26 +02:00
kijai aebeeb9160 Code cleanup 2025-12-02 14:49:42 +02:00
kijai e60eb995ee version 1.4.0 2025-12-02 00:46:01 +02:00
kijai b9f6c9aa50 Update __init__.py 2025-12-02 00:45:22 +02:00
kijai c1fbc93521 Add ViBTScheduler to use ViBT models
https://github.com/Yuanshi9815/ViBT/tree/main
2025-12-02 00:26:50 +02:00
kijai c4ca252fea Fix compiled LoRA application
Still needs to be unfused
2025-12-01 20:38:01 +02:00
kijai c4db00609a Use torch.chunk for chunking 2025-12-01 17:39:50 +02:00
kijai 5a52d6b92f Add VAE feat_cache offloading, VAE tqdm progress bar and memory usage report 2025-12-01 15:05:20 +02:00
kijai a652c55bf7 Remove LoRAs when fully offloading as well 2025-12-01 12:50:22 +02:00
kijai 196d39695f Fix sageattn_varlen 2025-12-01 01:11:06 +02:00
kijai e5be3e5263 Use torch custom_ops to avoid graph breaks with torch.compile
Hopefully finally fixes the torch.compile VRAM issues...
2025-12-01 00:29:27 +02:00
kijai a6071c7be5 Merge branch 'main' into steadydancer 2025-11-30 17:56:50 +02:00
kijai 0cba1edd4e Better just not compile this as it's causing issues 2025-11-30 17:53:28 +02:00
kijai a9cd073f29 Remove unnecessary recompile when using cfg 2025-11-30 17:52:56 +02:00
kijai 1e9e2be622 Avoid recompile here 2025-11-30 17:32:24 +02:00
kijai 99c3978da4 Reduce peak VRAM usage when not using torch.compile (and some even with it)
Found some intermediates that weren't freed which should reduce VRAM usage overall, and modified RoPE application outside torch compile for similar gains than when using torch.compile.
2025-11-30 17:14:53 +02:00
kijai 30bd7d46eb example update 2025-11-29 02:38:02 +02:00
kijai 0c9d5b8dcc context windows 2025-11-29 02:11:57 +02:00
kijai 66d44ec8db This doesn't really do anything useful 2025-11-28 21:15:25 +02:00
kijai 394c7c13d2 Add strength controls 2025-11-28 21:11:20 +02:00
kijai c9931364b3 Create wanvideo_I2V_steadydancer_testing.json 2025-11-28 20:34:55 +02:00
kijai e54fa5d059 Init 2025-11-28 20:32:16 +02:00
kijai 772642b4f1 Update nodes.py 2025-11-27 20:19:06 +02:00
kijai 472ed70757 Add node to preview image embeds 2025-11-27 20:08:23 +02:00
kijai 44feb24290 Allow TTM to work with 5B models 2025-11-20 11:43:53 +02:00
kijai fa7a967ee7 Fix stand-in 2025-11-17 11:56:02 +02:00
kijai ec161373f4 Update readme.md 2025-11-16 19:30:09 +02:00
kijai 0e3fd0b491 Create wanvideo2_2_I2V_A14B_TimeToMove_example.json 2025-11-16 19:26:51 +02:00
kijai f872460285 Fix TTM for dual sampler setups 2025-11-16 19:23:19 +02:00
kijai b826642a83 Add TTM support (Time To Move)
https://github.com/time-to-move/TTM
2025-11-16 17:52:49 +02:00
jamesjjcondon aa9f474958 Refactor data_mean and data_std initialization
Refactor data_mean and data_std to use register_buffer and remove nn.Buffer.

Tested with pytorch version: 2.4.0+cu121
xformers version: 0.0.27.post2
Set vram state to: LOW_VRAM
Device: cuda:0 NVIDIA GeForce RTX 3090 : cudaMallocAsync
Enabled pinned memory 122281.0
Using xformers attention
Python version: 3.10.12 (main, Aug 15 2025, 14:32:43) [GCC 11.4.0]
ComfyUI version: 0.3.68
ComfyUI frontend version: 1.28.8
2025-11-14 11:57:26 +10:30
kijai e3c2a1431b Squashed commit of the following:
commit f685ee33ac
Merge: bb5707f 4e31081
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Nov 13 16:37:38 2025 +0200

    Merge branch 'main' into bindweave

commit bb5707f601
Merge: acb662b ff26836
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Nov 11 18:53:19 2025 +0200

    Merge branch 'main' into bindweave

commit acb662b5af
Merge: 907c9e1 e926f7a
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Nov 11 11:44:26 2025 +0200

    Merge branch 'main' into bindweave

commit 907c9e1cdd
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Nov 10 21:02:58 2025 +0200

    Update nodes.py

commit e4a4d22537
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Nov 8 16:04:55 2025 +0200

    Update nodes_sampler.py

commit a3b2f67337
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Nov 8 16:03:00 2025 +0200

    Pad clip vision embeds like in original code

commit 1e00c8fb28
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Nov 8 12:21:11 2025 +0200

    Update nodes.py

commit ff16dce5c0
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Nov 7 01:15:13 2025 +0200

    Update nodes_sampler.py

commit f972b31bf2
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Nov 7 01:09:22 2025 +0200

    Update nodes.py

commit 3dacd6a719
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Nov 7 00:45:06 2025 +0200

    Update nodes_sampler.py

commit 7bf99791ad
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Nov 7 00:44:13 2025 +0200

    Update nodes.py

commit 7a5587b5af
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Nov 7 00:39:10 2025 +0200

    Let the user resize for QwenVL

    Seems to need smaller resolutions

commit d6cf172846
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Nov 6 23:41:24 2025 +0200

    Update nodes.py

commit cf86f4f0a4
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Nov 6 21:01:03 2025 +0200

    Update model.py

commit b1f8309a20
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Nov 6 19:55:54 2025 +0200

    Update nodes_model_loading.py

commit 8992c6af64
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Nov 6 19:22:56 2025 +0200

    Don't include padding for scheduler

commit e4084a961b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Nov 6 18:32:53 2025 +0200

    Update nodes.py

commit 3ec1edefbe
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Nov 6 17:35:48 2025 +0200

    init

    For testing, no idea if it works yet
2025-11-13 17:10:57 +02:00
wzxysf 6a37c0b2d6 Enable vae tiling with end frame 2025-10-24 22:23:31 +08:00
cmeka f653544f7b Merge branch 'kijai:main' into custom-sigmas 2025-10-21 12:48:02 -04:00
cmeka 425035d810 Fix custom sigmas for supported schedulers
Previously, only unipc, dpm++, and dpm++_sde schedulers preserved custom input sigmas exactly. Other schedulers such as (euler, lcm, deis,
  etc.) would transform or modify the sigmas through their set_timesteps() methods, causing inconsistent behavior.
2025-10-21 12:46:58 -04:00
91 changed files with 37934 additions and 12079 deletions
-1
View File
@@ -1 +0,0 @@
github: [kijai]
+149 -2
View File
@@ -2,6 +2,10 @@ import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
import torch import torch
import math import math
from einops import rearrange
from ..wanvideo.modules.model import WanRMSNorm, attention
from ..multitalk.multitalk import RotaryPositionalEmbedding1D, normalize_and_scale
class FeedForwardSwiGLU(nn.Module): class FeedForwardSwiGLU(nn.Module):
def __init__( def __init__(
@@ -22,7 +26,7 @@ class FeedForwardSwiGLU(nn.Module):
def forward(self, x): def forward(self, x):
return self.w2(F.silu(self.w1(x)) * self.w3(x)) return self.w2(F.silu(self.w1(x)) * self.w3(x))
class TimestepEmbedder(nn.Module): class TimestepEmbedder(nn.Module):
""" """
Embeds scalar timesteps into vector representations. Embeds scalar timesteps into vector representations.
@@ -62,4 +66,147 @@ class TimestepEmbedder(nn.Module):
if t_freq.dtype != dtype: if t_freq.dtype != dtype:
t_freq = t_freq.to(dtype) t_freq = t_freq.to(dtype)
t_emb = self.mlp(t_freq) t_emb = self.mlp(t_freq)
return t_emb return t_emb
class SingleStreamAttention(nn.Module):
def __init__(
self,
dim: int,
encoder_hidden_states_dim: int,
num_heads: int,
qkv_bias: bool,
qk_norm: bool,
attn_drop: float = 0.0,
proj_drop: float = 0.0,
eps: float = 1e-6,
class_range: int = 24,
class_interval: int = 4,
attention_mode: str = "sdpa",
) -> None:
super().__init__()
assert dim % num_heads == 0, "dim should be divisible by num_heads"
self.dim = dim
self.encoder_hidden_states_dim = encoder_hidden_states_dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.scale = self.head_dim**-0.5
self.q_linear = nn.Linear(dim, dim, bias=qkv_bias)
self.q_norm = WanRMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
self.attn_drop = nn.Dropout(attn_drop)
self.proj = nn.Linear(dim, dim)
self.proj_drop = nn.Dropout(proj_drop)
self.kv_linear = nn.Linear(encoder_hidden_states_dim, dim * 2, bias=qkv_bias)
self.k_norm = WanRMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
self.attention_mode = attention_mode
# multitalk related params
self.class_interval = class_interval
self.class_range = class_range
self.rope_h1 = (0, self.class_interval)
self.rope_h2 = (self.class_range - self.class_interval, self.class_range)
self.rope_bak = int(self.class_range // 2)
self.rope_1d = RotaryPositionalEmbedding1D(self.head_dim)
def _process_cross_attn(self, x, cond, frames_num=None, x_ref_attn_map=None):
N_t = frames_num
out_dtype = x.dtype
x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t)
# get q for hidden_state
B, N, C = x.shape
q = self.q_linear(x)
q_shape = (B, N, self.num_heads, self.head_dim)
q = q.view(q_shape).permute((0, 2, 1, 3)) # [B, H, N, D]
q = self.q_norm(q.to(self.q_norm.weight.dtype)).to(q.dtype)
# multitalk with rope1d pe
if x_ref_attn_map is not None:
max_values = x_ref_attn_map.max(1).values[:, None, None]
min_values = x_ref_attn_map.min(1).values[:, None, None]
max_min_values = torch.cat([max_values, min_values], dim=2)
human1_max_value, human1_min_value = max_min_values[0, :, 0].max(), max_min_values[0, :, 1].min()
human2_max_value, human2_min_value = max_min_values[1, :, 0].max(), max_min_values[1, :, 1].min()
human1 = normalize_and_scale(x_ref_attn_map[0], (human1_min_value, human1_max_value), (self.rope_h1[0], self.rope_h1[1]))
human2 = normalize_and_scale(x_ref_attn_map[1], (human2_min_value, human2_max_value), (self.rope_h2[0], self.rope_h2[1]))
back = torch.full((x_ref_attn_map.size(1),), self.rope_bak, dtype=human1.dtype).to(human1.device)
max_indices = x_ref_attn_map.argmax(dim=0)
normalized_map = torch.stack([human1, human2, back], dim=1)
normalized_pos = normalized_map[range(x_ref_attn_map.size(1)), max_indices]
q = rearrange(q, "(B N_t) H S C -> B H (N_t S) C", N_t=N_t)
q = self.rope_1d(q, normalized_pos)
q = rearrange(q, "B H (N_t S) C -> (B N_t) H S C", N_t=N_t)
# get kv from encoder_hidden_states
_, N_a, _ = cond.shape
encoder_kv = self.kv_linear(cond)
encoder_kv_shape = (B, N_a, 2, self.num_heads, self.head_dim)
encoder_kv = encoder_kv.view(encoder_kv_shape).permute((2, 0, 3, 1, 4))
encoder_k, encoder_v = encoder_kv.unbind(0)
encoder_k = self.k_norm(encoder_k.to(self.k_norm.weight.dtype)).to(encoder_k.dtype)
# multitalk with rope1d pe
if x_ref_attn_map is not None:
per_frame = torch.zeros(N_a, dtype=encoder_k.dtype).to(encoder_k.device)
per_frame[:per_frame.size(0)//2] = (self.rope_h1[0] + self.rope_h1[1]) / 2
per_frame[per_frame.size(0)//2:] = (self.rope_h2[0] + self.rope_h2[1]) / 2
encoder_pos = torch.concat([per_frame]*N_t, dim=0)
encoder_k = rearrange(encoder_k, "(B N_t) H S C -> B H (N_t S) C", N_t=N_t)
encoder_k = self.rope_1d(encoder_k, encoder_pos)
encoder_k = rearrange(encoder_k, "B H (N_t S) C -> (B N_t) H S C", N_t=N_t)
# Input tensors must be in format ``[B, M, H, K]``, where B is the batch size, M \
# the sequence length, H the number of heads, and K the embeding size per head
q = rearrange(q, "B H M K -> B M H K")
encoder_k = rearrange(encoder_k, "B H M K -> B M H K")
encoder_v = rearrange(encoder_v, "B H M K -> B M H K")
x = attention(q, encoder_k, encoder_v, attention_mode=self.attention_mode)
x = rearrange(x, "B M H K -> B H M K")
# linear transform
x_output_shape = (B, N, C)
x = x.transpose(1, 2)
x = x.reshape(x_output_shape)
x = self.proj(x)
x = self.proj_drop(x)
# reshape x to origin shape
x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t)
return x.type(out_dtype)
def forward(self, x, cond, num_latent_frames=None, num_cond_latents=None, x_ref_attn_map=None, human_num=None):
B, N, C = x.shape
if (num_cond_latents is None or num_cond_latents == 0):
# text to video
output = self._process_cross_attn(x, cond, num_latent_frames, x_ref_attn_map)
return None, output
elif num_cond_latents is not None and num_cond_latents > 0:
# image to video or video continuation
num_cond_latents_thw = num_cond_latents * (N // num_latent_frames)
x_noise = x[:, num_cond_latents_thw:]
cond = rearrange(cond, "(B N_t) M C -> B N_t M C", B=B)
cond = cond[:, num_cond_latents:]
cond = rearrange(cond, "B N_t M C -> (B N_t) M C")
frames_num = num_latent_frames - num_cond_latents
if human_num is not None and human_num == 2:
# multitalk mode
output_noise = self._process_cross_attn(x_noise, cond, frames_num, x_ref_attn_map)
else:
# singletalk mode
output_noise = self._process_cross_attn(x_noise, cond, frames_num)
output_cond = torch.zeros((B, num_cond_latents_thw, C), dtype=output_noise.dtype, device=output_noise.device)
return output_cond, output_noise
else:
raise NotImplementedError
+308
View File
@@ -0,0 +1,308 @@
import torch
import torch.nn.functional as F
from ..utils import log
import comfy.model_management as mm
from comfy_api.latest import io
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="WanVideoLongCatAvatarExtendEmbeds",
category="WanVideoWrapper",
inputs=[
io.Latent.Input("prev_latents", tooltip="Full previous latents to be used to continue generation, continuation frames are selected based on 'overlap' parameter"),
io.Custom("MULTITALK_EMBEDS").Input("audio_embeds", tooltip="Full length audio embeddings"),
io.Int.Input("num_frames", default=93, min=1, max=256, step=1, tooltip="Number of new frames to generate"),
io.Int.Input("overlap", default=13, min=0, max=16, step=1, tooltip="Number of overlapping frames from previous latents for video continuation, set to 0 for T2V"),
io.Int.Input("frames_processed", default=0, min=0, max=10000, step=1, tooltip="Number of frames already processed in the video, used to select audio features"),
io.Combo.Input("if_not_enough_audio", ["pad_with_start", "mirror_from_end"], default="pad_with_start", tooltip="What to do if there are not enough frames in pose_images for the window"),
io.Int.Input("ref_frame_index", default=10, min=0, max=1000, step=1, tooltip="Values between 0 - 24 ensures better consistency, while selecting other ranges (e.g., -10 or 30) helps reduce repeated actions"),
io.Int.Input("ref_mask_frame_range", default=3, min=0, max=20, step=1, tooltip="Larger range can further help mitigate repeated actions, but excessively large values may introduce artifacts"),
io.Latent.Input("ref_latent", optional=True, tooltip="Reference latent used for consistency, generally should be either the init image, or first latent from first generation"),
io.Latent.Input("samples", optional=True, tooltip="For the sampler 'samples' input, used for slicing samples per window for vid2vid"),
io.Custom("IMAGE").Input("prev_images", optional=True, tooltip="LongCat-Avatar-1.5: decoded frames from the previous segment. When provided together with `vae`, the trailing `overlap` frames are re-encoded through the VAE and used as the overlap conditioning (matches v1.5's use_vcond=False behavior). Leave disconnected for v1.0."),
io.Custom("WANVAE").Input("vae", optional=True, tooltip="LongCat-Avatar-1.5: VAE used to re-encode `prev_images` for the overlap region. Only used when `prev_images` is also provided."),
],
outputs=[
io.Custom("WANVIDIMAGE_EMBEDS").Output(display_name="image_embeds", tooltip="Embeds for WanVideo LongCat Avatar generation"),
io.Latent.Output(display_name="samples_slice", tooltip="Sliced latent samples for the new frames"),
],
)
@classmethod
def execute(cls, prev_latents, audio_embeds, num_frames, overlap, if_not_enough_audio, frames_processed, ref_frame_index, ref_mask_frame_range, ref_latent=None, samples=None, prev_images=None, vae=None) -> io.NodeOutput:
new_audio_embed = audio_embeds.copy()
audio_features = torch.stack(new_audio_embed["audio_features"])
num_audio_features = audio_features.shape[1]
if audio_features.shape[1] < frames_processed + num_frames:
deficit = frames_processed + num_frames - audio_features.shape[1]
if if_not_enough_audio == "pad_with_start":
pad = audio_features[:, :1].repeat(1, deficit, 1, 1)
audio_features = torch.cat([audio_features, pad], dim=1)
elif if_not_enough_audio == "mirror_from_end":
to_add = audio_features[:, -deficit:, :].flip(dims=[1])
audio_features = torch.cat([audio_features, to_add], dim=1)
log.warning(f"Not enough audio features, padded with strategy '{if_not_enough_audio}' from {num_audio_features} to {audio_features.shape[1]} frames")
ref_target_masks = new_audio_embed.get("ref_target_masks", None)
if ref_target_masks is not None:
new_audio_embed["ref_target_masks"] = ref_target_masks[:, frames_processed:frames_processed+num_frames, :]
prev_samples = prev_latents["samples"].clone()
if overlap != 0:
latent_overlap = (overlap - 1) // 4 + 1
if prev_images is not None and vae is not None:
# LongCat-Avatar-1.5 path: re-encodes instead of just slicing
img = prev_images[-overlap:]
if img.shape[-1] == 4:
img = img[..., :3]
img = img.to(vae.dtype).to(device) * 2.0 - 1.0
img = img.permute(3, 0, 1, 2).unsqueeze(0).contiguous() # [T, H, W, C] -> [B, C, T, H, W]
vae.to(device)
prev_samples = vae.encode(img, device=device).to(prev_samples)
vae.to(offload_device)
mm.soft_empty_cache()
log.info(f"Re-encoded {overlap} overlap frames -> latent shape {tuple(prev_samples.shape)}")
else:
prev_samples = prev_samples[:, :, -latent_overlap:]
ref_sample = None
if ref_latent is not None:
ref_sample = ref_latent["samples"][0, :, :1].clone()
log.info(f"Previous latents shape: {prev_samples.shape}, using last {latent_overlap} latent frames for overlap.")
new_latent_frames = (num_frames - 1) // 4 + 1
target_shape = (16, new_latent_frames, prev_samples.shape[-2], prev_samples.shape[-1])
audio_stride = new_audio_embed.get("audio_stride", 2)
indices = torch.arange(2 * 2 + 1) - 2
if frames_processed == 0:
audio_start_idx = 0
else:
audio_start_idx = (frames_processed - overlap) * audio_stride
audio_end_idx = audio_start_idx + num_frames * audio_stride
log.info(f"Extracting audio embeddings from index {audio_start_idx} to {audio_end_idx}")
audio_embs = []
for human_idx in range(len(audio_features)):
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
center_indices = torch.clamp(center_indices, min=0, max=audio_features[human_idx].shape[0] - 1)
audio_emb = audio_features[human_idx][center_indices].unsqueeze(0).to(device)
audio_embs.append(audio_emb)
audio_emb = torch.cat(audio_embs, dim=0)
new_audio_embed["audio_features"] = None
new_audio_embed["audio_emb_slice"] = audio_emb
longcat_avatar_options = {
"longcat_ref_latent": ref_sample,
"ref_frame_index": ref_frame_index,
"ref_mask_frame_range": ref_mask_frame_range,
}
embeds = {
"target_shape": target_shape,
"num_frames": num_frames,
"extra_latents": [{"samples": prev_samples, "index": 0}] if overlap != 0 else None,
"multitalk_embeds": new_audio_embed,
"longcat_avatar_options": longcat_avatar_options,
}
samples_slice = None
if samples is not None:
latent_start_index = (frames_processed - 1) // 4 + 1 if frames_processed > 0 else 0
latent_end_index = latent_start_index + new_latent_frames
samples_slice = samples.copy()
samples_slice["samples"] = samples["samples"][:, :, latent_start_index:latent_end_index].clone()
return io.NodeOutput(embeds, samples_slice)
class LongCatAvatarWhisperEmbeds:
"""Audio embeds for LongCat-Video-Avatar-1.5 (Whisper-large-v3).
Produces a MULTITALK_EMBEDS dict whose audio_features are shaped [T, 5, 1280]
(5 grouped Whisper layers, 1280-d hidden state), matching the audio stream
the v1.5 AudioProjModel expects. audio_stride is set to 1 to signal v1.5
timing to the consumer nodes (vs. 2 for the v1.0 wav2vec2 path).
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"whisper_model": ("WHISPERMODEL",),
"audio_1": ("AUDIO",),
"normalize_loudness": ("BOOLEAN", {"default": True, "tooltip": "Normalize audio loudness to -23 LUFS before encoding (matches the v1.5 reference pipeline)"}),
"num_frames": ("INT", {"default": 93, "min": 1, "max": 10000, "step": 1, "tooltip": "Total frame count to generate; bounds how much audio is consumed"}),
"fps": ("FLOAT", {"default": 25.0, "min": 1.0, "max": 60.0, "step": 0.1, "tooltip": "Target video fps. LongCat-Video-Avatar-1.5 is trained at 25 fps."}),
"audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the audio conditioning"}),
"audio_cfg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "When not 1.0, an extra model pass without audio conditioning is done"}),
"multi_audio_type": (["para", "add"], {"default": "para", "tooltip": "'para' overlays speakers in parallel (equal length); 'add' concatenates speakers sequentially with silence padding"}),
},
"optional": {
"audio_2": ("AUDIO",),
"audio_3": ("AUDIO",),
"audio_4": ("AUDIO",),
"ref_target_masks": ("MASK", {"tooltip": "Per-speaker semantic mask(s) in pixel space, one per speaker"}),
},
}
RETURN_TYPES = ("MULTITALK_EMBEDS", "AUDIO", "INT",)
RETURN_NAMES = ("multitalk_embeds", "audio", "num_frames",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, whisper_model, audio_1, normalize_loudness, num_frames, fps,
audio_scale, audio_cfg_scale, multi_audio_type,
audio_2=None, audio_3=None, audio_4=None, ref_target_masks=None):
import torchaudio
import numpy as np
from ..multitalk.nodes import loudness_norm
model = whisper_model["model"]
feature_extractor = whisper_model["feature_extractor"]
dtype = whisper_model["dtype"]
sr = 16000
MEL_CHUNK = 750 * 640 # 480000 samples = 30s at 16kHz; matches Whisper's chunk_length
ENC_CHUNK = 3000 # encoder window in mel frames
ENC_FPS = 50 # whisper encoder output frames per second
def linear_interp(features, output_len):
features = features.transpose(1, 2) # [B, D, T]
out = F.interpolate(features, size=output_len, align_corners=True, mode='linear')
return out.transpose(1, 2)
audio_inputs = [a for a in [audio_1, audio_2, audio_3, audio_4] if a is not None]
audio_features_list = []
seq_lengths = []
audio_outputs = []
end_time = num_frames / float(fps)
end_sample = int(end_time * sr)
for audio in audio_inputs:
audio_input = audio["waveform"]
sample_rate = audio["sample_rate"]
if sample_rate != sr:
audio_input = torchaudio.functional.resample(audio_input, sample_rate, sr)
audio_input = audio_input[0][0]
audio_segment = audio_input[:end_sample].cpu().numpy().astype(np.float32)
if normalize_loudness:
audio_segment = loudness_norm(audio_segment, sr=sr)
audio_duration = len(audio_segment) / sr
video_length = int(audio_duration * fps)
if video_length < 1:
continue
mel_chunks = []
for i in range(0, len(audio_segment), MEL_CHUNK):
mel = feature_extractor(audio_segment[i:i + MEL_CHUNK], sampling_rate=sr,
return_tensors="pt").input_features
mel_chunks.append(mel)
mel_features = torch.cat(mel_chunks, dim=-1).to(device=device, dtype=dtype)
model.to(device)
enc_chunks = []
with torch.no_grad():
for i in range(0, mel_features.shape[-1], ENC_CHUNK):
chunk = mel_features[:, :, i:i + ENC_CHUNK]
chunk_hs = model.encoder(chunk, output_hidden_states=True).hidden_states
enc_chunks.append(torch.stack(chunk_hs, dim=2)) # [1, T_enc, n_layers+1, D]
model.to(offload_device)
audio_prompts = torch.cat(enc_chunks, dim=1)
audio_prompts = audio_prompts[:, :video_length * 2]
feat0 = linear_interp(audio_prompts[:, :, 0:8].mean(dim=2), video_length)
feat1 = linear_interp(audio_prompts[:, :, 8:16].mean(dim=2), video_length)
feat2 = linear_interp(audio_prompts[:, :, 16:24].mean(dim=2), video_length)
feat3 = linear_interp(audio_prompts[:, :, 24:32].mean(dim=2), video_length)
feat4 = linear_interp(audio_prompts[:, :, 32], video_length)
audio_emb = torch.stack([feat0, feat1, feat2, feat3, feat4], dim=2)[0] # [T, 5, 1280]
audio_features_list.append(audio_emb.cpu().detach())
seq_lengths.append(audio_emb.shape[0])
waveform_tensor = torch.from_numpy(audio_segment).float().unsqueeze(0).unsqueeze(0)
audio_outputs.append({"waveform": waveform_tensor, "sample_rate": sr})
if len(audio_features_list) == 0:
raise RuntimeError("No valid Whisper audio embeddings extracted, please check inputs")
if len(audio_features_list) > 1:
if multi_audio_type == "para":
max_len = max(seq_lengths)
padded = []
for emb in audio_features_list:
if emb.shape[0] < max_len:
pad = torch.zeros(max_len - emb.shape[0], *emb.shape[1:], dtype=emb.dtype)
emb = torch.cat([emb, pad], dim=0)
padded.append(emb)
audio_features_list = padded
else: # "add"
total_len = sum(seq_lengths)
full_list = []
offset = 0
for emb, length in zip(audio_features_list, seq_lengths):
full = torch.zeros(total_len, *emb.shape[1:], dtype=emb.dtype)
full[offset:offset + length] = emb
full_list.append(full)
offset += length
audio_features_list = full_list
multitalk_embeds = {
"audio_features": audio_features_list,
"audio_scale": audio_scale,
"audio_cfg_scale": audio_cfg_scale,
"ref_target_masks": ref_target_masks,
"audio_stride": 1,
"audio_encoder_type": "whisper",
}
if len(audio_outputs) == 1:
out_audio = audio_outputs[0]
elif multi_audio_type == "para":
max_len = max(a["waveform"].shape[-1] for a in audio_outputs)
mixed = torch.zeros(1, 1, max_len, dtype=audio_outputs[0]["waveform"].dtype)
for a in audio_outputs:
w = a["waveform"]
if w.shape[-1] < max_len:
w = F.pad(w, (0, max_len - w.shape[-1]))
mixed += w
out_audio = {"waveform": mixed, "sample_rate": sr}
else:
total_len = sum(a["waveform"].shape[-1] for a in audio_outputs)
mixed = torch.zeros(1, 1, total_len, dtype=audio_outputs[0]["waveform"].dtype)
offset = 0
for a in audio_outputs:
w = a["waveform"]
mixed[:, :, offset:offset + w.shape[-1]] += w
offset += w.shape[-1]
out_audio = {"waveform": mixed, "sample_rate": sr}
return (multitalk_embeds, out_audio, num_frames)
NODE_CLASS_MAPPINGS = {
"WanVideoLongCatAvatarExtendEmbeds": WanVideoLongCatAvatarExtendEmbeds,
"LongCatAvatarWhisperEmbeds": LongCatAvatarWhisperEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoLongCatAvatarExtendEmbeds": "WanVideo LongCat Avatar Extend Embeds",
"LongCatAvatarWhisperEmbeds": "LongCat Avatar Whisper Embeds (v1.5)",
}
+199
View File
@@ -0,0 +1,199 @@
import torch
import torch.nn as nn
from einops import rearrange
from ..wanvideo.modules.attention import attention
def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor):
return (x * (1 + scale) + shift)
def sinusoidal_embedding_1d(dim, position):
sinusoid = torch.outer(position.type(torch.float64), torch.pow(
10000, -torch.arange(dim//2, dtype=torch.float64, device=position.device).div(dim//2)))
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
return x.to(position.dtype)
def precompute_freqs_cis_3d(dim: int, end: int = 1024, theta: float = 10000.0):
# 3d rope precompute
f_freqs_cis = precompute_freqs_cis(dim - 2 * (dim // 3), end, theta)
h_freqs_cis = precompute_freqs_cis(dim // 3, end, theta)
w_freqs_cis = precompute_freqs_cis(dim // 3, end, theta)
return f_freqs_cis, h_freqs_cis, w_freqs_cis
def precompute_freqs_cis(dim: int, end: int = 1024, theta: float = 10000.0):
# 1d rope precompute
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)
[: (dim // 2)].double() / dim))
freqs = torch.outer(torch.arange(end, device=freqs.device), freqs)
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64
return freqs_cis
def rope_apply(x, freqs, num_heads):
x = rearrange(x, "b s (n d) -> b s n d", n=num_heads)
x_out = torch.view_as_complex(x.to(torch.float64).reshape(
x.shape[0], x.shape[1], x.shape[2], -1, 2))
x_out = torch.view_as_real(x_out * freqs).flatten(2)
return x_out.to(x.dtype)
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-5):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def norm(self, x):
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
def forward(self, x):
dtype = x.dtype
return self.norm(x.float()).to(dtype) * self.weight
class AttentionModule(nn.Module):
def __init__(self, num_heads, head_dim):
super().__init__()
self.num_heads = num_heads
self.head_dim = head_dim
def forward(self, q, k, v):
b, n, d = q.size(0), self.num_heads, self.head_dim
x = attention(
q.view(b, -1, n, d),
k.view(b, -1, n, d),
v.view(b, -1, n, d)
)
return x.flatten(2)
class SelfAttention(nn.Module):
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6):
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.q = nn.Linear(dim, dim)
self.k = nn.Linear(dim, dim)
self.v = nn.Linear(dim, dim)
self.o = nn.Linear(dim, dim)
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
self.attn = AttentionModule(self.num_heads, self.head_dim)
def forward(self, x, freqs):
q = self.norm_q(self.q(x))
k = self.norm_k(self.k(x))
v = self.v(x)
q = rope_apply(q, freqs, self.num_heads)
k = rope_apply(k, freqs, self.num_heads)
x = self.attn(q, k, v)
return self.o(x)
class CrossAttention(nn.Module):
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6, clip_fea: torch.Tensor = None):
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.q = nn.Linear(dim, dim)
self.k = nn.Linear(dim, dim)
self.v = nn.Linear(dim, dim)
self.o = nn.Linear(dim, dim)
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
self.k_img = nn.Linear(dim, dim)
self.v_img = nn.Linear(dim, dim)
self.norm_k_img = RMSNorm(dim, eps=eps)
self.attn = AttentionModule(self.num_heads, self.head_dim)
def forward(self, x: torch.Tensor, y: torch.Tensor, clip_fea: torch.Tensor = None):
ctx = y
q = self.norm_q(self.q(x))
k = self.norm_k(self.k(ctx))
v = self.v(ctx)
x = self.attn(q, k, v)
if clip_fea is not None:
k_img = self.norm_k_img(self.k_img(clip_fea))
v_img = self.v_img(clip_fea)
y = self.attn(q, k_img, v_img)
x = x + y
return self.o(x)
class GateModule(nn.Module):
def __init__(self,):
super().__init__()
def forward(self, x, gate, residual):
return x + gate * residual
class DiTBlock(nn.Module):
def __init__(self, dim: int, num_heads: int, ffn_dim: int, eps: float = 1e-6):
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.ffn_dim = ffn_dim
self.self_attn = SelfAttention(dim, num_heads, eps)
self.cross_attn = CrossAttention(dim, num_heads, eps)
self.norm1 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
self.norm2 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
self.norm3 = nn.LayerNorm(dim, eps=eps)
self.ffn = nn.Sequential(nn.Linear(dim, ffn_dim), nn.GELU(
approximate='tanh'), nn.Linear(ffn_dim, dim))
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
self.gate = GateModule()
def forward(self, x, context, t_mod, freqs, clip_fea=None):
has_seq = len(t_mod.shape) == 4
chunk_dim = 2 if has_seq else 1
# msa: multi-head self-attention mlp: multi-layer perceptron
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
self.modulation.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod).chunk(6, dim=chunk_dim)
if has_seq:
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
shift_msa.squeeze(2), scale_msa.squeeze(2), gate_msa.squeeze(2),
shift_mlp.squeeze(2), scale_mlp.squeeze(2), gate_mlp.squeeze(2),
)
input_x = modulate(self.norm1(x), shift_msa, scale_msa)
x = self.gate(x, gate_msa, self.self_attn(input_x, freqs))
x = x + self.cross_attn(self.norm3(x), context, clip_fea=clip_fea)
input_x = modulate(self.norm2(x), shift_mlp, scale_mlp)
x = self.gate(x, gate_mlp, self.ffn(input_x))
return x
class WanModelDualControl(torch.nn.Module):
def __init__(self, dim: int, ffn_dim: int, eps: float, num_heads: int, control_layers = 12):
super().__init__()
self.control_layers = control_layers
self.control_blocks_dense = nn.ModuleList([
DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps)
for _ in range(self.control_layers)
])
self.control_blocks_sparse = nn.ModuleList([
DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps)
for _ in range(self.control_layers)
])
self.control_initial_combine_linear_dense = torch.nn.Linear(dim, dim//2)
self.control_initial_combine_linear_sparse = torch.nn.Linear(dim, dim//2)
self.control_text_linear = torch.nn.Linear(dim, dim//2)
self.control_t_mod = torch.nn.Linear(dim, dim//2)
self.control_combine_linears = torch.nn.ModuleList([torch.nn.Linear(dim//2, dim) for _ in range(self.control_layers)])
head_dim = dim // num_heads
self.freqs = precompute_freqs_cis_3d(head_dim)
+88
View File
@@ -0,0 +1,88 @@
import torch
from ..utils import log
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoAddDualControlEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"vae": ("WANVAE", {"tooltip": "VAE model"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
"first_frame_noise_level": ("FLOAT", {"default": 0.925926, "min": 0.0, "max": 1.0, "step": 0.000001, "tooltip": "Noise level for the first frame when using previous frames"}),
},
"optional": {
"dense": ("IMAGE", {"tooltip": "Dense control signal (depth) video input"}),
"sparse": ("IMAGE", {"tooltip": "Sparse control signal (tracks) video input"}),
"prev_images": ("IMAGE", {"tooltip": "Previous frames for temporal consistency, default is 8 frames"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, vae, strength, start_percent, end_percent, first_frame_noise_level, dense=None, sparse=None, prev_images=None):
updated = dict(embeds)
updated.setdefault("dual_control", {})
if dense is None and sparse is None:
raise ValueError("At least one of dense or sparse inputs must be provided.")
num_frames = dense.shape[0] if dense is not None else sparse.shape[0]
height = dense.shape[1] if dense is not None else sparse.shape[1]
width = dense.shape[2] if dense is not None else sparse.shape[2]
msk = torch.ones(1, num_frames, height//8, width//8, device=device)
msk[:, 1:] = 0
msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
msk = msk.view(1, msk.shape[1] // 4, 4, height//8, width//8)
msk = msk.transpose(1, 2)
dense_input_latent = sparse_input_latent = None
vae.to(device)
if dense is not None:
dense_images = 1 - dense[..., :3] # Invert colors for depth to match the usual range in comfy
dense_images = dense_images.permute(3, 0, 1, 2) * 2 - 1
dense_video_latent = vae.encode([dense_images.to(device, vae.dtype)], device, tiled=False)
dense_first = (dense_images[:, :1]).to(device, vae.dtype)
vae_input_dense = torch.cat([dense_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1)
dense_concat_latent = vae.encode([vae_input_dense], device, tiled=False)
dense_concat_latent = torch.cat([msk, dense_concat_latent], dim=1)
dense_input_latent = torch.cat([dense_video_latent, dense_concat_latent],dim=1)
if sparse is not None:
sparse_images = sparse[..., :3].permute(3, 0, 1, 2) * 2 - 1
sparse_video_latent = vae.encode([sparse_images.to(device, vae.dtype)], device, tiled=False)
sparse_first = (sparse_images[:, :1]).to(device, vae.dtype)
vae_input_sparse = torch.cat([sparse_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1)
sparse_concat_latent = vae.encode([vae_input_sparse], device, tiled=False)
sparse_concat_latent = torch.cat([msk, sparse_concat_latent], dim=1)
sparse_input_latent = torch.cat([sparse_video_latent, sparse_concat_latent],dim=1)
if prev_images is not None:
prev_images = prev_images[..., :3].permute(3, 0, 1, 2) * 2 - 1
prev_video_latent = vae.encode([prev_images.to(device, vae.dtype)], device, tiled=False)
updated["dual_control"]["prev_latent"] = prev_video_latent[0]
vae.to(offload_device)
updated["dual_control"]["dense_input_latent"] = dense_input_latent
updated["dual_control"]["sparse_input_latent"] = sparse_input_latent
updated["dual_control"]["strength"] = strength
updated["dual_control"]["start_percent"] = start_percent
updated["dual_control"]["end_percent"] = end_percent
updated["dual_control"]["first_frame_noise_level"] = first_frame_noise_level
return (updated,)
NODE_CLASS_MAPPINGS = {
"WanVideoAddDualControlEmbeds": WanVideoAddDualControlEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddDualControlEmbeds": "WanVideo Add Dual Control Embeds",
}
+83 -12
View File
@@ -31,7 +31,7 @@ def p3d_to_p2d(point_3d, height, width): # point3d n*1024*3
def get_pose_images(smpl_data, offset): def get_pose_images(smpl_data, offset):
pose_images = [] pose_images = []
for data in smpl_data: for data in smpl_data:
if isinstance(data, np.ndarray): if isinstance(data, np.ndarray):
joints3d = data joints3d = data
else: else:
@@ -43,28 +43,33 @@ def get_pose_images(smpl_data, offset):
return pose_images return pose_images
def get_control_conditions(poses, h, w): def get_control_conditions(poses, h, w, stick_width=1.0, point_radius=2, style="original"):
video_transforms = transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True)
control_images = [] control_images = []
for idx, pose in enumerate(poses): for idx, pose in enumerate(poses):
canvas = np.zeros(shape=(h, w, 3), dtype=np.uint8) canvas = np.zeros(shape=(h, w, 3), dtype=np.uint8)
try: try:
joints3d = p3d_to_p2d(pose, h, w) joints3d = p3d_to_p2d(pose, h, w)
canvas = draw_3d_points( if style == "original":
canvas, canvas = draw_3d_points(
joints3d[0], canvas,
stickwidth=int(h / 350), joints3d[0],
) stickwidth=int(h / 350 * stick_width),
r=point_radius,
)
elif style == "scail":
canvas = draw_3d_points_scail(
canvas,
joints3d[0],
stickwidth=int(h / 350 * stick_width),
r=point_radius,
)
resized_canvas = cv2.resize(canvas, (w, h)) resized_canvas = cv2.resize(canvas, (w, h))
# Image.fromarray(resized_canvas).save(f'tmp/{idx}_pose.jpg') # Image.fromarray(resized_canvas).save(f'tmp/{idx}_pose.jpg')
control_images.append(resized_canvas) control_images.append(resized_canvas)
except Exception as e: except Exception:
print("wrong:", e)
control_images.append(Image.fromarray(canvas)) control_images.append(Image.fromarray(canvas))
control_pixel_values = np.array(control_images) control_pixel_values = np.array(control_images)
control_pixel_values = torch.from_numpy(control_pixel_values).contiguous() / 255. control_pixel_values = torch.from_numpy(control_pixel_values).contiguous() / 255.
print("control_pixel_values.shape", control_pixel_values.shape)
#control_pixel_values = video_transforms(control_pixel_values)
return control_pixel_values return control_pixel_values
@@ -140,3 +145,69 @@ def draw_3d_points(canvas, points, stickwidth=2, r=2, draw_line=True):
cv2.fillConvexPoly(canvas, polygon, connection_colors[i%17]) cv2.fillConvexPoly(canvas, polygon, connection_colors[i%17])
return canvas return canvas
def draw_3d_points_scail(canvas, points, stickwidth=2, r=2, draw_line=True):
connetions = [
[15,12],[12, 16],[16, 18],[18, 20],[20, 22], # 0-4: Left arm chain
[12,17],[17,19],[19,21], # 5-7: Right arm chain
[21,23], # 8: Right hand
[12,1],[1,4],[4,7], # 9-11: Neck to left leg (hip, thigh, shin)
[12,2],[2,5],[5,8], # 12-14: Neck to right leg (hip, thigh, shin)
]
# Warm colors for right side, cool colors for left side
connection_colors = [
[180, 180, 180], # 0: [15,12] - L. clavicle (Bright Cyan)
[0, 200, 255], # 1: [12,16] - L. shoulder (Bright Cyan)
[0, 120, 255], # 2: [16,18] - L. upper arm (Bright Blue)
[0, 60, 255], # 3: [18,20] - L. forearm (Deep Blue)
[60, 0, 255], # 4: [20,22] - L. hand (Blue-Purple)
[255, 0, 0], # 5: [12,17] - R. clavicle (Bright Red)
[255, 100, 0], # 6: [17,19] - R. upper arm (Bright Orange)
[255, 180, 0], # 7: [19,21] - R. forearm (Golden Orange)
[255, 255, 0], # 8: [21,23] - R. hand (Bright Yellow)
[30, 27, 160], # 9: [12,1] - Neck to L. hip (purple-blue)
[73, 27, 177], # 10: [1,4] - L. thigh (purple)
[145, 27, 194], # 11: [4,7] - L. shin (magenta)
[200, 255, 100], # 12: [12,2] - Neck to R. hip (yellow)
[54, 201, 52], # 13: [2,5] - R. thigh (green)
[30, 176, 85], # 14: [5,8] - R. shin (green)
]
# draw line
if draw_line:
# Collect all joints that are part of connections
joints_in_use = set()
for connection in connetions:
joints_in_use.add(connection[0])
joints_in_use.add(connection[1])
for i in range(len(connetions)):
point1_idx, point2_idx = connetions[i][0:2]
point1 = points[point1_idx]
point2 = points[point2_idx]
x1, y1 = int(point1[0]), int(point1[1])
x2, y2 = int(point2[0]), int(point2[1])
cv2.line(canvas, (x1, y1), (x2, y2), connection_colors[i], stickwidth)
# draw points for joints that have connections
joints_in_use = set()
for connection in connetions:
joints_in_use.add(connection[0])
joints_in_use.add(connection[1])
for joint_idx in joints_in_use:
if joint_idx >= len(points):
continue
x, y = points[joint_idx][0:2]
x, y = int(x), int(y)
# Use the color from the first connection involving this joint
joint_color = [180, 180, 180] # default grey
for i, connection in enumerate(connetions):
if connection[0] == joint_idx or connection[1] == joint_idx:
joint_color = connection_colors[i]
break
cv2.circle(canvas, (x, y), r, joint_color, thickness=-1)
return canvas
View File
+146 -48
View File
@@ -1,34 +1,56 @@
import os import os
import torch import torch
import gc from ..utils import log
from ..utils import log, dict_to_device
import numpy as np import numpy as np
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
import comfy.model_management as mm import comfy.model_management as mm
from comfy.utils import load_torch_file from comfy.utils import load_torch_file
import folder_paths import folder_paths
script_directory = os.path.dirname(os.path.abspath(__file__)) script_directory = os.path.dirname(os.path.abspath(__file__))
device = mm.get_torch_device() device = mm.get_torch_device()
offload_device = mm.unet_offload_device() offload_device = mm.unet_offload_device()
local_model_path = os.path.join(folder_paths.models_dir, "nlf", "nlf_l_multi_0.3.2.torchscript") local_model_path = os.path.join(folder_paths.models_dir, "nlf", "nlf_l_multi_0.3.2.torchscript")
folder_paths.add_model_folder_path("nlf", os.path.join(folder_paths.models_dir, "nlf"))
from .motion4d import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder from .motion4d import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder
from .mtv import prepare_motion_embeddings
def check_jit_script_function():
if torch.jit.script.__name__ != "script":
# Get more details about what modified it
module = torch.jit.script.__module__
qualname = getattr(torch.jit.script, '__qualname__', 'unknown')
code_file = None
try:
code_file = torch.jit.script.__code__.co_filename
code_line = torch.jit.script.__code__.co_firstlineno
log.warning(f"torch.jit.script has been modified by another custom node.\n"
f" Function name: {torch.jit.script.__name__}\n"
f" Module: {module}\n"
f" Qualified name: {qualname}\n"
f" Defined in: {code_file}:{code_line}\n"
f"This may cause issues with the NLF model.")
except Exception:
log.warning("--------------------------------")
log.warning(f"torch.jit.script function is: {torch.jit.script.__name__} from module {module}, "
f"this has been modified by another custom node. This may cause issues with the NLF model.")
log.warning("--------------------------------")
model_list = [
"https://github.com/isarandi/nlf/releases/download/v0.3.2/nlf_l_multi_0.3.2.torchscript",
"https://github.com/isarandi/nlf/releases/download/v0.2.2/nlf_l_multi_0.2.2.torchscript",
]
class DownloadAndLoadNLFModel: class DownloadAndLoadNLFModel:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return { return {
"required": { "required": {
"url": ( "url": (model_list, {"default": "https://github.com/isarandi/nlf/releases/download/v0.3.2/nlf_l_multi_0.3.2.torchscript"}),
[ },
"https://github.com/isarandi/nlf/releases/download/v0.3.2/nlf_l_multi_0.3.2.torchscript" "optional": {
], "warmup": ("BOOLEAN", {"default": True, "tooltip": "Whether to warmup the model after loading"}),
)
}, },
} }
@@ -37,8 +59,11 @@ class DownloadAndLoadNLFModel:
FUNCTION = "loadmodel" FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
def loadmodel(self, url): def loadmodel(self, url, warmup=True):
if url not in model_list:
raise ValueError(f"URL {url} is not in the list of allowed models.")
check_jit_script_function()
if not os.path.exists(local_model_path): if not os.path.exists(local_model_path):
log.info(f"Downloading NLF model to: {local_model_path}") log.info(f"Downloading NLF model to: {local_model_path}")
import requests import requests
@@ -52,6 +77,20 @@ class DownloadAndLoadNLFModel:
model = torch.jit.load(local_model_path).eval() model = torch.jit.load(local_model_path).eval()
if warmup:
log.info("Warming up NLF model...")
dummy_input = torch.zeros(1, 3, 256, 256, device=device)
jit_profiling_prev_state = torch._C._jit_set_profiling_executor(True)
try:
for _ in range(2):
_ = model.detect_smpl_batched(dummy_input)
finally:
torch._C._jit_set_profiling_executor(jit_profiling_prev_state)
log.info("NLF model warmed up")
model = model.to(offload_device)
return (model,) return (model,)
class LoadNLFModel: class LoadNLFModel:
@@ -59,8 +98,12 @@ class LoadNLFModel:
def INPUT_TYPES(s): def INPUT_TYPES(s):
return { return {
"required": { "required": {
"path": ("STRING", {"default": local_model_path}), "nlf_model": (folder_paths.get_filename_list("nlf"), {"tooltip": "These models are loaded from the 'ComfyUI/models/nlf' -folder",}),
}, },
"optional": {
"warmup": ("BOOLEAN", {"default": True, "tooltip": "Whether to warmup the model after loading"}),
},
} }
RETURN_TYPES = ("NLFMODEL",) RETURN_TYPES = ("NLFMODEL",)
@@ -68,8 +111,22 @@ class LoadNLFModel:
FUNCTION = "loadmodel" FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
def loadmodel(self, path): def loadmodel(self, nlf_model, warmup=True):
model = torch.jit.load(path).eval() check_jit_script_function()
model = torch.jit.load(folder_paths.get_full_path_or_raise("nlf", nlf_model)).eval()
if warmup:
log.info("Warming up NLF model...")
dummy_input = torch.zeros(1, 3, 256, 256, device=device)
jit_profiling_prev_state = torch._C._jit_set_profiling_executor(True)
try:
for _ in range(2):
_ = model.detect_smpl_batched(dummy_input)
finally:
torch._C._jit_set_profiling_executor(jit_profiling_prev_state)
log.info("NLF model warmed up")
model = model.to(offload_device)
return model, return model,
@@ -108,7 +165,7 @@ class LoadVQVAE:
frame_upsample_rate=[2.0, 2.0], frame_upsample_rate=[2.0, 2.0],
joint_upsample_rate=[1.0, 1.0] joint_upsample_rate=[1.0, 1.0]
) )
vqvae = SMPL_VQVAE(motion_encoder, motion_decoder, motion_quant).to(device) vqvae = SMPL_VQVAE(motion_encoder, motion_decoder, motion_quant).to(device)
vqvae.load_state_dict(vae_sd, strict=True) vqvae.load_state_dict(vae_sd, strict=True)
@@ -131,15 +188,6 @@ class MTVCrafterEncodePoses:
def encode(self, vqvae, poses): def encode(self, vqvae, poses):
# import pickle
# with open(os.path.join(script_directory, "data", "sampled_data.pkl"), 'rb') as f:
# data_list = pickle.load(f)
# if not isinstance(data_list, list):
# data_list = [data_list]
# print(data_list)
# smpl_poses = data_list[1]['pose']
global_mean = np.load(os.path.join(script_directory, "data", "mean.npy")) #global_mean.shape: (24, 3) global_mean = np.load(os.path.join(script_directory, "data", "mean.npy")) #global_mean.shape: (24, 3)
global_std = np.load(os.path.join(script_directory, "data", "std.npy")) global_std = np.load(os.path.join(script_directory, "data", "std.npy"))
@@ -153,7 +201,7 @@ class MTVCrafterEncodePoses:
vqvae.to(device) vqvae.to(device)
motion_tokens, vq_loss = vqvae(norm_poses.to(device), return_vq=True) motion_tokens, vq_loss = vqvae(norm_poses.to(device), return_vq=True)
recon_motion = vqvae(norm_poses.to(device))[0][0].to(dtype=torch.float32).cpu().detach() * global_std + global_mean recon_motion = vqvae(norm_poses.to(device))[0][0].to(dtype=torch.float32).cpu().detach() * global_std + global_mean
vqvae.to(offload_device) vqvae.to(offload_device)
@@ -162,7 +210,7 @@ class MTVCrafterEncodePoses:
'global_mean': global_mean, 'global_mean': global_mean,
'global_std': global_std 'global_std': global_std
} }
return poses_dict, recon_motion return poses_dict, recon_motion
@@ -173,32 +221,74 @@ class NLFPredict:
"model": ("NLFMODEL",), "model": ("NLFMODEL",),
"images": ("IMAGE", {"tooltip": "Input images for the model"}), "images": ("IMAGE", {"tooltip": "Input images for the model"}),
}, },
"optional": {
"per_batch": ("INT", {"default": -1, "min": -1, "max": 10000, "step": 1, "tooltip": "How many images to process at once. -1 means all at once."}),
}
} }
RETURN_TYPES = ("NLFPRED", ) RETURN_TYPES = ("NLFPRED", "BBOX",)
RETURN_NAMES = ("pose_results",) RETURN_NAMES = ("pose_results", "bboxes")
FUNCTION = "predict" FUNCTION = "predict"
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
def predict(self, model, images): def predict(self, model, images, per_batch=-1):
model.to(device)
pred = model.detect_smpl_batched(images.permute(0, 3, 1, 2).to(device))
model.to(offload_device)
pred = dict_to_device(pred, offload_device) check_jit_script_function()
model = model.to(device)
num_images = images.shape[0]
# Determine batch size
if per_batch == -1:
batch_size = num_images
else:
batch_size = per_batch
# Initialize result containers
all_boxes = []
all_joints3d_nonparam = []
# Process in batches
for i in range(0, num_images, batch_size):
end_idx = min(i + batch_size, num_images)
batch_images = images[i:end_idx]
jit_profiling_prev_state = torch._C._jit_set_profiling_executor(True)
try:
pred = model.detect_smpl_batched(batch_images.permute(0, 3, 1, 2).to(device))
finally:
torch._C._jit_set_profiling_executor(jit_profiling_prev_state)
# Collect boxes and joints from this batch
if 'boxes' in pred:
all_boxes.extend(pred['boxes'])
if 'joints3d_nonparam' in pred:
all_joints3d_nonparam.extend(pred['joints3d_nonparam'])
model = model.to(offload_device)
# Move collected results to offload device
all_boxes = [box.to(offload_device) for box in all_boxes]
all_joints3d_nonparam = [joints.to(offload_device) for joints in all_joints3d_nonparam]
# Maintain the original nested format: wrap in a list to match expected structure
pose_results = { pose_results = {
'joints3d_nonparam': [], 'joints3d_nonparam': [all_joints3d_nonparam],
} }
# Collect pose data
for key in pose_results.keys(): # Convert bboxes to list format: [x_min, y_min, x_max, y_max] for each detection
if key in pred: # Each box tensor is shape (1, 5) with [x_min, y_min, x_max, y_max, confidence]
pose_results[key].append(pred[key]) formatted_boxes = []
for box in all_boxes:
# Handle empty detections (no person detected in frame)
if box.numel() == 0 or box.shape[0] == 0:
formatted_boxes.append([0.0, 0.0, 0.0, 0.0])
else: else:
pose_results[key].append(None) # Extract first 4 values (x_min, y_min, x_max, y_max), drop confidence
bbox_values = box[0, :4].cpu().tolist()
return (pose_results,) formatted_boxes.append(bbox_values)
return (pose_results, formatted_boxes)
class DrawNLFPoses: class DrawNLFPoses:
@classmethod @classmethod
@@ -208,25 +298,32 @@ class DrawNLFPoses:
"width": ("INT", {"default": 512}), "width": ("INT", {"default": 512}),
"height": ("INT", {"default": 512}), "height": ("INT", {"default": 512}),
}, },
} "optional": {
"stick_width": ("FLOAT", {"default": 4.0, "min": 0.0, "max": 1000.0, "step": 0.01, "tooltip": "Stick width multiplier"}),
"point_radius": ("INT", {"default": 5, "min": 1, "max": 10, "step": 1, "tooltip": "Point radius for drawing the pose"}),
"style": (["original", "scail"], {"default": "original", "tooltip": "style of the pose drawing"}),
}
}
RETURN_TYPES = ("IMAGE", ) RETURN_TYPES = ("IMAGE", )
RETURN_NAMES = ("image",) RETURN_NAMES = ("image",)
FUNCTION = "predict" FUNCTION = "predict"
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
def predict(self, poses, width, height): def predict(self, poses, width, height, stick_width=1.0, point_radius=2, style="original"):
from .draw_pose import get_control_conditions from .draw_pose import get_control_conditions
print(type(poses))
if isinstance(poses, dict): if isinstance(poses, dict):
pose_input = poses['joints3d_nonparam'][0] if 'joints3d_nonparam' in poses else poses pose_input = poses['joints3d_nonparam'][0] if 'joints3d_nonparam' in poses else poses
else: else:
pose_input = poses pose_input = poses
control_conditions = get_control_conditions(pose_input, height, width)
control_conditions = get_control_conditions(pose_input, height, width, stick_width=stick_width, point_radius=point_radius, style=style)
return (control_conditions,) return (control_conditions,)
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"LoadNLFModel": LoadNLFModel,
"DownloadAndLoadNLFModel": DownloadAndLoadNLFModel, "DownloadAndLoadNLFModel": DownloadAndLoadNLFModel,
"NLFPredict": NLFPredict, "NLFPredict": NLFPredict,
"DrawNLFPoses": DrawNLFPoses, "DrawNLFPoses": DrawNLFPoses,
@@ -234,6 +331,7 @@ NODE_CLASS_MAPPINGS = {
"MTVCrafterEncodePoses": MTVCrafterEncodePoses "MTVCrafterEncodePoses": MTVCrafterEncodePoses
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"LoadNLFModel": "Load NLF Model",
"DownloadAndLoadNLFModel": "(Download)Load NLF Model", "DownloadAndLoadNLFModel": "(Download)Load NLF Model",
"NLFPredict": "NLF Predict", "NLFPredict": "NLF Predict",
"DrawNLFPoses": "Draw NLF Poses", "DrawNLFPoses": "Draw NLF Poses",
+13 -6
View File
@@ -75,14 +75,21 @@ class VAE(nn.Module):
super().__init__() super().__init__()
if data_dim == 80: if data_dim == 80:
self.data_mean = nn.Buffer(torch.tensor(DATA_MEAN_80D, dtype=torch.float32)) data_mean = torch.tensor(DATA_MEAN_80D, dtype=torch.float32)
self.data_std = nn.Buffer(torch.tensor(DATA_STD_80D, dtype=torch.float32)) data_std = torch.tensor(DATA_STD_80D, dtype=torch.float32)
elif data_dim == 128: elif data_dim == 128:
self.data_mean = nn.Buffer(torch.tensor(DATA_MEAN_128D, dtype=torch.float32)) data_mean = torch.tensor(DATA_MEAN_128D, dtype=torch.float32)
self.data_std = nn.Buffer(torch.tensor(DATA_STD_128D, dtype=torch.float32)) data_std = torch.tensor(DATA_STD_128D, dtype=torch.float32)
else:
raise ValueError(f"Unsupported data_dim={data_dim}, expected 80 or 128")
self.data_mean = self.data_mean.view(1, -1, 1) # match old shape: (1, channels, 1)
self.data_std = self.data_std.view(1, -1, 1) data_mean = data_mean.view(1, -1, 1)
data_std = data_std.view(1, -1, 1)
# register as buffers so they move with .to(device) / .cuda()
self.register_buffer("data_mean", data_mean)
self.register_buffer("data_std", data_std)
self.encoder = Encoder1D( self.encoder = Encoder1D(
dim=hidden_dim, dim=hidden_dim,
+96
View File
@@ -0,0 +1,96 @@
import torch
from ..utils import log
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoAddSCAILReferenceEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"vae": ("WANVAE", {"tooltip": "VAE model"}),
"ref_image": ("IMAGE",),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
},
"optional": {
"clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, vae, ref_image, strength, start_percent, end_percent, clip_embeds=None):
updated = dict(embeds)
vae.to(device)
ref_image_in = (ref_image[..., :3].permute(3, 0, 1, 2) * 2 - 1).to(device, vae.dtype)
ref_latent = vae.encode([ref_image_in], device, tiled=False)[0]
log.info(f"SCAIL ref_latent shape: {ref_latent.shape}")
ref_mask = torch.ones_like(ref_latent[:4])
ref_latent = torch.cat([ref_latent, ref_mask], dim=0)
vae.to(offload_device)
updated.setdefault("scail_embeds", {})
updated["scail_embeds"]["ref_latent_pos"] = ref_latent * strength
updated["scail_embeds"]["ref_latent_neg"] = torch.zeros_like(ref_latent)
updated["scail_embeds"]["ref_start_percent"] = start_percent
updated["scail_embeds"]["ref_end_percent"] = end_percent
updated["clip_context"] = clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None
return (updated,)
class WanVideoAddSCAILPoseEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"vae": ("WANVAE", {"tooltip": "VAE model"}),
"pose_images": ("IMAGE", {"tooltip": "Pose images for the entire video"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the pose control"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the pose control application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the pose control application"}),
},
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, vae, pose_images, strength, start_percent=0.0, end_percent=1.0):
updated = dict(embeds)
vae.to(device)
pose_images_in = (pose_images[..., :3].permute(3, 0, 1, 2) * 2 - 1).to(device, vae.dtype)
pose_latent = vae.encode([pose_images_in], device, tiled=False)[0]
pose_mask = torch.ones_like(pose_latent[:4])
pose_latent = torch.cat([pose_latent, pose_mask], dim=0)
log.info(f"SCAIL pose_latent shape: {pose_latent.shape}")
vae.to(offload_device)
updated.setdefault("scail_embeds", {})
updated["scail_embeds"]["pose_latent"] = pose_latent
updated["scail_embeds"]["pose_strength"] = strength
updated["scail_embeds"]["pose_start_percent"] = start_percent
updated["scail_embeds"]["pose_end_percent"] = end_percent
return (updated,)
NODE_CLASS_MAPPINGS = {
"WanVideoAddSCAILPoseEmbeds": WanVideoAddSCAILPoseEmbeds,
"WanVideoAddSCAILReferenceEmbeds": WanVideoAddSCAILReferenceEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddSCAILReferenceEmbeds": "WanVideo Add SCAIL Reference Embeds",
"WanVideoAddSCAILPoseEmbeds": "WanVideo Add SCAIL Pose Embeds",
}
Binary file not shown.
Binary file not shown.
+207
View File
@@ -0,0 +1,207 @@
import json
import torch
import torchvision.transforms.functional as TF
from ..utils import log
from .trajectory import create_pos_feature_map, draw_tracks_on_video, replace_feature
import os
from comfy import model_management as mm
device = mm.get_torch_device()
script_directory = os.path.dirname(os.path.abspath(__file__))
VAE_STRIDE = (4, 8, 8) # t, h, w
class WanVideoWanDrawWanMoveTracks:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"images": ("IMAGE",),
"tracks": ("TRACKS",),
},
"optional": {
"line_resolution": ("INT", {"default": 24, "min": 4, "max": 64, "step": 1, "tooltip": "Number of points to use for each line segment"}),
"circle_size": ("INT", {"default": 10, "min": 1, "max": 20, "step": 1, "tooltip": "Size of the circle to draw for each track point"}),
"opacity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Opacity of the circle to draw for each track point"}),
"line_width": ("INT", {"default": 14, "min": 1, "max": 50, "step": 1, "tooltip": "Width of the line to draw for each track"}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "execute"
CATEGORY = "WanVideoWrapper"
def execute(self, images, tracks, line_resolution=24, circle_size=10, opacity=0.5, line_width=14):
if tracks is None or "track_path" not in tracks:
log.warning("WanVideoWanDrawWanMoveTracks: No tracks provided.")
return (images.float().cpu(), )
track = tracks["track_path"].unsqueeze(0)
track_visibility = tracks["track_visibility"].unsqueeze(0)
images_in = images * 255.0
if images_in.shape[0] != track.shape[1]:
repeat_count = track.shape[1] // images.shape[0]
images_in = images_in.repeat(repeat_count, 1, 1, 1)
track_video = draw_tracks_on_video(images_in, track, track_visibility, track_frame=line_resolution, circle_size=circle_size, opacity=opacity, line_width=line_width)
track_video = torch.stack([TF.to_tensor(frame) for frame in track_video], dim=0).movedim(1, -1)
return (track_video.float().cpu(), )
class WanVideoAddWanMoveTracks:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"image_embeds": ("WANVIDIMAGE_EMBEDS",),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
},
"optional": {
"track_mask": ("MASK",),
"track_coords": ("STRING", {"forceInput": True, "tooltip": "JSON string or list of JSON strings representing the tracks"}),
"tracks": ("TRACKS", {"tooltip": "Alternatively use Comfy Tracks dictionary"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "TRACKS")
RETURN_NAMES = ("image_embeds", "tracks")
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, image_embeds, track_coords=None, tracks=None, strength=1.0, track_mask=None):
updated = dict(image_embeds)
track_visibility = None
target_shape = image_embeds.get("target_shape")
if target_shape is not None:
height = target_shape[2] * VAE_STRIDE[1]
width = target_shape[3] * VAE_STRIDE[2]
else:
height = image_embeds["lat_h"] * VAE_STRIDE[1]
width = image_embeds["lat_w"] * VAE_STRIDE[2]
num_frames = image_embeds["num_frames"]
if track_coords is not None:
tracks_data = parse_json_tracks(track_coords)
track_list = [
[[track[frame]['x'], track[frame]['y']] for track in tracks_data]
for frame in range(len(tracks_data[0]))
]
track = torch.tensor(track_list, dtype=torch.float32, device=device) # shape: (frames, num_tracks, 2)
elif tracks is not None and "track_path" in tracks:
track = tracks["track_path"]
if track_mask is None:
track_visibility = tracks.get("track_visibility", None)
track = track[:num_frames]
num_tracks = track.shape[-2]
if track_visibility is None:
if track_mask is None:
track_visibility = torch.ones((num_frames, num_tracks), dtype=torch.bool, device=device)
else:
track_visibility = (track_mask > 0).any(dim=(1, 2)).unsqueeze(-1)
feature_map, track_pos = create_pos_feature_map(track, track_visibility, VAE_STRIDE, height, width, 16, track_num=num_tracks, device=device)
updated.setdefault("wanmove_embeds", {})
updated["wanmove_embeds"]["track_pos"] = track_pos
updated["wanmove_embeds"]["strength"] = strength
tracks_dict = {
"track_path": track,
"track_visibility": track_visibility,
}
return (updated, tracks_dict,)
def parse_json_tracks(tracks):
tracks_data = []
try:
# If tracks is a string, try to parse it as JSON
if isinstance(tracks, str):
parsed = json.loads(tracks.replace("'", '"'))
tracks_data.extend(parsed)
else:
# If tracks is a list of strings, parse each one
for track_str in tracks:
parsed = json.loads(track_str.replace("'", '"'))
tracks_data.append(parsed)
# Check if we have a single track (dict with x,y) or a list of tracks
if tracks_data and isinstance(tracks_data[0], dict) and 'x' in tracks_data[0]:
# Single track detected, wrap it in a list
tracks_data = [tracks_data]
elif tracks_data and isinstance(tracks_data[0], list) and tracks_data[0] and isinstance(tracks_data[0][0], dict) and 'x' in tracks_data[0][0]:
# Already a list of tracks, nothing to do
pass
else:
# Unexpected format
log.warning(f"Warning: Unexpected track format: {type(tracks_data[0])}")
except json.JSONDecodeError as e:
log.warning(f"Error parsing tracks JSON: {e}")
tracks_data = []
return tracks_data
import node_helpers
class WanMove_native:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"positive": ("CONDITIONING",),
"track_coords": ("STRING", {"forceInput": True, "tooltip": "JSON string or list of JSON strings representing the tracks"}),
},
"optional": {
"track_mask": ("MASK",),
}
}
RETURN_TYPES = ("CONDITIONING", "TRACKS")
RETURN_NAMES = ("positive", "tracks")
FUNCTION = "patchcond"
CATEGORY = "WanVideoWrapper"
DEPRECATED = True
def patchcond(self, positive, track_coords, track_mask=None):
concat_latent_image = positive[0][1]["concat_latent_image"]
B, C, T, H, W = concat_latent_image.shape
num_frames = (T-1) * 4 + 1
width = W * 8
height = H * 8
tracks_data = parse_json_tracks(track_coords)
track_list = [
[[track[frame]['x'], track[frame]['y']] for track in tracks_data]
for frame in range(len(tracks_data[0]))
]
track = torch.tensor(track_list, dtype=torch.float32, device=device) # shape: (frames, num_tracks, 2)
track = track[:num_frames]
num_tracks = track.shape[-2]
if track_mask is None:
track_visibility = torch.ones((num_frames, num_tracks), dtype=torch.bool, device=device)
else:
track_visibility = (track_mask > 0).any(dim=(1, 2)).unsqueeze(-1)
feature_map, track_pos = create_pos_feature_map(track, track_visibility, VAE_STRIDE, height, width, 16, track_num=num_tracks, device=device)
wanmove_cond = replace_feature(concat_latent_image, track_pos.unsqueeze(0))
positive = node_helpers.conditioning_set_values(positive, {"concat_latent_image": wanmove_cond})
tracks_dict = {
"track_path": track,
"track_visibility": track_visibility,
}
return (positive, tracks_dict)
NODE_CLASS_MAPPINGS = {
"WanVideoAddWanMoveTracks": WanVideoAddWanMoveTracks,
"WanVideoWanDrawWanMoveTracks": WanVideoWanDrawWanMoveTracks,
"WanMove_native": WanMove_native,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddWanMoveTracks": "WanVideo Add WanMove Tracks",
"WanVideoWanDrawWanMoveTracks": "WanVideo Draw WanMove Tracks",
"WanMove_native": "WanMove Native",
}
+340
View File
@@ -0,0 +1,340 @@
# https://github.com/ali-vilab/Wan-Move/blob/main/wan/modules/trajectory.py
import numpy as np
import torch
from PIL import Image, ImageDraw
SKIP_ZERO = False
def get_pos_emb(
pos_k: torch.Tensor,
pos_emb_dim: int,
theta_func: callable = lambda i, d: torch.pow(10000, torch.mul(2, torch.div(i.to(torch.float32), d))),
device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu"),
dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
"""
Generate batch position embeddings.
Args:
pos_k (torch.Tensor): A 1D tensor containing positions for which to generate embeddings.
pos_emb_dim (int): The dimension of position embeddings.
theta_func (callable): Function to compute thetas based on position and embedding dimensions.
device (torch.device): Device to store the position embeddings.
dtype (torch.dtype): Desired data type for computations.
Returns:
torch.Tensor: The position embeddings with shape (batch_size, pos_emb_dim).
"""
assert pos_emb_dim % 2 == 0, "The dimension of position embeddings must be even."
pos_k = pos_k.to(device, dtype)
if SKIP_ZERO:
pos_k = pos_k + 1
batch_size = pos_k.size(0)
denominator = torch.arange(0, pos_emb_dim // 2, device=device, dtype=dtype)
# Expand denominator to match the shape needed for broadcasting
denominator_expanded = denominator.view(1, -1).expand(batch_size, -1)
thetas = theta_func(denominator_expanded, pos_emb_dim)
# Ensure pos_k is in the correct shape for broadcasting
pos_k_expanded = pos_k.view(-1, 1).to(dtype)
sin_thetas = torch.sin(torch.div(pos_k_expanded, thetas))
cos_thetas = torch.cos(torch.div(pos_k_expanded, thetas))
# Concatenate sine and cosine embeddings along the last dimension
pos_emb = torch.cat([sin_thetas, cos_thetas], dim=-1)
return pos_emb
def create_pos_feature_map(
pred_tracks: torch.Tensor, # [T, N, 2]
pred_visibility: torch.Tensor, # [T, N]
downsample_ratios: list[int],
height: int,
width: int,
pos_emb_dim: int,
track_num: int = -1,
t_down_strategy: str = "sample",
device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu"),
dtype: torch.dtype = torch.float32,
):
"""
Create a feature map from the predicted tracks.
Args:
- pred_tracks: torch.Tensor, the predicted tracks, [T, N, 2]
- pred_visibility: torch.Tensor, the predicted visibility, [T, N]
- downsample_ratios: list[int], the ratios for downsampling time, height, and width
- height: int, the height of the feature map
- width: int, the width of the feature map
- pos_emb_dim: int, the dimension of the position embeddings
- track_num: int, the number of tracks to use
- t_down_strategy: str, the strategy for downsampling time dimension
- device: torch.device, the device
- dtype: torch.dtype, the data type
Returns:
- feature_map: torch.Tensor, the feature map, [T', H', W', pos_emb_dim]
- track_pos: torch.Tensor, the position embeddings, [N, T', 2], 2 = height, width
"""
assert t_down_strategy in ["sample", "average"], "Invalid strategy for downsampling time dimension."
t, n, _ = pred_tracks.shape
t_down, h_down, w_down = downsample_ratios
feature_map = torch.zeros((t-1) // t_down + 1, height // h_down, width // w_down, pos_emb_dim, device=device, dtype=dtype)
track_pos = - torch.ones(n, (t-1) // t_down + 1, 2, dtype=torch.long)
if track_num == -1:
track_num = n
tracks_idx = torch.randperm(n)[:track_num]
tracks = pred_tracks[:, tracks_idx]
visibility = pred_visibility[:, tracks_idx]
#tracks_embs = get_pos_emb(torch.randperm(n)[:track_num], pos_emb_dim, device=device, dtype=dtype)
for t_idx in range(0, t, t_down):
if t_down_strategy == "sample" or t_idx == 0:
cur_tracks = tracks[t_idx] # [N, 2]
cur_visibility = visibility[t_idx] # [N]
else:
cur_tracks = tracks[t_idx:t_idx+t_down].mean(dim=0)
cur_visibility = torch.any(visibility[t_idx:t_idx+t_down], dim=0)
for i in range(track_num):
if not cur_visibility[i] or cur_tracks[i][0] < 0 or cur_tracks[i][1] < 0 or cur_tracks[i][0] >= width or cur_tracks[i][1] >= height:
continue
x, y = cur_tracks[i]
x, y = int(x // w_down), int(y // h_down)
#feature_map[t_idx // t_down, y, x] += tracks_embs[i]
track_pos[i, t_idx // t_down, 0], track_pos[i, t_idx // t_down, 1] = y, x
return feature_map, track_pos
def replace_feature(
vae_feature: torch.Tensor, # [B, C', T', H', W']
track_pos: torch.Tensor, # [B, N, T', 2]
strength: float = 1.0,
) -> torch.Tensor:
b, _, t, h, w = vae_feature.shape
assert b == track_pos.shape[0], "Batch size mismatch."
n = track_pos.shape[1]
# Shuffle the trajectory order
track_pos = track_pos[:, torch.randperm(n), :, :]
# Extract coordinates at time steps ≥ 1 and generate a valid mask
current_pos = track_pos[:, :, 1:, :] # [B, N, T-1, 2]
mask = (current_pos[..., 0] >= 0) & (current_pos[..., 1] >= 0) # [B, N, T-1]
# Get all valid indices
valid_indices = mask.nonzero(as_tuple=False) # [num_valid, 3]
num_valid = valid_indices.shape[0]
if num_valid == 0:
return vae_feature
# Decompose valid indices into each dimension
batch_idx = valid_indices[:, 0]
track_idx = valid_indices[:, 1]
t_rel = valid_indices[:, 2]
t_target = t_rel + 1 # Convert to original time step indices
# Extract target position coordinates
h_target = current_pos[batch_idx, track_idx, t_rel, 0].long() # Ensure integer indices
w_target = current_pos[batch_idx, track_idx, t_rel, 1].long()
# Extract source position coordinates (t=0)
h_source = track_pos[batch_idx, track_idx, 0, 0].long()
w_source = track_pos[batch_idx, track_idx, 0, 1].long()
# Get source features and assign to target positions
src_features = vae_feature[batch_idx, :, 0, h_source, w_source]
dst_features = vae_feature[batch_idx, :, t_target, h_target, w_target]
vae_feature[batch_idx, :, t_target, h_target, w_target] = dst_features + (src_features - dst_features) * strength
return vae_feature
def get_video_track_video(
model,
video_tensor: torch.Tensor, # [T, C, H, W]
downsample_ratios: list[int],
pos_emb_dim: int,
grid_size: int = 32,
track_num: int = -1,
t_down_strategy: str = "sample",
device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu"),
dtype: torch.dtype = torch.float32,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Get the track video from the video tensor.
Args:
- model: torch.nn.Module, the model for tracking, CoTracker
- video_tensor: torch.Tensor, the video tensor, [T, C, H, W]
- downsample_ratios: list[int], the ratios for downsampling time, height, and width
- height: int, the height of the feature map
- width: int, the width of the feature map
- pos_emb_dim: int, the dimension of the position embeddings
- grid_size: int, the size of the grid
- track_num: int, the number of tracks to use
- t_down_strategy: str, the strategy for downsampling time dimension
- device: torch.device, the device
- dtype: torch.dtype, the data type
Returns:
- track_video: torch.Tensor, the track video, [pos_emb_dim, T', H', W']
- track_pos: torch.Tensor, the position embeddings, [N, T', 2], 2 = height, width
- pred_tracks: the predicted point trajectories
- pred_visibility: visibility of the predicted point trajectories
"""
t, c, height, width = video_tensor.shape
with (
torch.autocast(device_type=device.type, dtype=dtype),
torch.no_grad(),
):
pred_tracks, pred_visibility = model(
video_tensor.unsqueeze(0),
grid_size=grid_size,
backward_tracking=False,
)
track_video, track_pos = create_pos_feature_map(
pred_tracks[0], pred_visibility[0], downsample_ratios, height, width, pos_emb_dim, track_num, t_down_strategy, device, dtype
)
return track_video.permute(3, 0, 1, 2), track_pos, pred_tracks, pred_visibility
# ---------------------------
# Visualize functions
# --------------------------
def add_weighted(rgb, track):
rgb = np.array(rgb) # [H, W, C] "RGB"
track = np.array(track) # [H, W, C] "RGBA"
# Compute weights from the alpha channel
alpha = track[:, :, 3] / 255.0
# Expand alpha to 3 channels to match RGB
alpha = np.stack([alpha] * 3, axis=-1)
# Blend the two images
blend_img = track[:, :, :3] * alpha + rgb * (1 - alpha)
return Image.fromarray(blend_img.astype(np.uint8))
def draw_tracks_on_video(video, tracks, visibility=None, track_frame=24, circle_size=12, opacity=0.5, line_width=16):
color_map = [(102, 153, 255), (0, 255, 255), (255, 255, 0), (255, 102, 204), (0, 255, 0)]
video = video.byte().cpu().numpy() # (81, 480, 832, 3)
tracks = tracks[0].long().detach().cpu().numpy()
if visibility is not None:
visibility = visibility[0].detach().cpu().numpy()
num_frames, height, width = video.shape[:3]
num_tracks = tracks.shape[1]
alpha_opacity = int(255 * opacity)
output_frames = []
for t in range(num_frames):
frame_rgb = video[t].astype(np.float32)
# Create a single RGBA overlay for all tracks in this frame
overlay = Image.new("RGBA", (width, height), (0, 0, 0, 0))
draw_overlay = ImageDraw.Draw(overlay)
polyline_data = []
# Draw all circles on a single overlay
for n in range(num_tracks):
if visibility is not None and visibility[t, n] == 0:
continue
track_coord = tracks[t, n]
color = color_map[n % len(color_map)]
circle_color = color + (alpha_opacity,)
draw_overlay.ellipse(
(
track_coord[0] - circle_size,
track_coord[1] - circle_size,
track_coord[0] + circle_size,
track_coord[1] + circle_size
),
fill=circle_color
)
# Store polyline data for batch processing
tracks_coord = tracks[max(t - track_frame, 0):t + 1, n]
if len(tracks_coord) > 1:
polyline_data.append((tracks_coord, color))
# Blend circles overlay once
overlay_np = np.array(overlay)
alpha = overlay_np[:, :, 3:4] / 255.0
frame_rgb = overlay_np[:, :, :3] * alpha + frame_rgb * (1 - alpha)
# Draw all polylines on a single overlay
if polyline_data:
polyline_overlay = Image.new("RGBA", (width, height), (0, 0, 0, 0))
for tracks_coord, color in polyline_data:
_draw_gradient_polyline_on_overlay(polyline_overlay, line_width, tracks_coord, color, opacity)
# Blend polylines overlay once
polyline_np = np.array(polyline_overlay)
alpha = polyline_np[:, :, 3:4] / 255.0
frame_rgb = polyline_np[:, :, :3] * alpha + frame_rgb * (1 - alpha)
output_frames.append(Image.fromarray(frame_rgb.astype(np.uint8)))
return output_frames
def _draw_gradient_polyline_on_overlay(overlay, line_width, points, start_color, opacity=1.0):
"""
Draw a gradient polyline directly onto an existing RGBA overlay image.
This is an optimized version that doesn't create new images.
"""
draw = ImageDraw.Draw(overlay, 'RGBA')
points = points[::-1]
# Compute total length
total_length = 0
segment_lengths = []
for i in range(len(points) - 1):
dx = points[i + 1][0] - points[i][0]
dy = points[i + 1][1] - points[i][1]
length = (dx * dx + dy * dy) ** 0.5
segment_lengths.append(length)
total_length += length
if total_length == 0:
return
accumulated_length = 0
# Draw the gradient polyline
for idx, (start_point, end_point) in enumerate(zip(points[:-1], points[1:])):
segment_length = segment_lengths[idx]
steps = max(int(segment_length), 1)
for i in range(steps):
current_length = accumulated_length + (i / steps) * segment_length
ratio = current_length / total_length
alpha = int(255 * (1 - ratio) * opacity)
color = (*start_color, alpha)
x = int(start_point[0] + (end_point[0] - start_point[0]) * i / steps)
y = int(start_point[1] + (end_point[1] - start_point[1]) * i / steps)
dynamic_line_width = max(int(line_width * (1 - ratio)), 1)
draw.line([(x, y), (x + 1, y)], fill=color, width=dynamic_line_width)
accumulated_length += segment_length
+61 -114
View File
@@ -1,128 +1,75 @@
try: try:
from .utils import check_duplicate_nodes, log from .utils import check_duplicate_nodes, log, color_text
duplicate_dirs = check_duplicate_nodes() duplicate_dirs = check_duplicate_nodes()
if duplicate_dirs: if duplicate_dirs:
warning_msg = f"WARNING: Found {len(duplicate_dirs)} other WanVideoWrapper directories:\n" warning_msg = f"WARNING: Found {len(duplicate_dirs)} other WanVideoWrapper directories:\n"
for dir_path in duplicate_dirs: for dir_path in duplicate_dirs:
warning_msg += f" - {dir_path}\n" warning_msg += f" - {color_text(dir_path, 'yellow')}\n"
log.warning(warning_msg + "Please remove duplicates to avoid possible conflicts.") log.warning(color_text(warning_msg + "Please remove duplicates to avoid possible conflicts.", "red"))
except: except Exception:
pass pass
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS from .utils import log
from .recammaster.nodes import NODE_CLASS_MAPPINGS as RECAM_MASTER_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS
from .skyreels.nodes import NODE_CLASS_MAPPINGS as SKYREELS_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as SKYREELS_NODE_DISPLAY_NAME_MAPPINGS
from .fantasytalking.nodes import NODE_CLASS_MAPPINGS as FANTASYTALKING_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FANTASYTALKING_NODE_DISPLAY_NAME_MAPPINGS
from .nodes_sampler import NODE_CLASS_MAPPINGS as SAMPLER_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as SAMPLER_NODE_DISPLAY_NAME_MAPPINGS
from .fun_camera.nodes import NODE_CLASS_MAPPINGS as FUN_CAMERA_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FUN_CAMERA_NODE_DISPLAY_NAME_MAPPINGS
from .uni3c.nodes import NODE_CLASS_MAPPINGS as UNI3C_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNI3C_NODE_DISPLAY_NAME_MAPPINGS
from .controlnet.nodes import NODE_CLASS_MAPPINGS as CONTROLNET_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as CONTROLNET_NODE_DISPLAY_NAME_MAPPINGS
from .ATI.nodes import NODE_CLASS_MAPPINGS as ATI_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as ATI_NODE_DISPLAY_NAME_MAPPINGS
from .multitalk.nodes import NODE_CLASS_MAPPINGS as MULTITALK_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MULTITALK_NODE_DISPLAY_NAME_MAPPINGS
from .nodes_model_loading import NODE_CLASS_MAPPINGS as MODEL_LOADING_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MODEL_LOADING_NODE_DISPLAY_NAME_MAPPINGS
from .nodes_utility import NODE_CLASS_MAPPINGS as UTILITY_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UTILITY_NODE_DISPLAY_NAME_MAPPINGS
from .cache_methods.nodes_cache import NODE_CLASS_MAPPINGS as NODE_CACHE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as NODE_CACHE_DISPLAY_NAME_MAPPINGS
from .nodes_deprecated import NODE_CLASS_MAPPINGS as DEPRECATED_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS
from .s2v.nodes import NODE_CLASS_MAPPINGS as S2V_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as S2V_NODE_DISPLAY_NAME_MAPPINGS
from .FlashVSR.flashvsr_nodes import NODE_CLASS_MAPPINGS as FLASHVSR_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FLASHVSR_NODE_DISPLAY_NAME_MAPPINGS
from .mocha.nodes import NODE_CLASS_MAPPINGS as MOCHA_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MOCHA_NODE_DISPLAY_NAME_MAPPINGS
try: NODE_CLASS_MAPPINGS = {}
from .qwen.qwen import NODE_CLASS_MAPPINGS as QWEN_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as QWEN_NODE_DISPLAY_NAME_MAPPINGS NODE_DISPLAY_NAME_MAPPINGS = {}
except Exception as e:
log.warning(f"WanVideoWrapper WARNING: Qwen nodes not available due to error in importing them: {e}")
QWEN_NODE_CLASS_MAPPINGS = {}
QWEN_NODE_DISPLAY_NAME_MAPPINGS = {}
# Required modules (will raise on import failure)
REQUIRED_MODULES = [
(".nodes", "Main"),
(".nodes_sampler", "Sampler"),
(".nodes_model_loading", "ModelLoading"),
(".nodes_utility", "Utility"),
(".cache_methods.nodes_cache", "Cache"),
]
try: # Optional modules (will warn on import failure)
from .fantasyportrait.nodes import NODE_CLASS_MAPPINGS as FANTASYPORTRAIT_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FANTASYPORTRAIT_NODE_DISPLAY_NAME_MAPPINGS OPTIONAL_MODULES = [
except Exception as e: (".nodes_deprecated", "Deprecated"),
log.warning(f"WanVideoWrapper WARNING: FantasyPortrait nodes not available due to error in importing them: {e}") (".s2v.nodes", "S2V"),
FANTASYPORTRAIT_NODE_CLASS_MAPPINGS = {} (".FlashVSR.flashvsr_nodes", "FlashVSR"),
FANTASYPORTRAIT_NODE_DISPLAY_NAME_MAPPINGS = {} (".mocha.nodes", "Mocha"),
(".fun_camera.nodes", "FunCamera"),
(".uni3c.nodes", "Uni3C"),
(".controlnet.nodes", "ControlNet"),
(".ATI.nodes", "ATI"),
(".multitalk.nodes", "MultiTalk"),
(".recammaster.nodes", "RecamMaster"),
(".skyreels.nodes", "SkyReels"),
(".fantasytalking.nodes", "FantasyTalking"),
(".qwen.qwen", "Qwen"),
(".fantasyportrait.nodes", "FantasyPortrait"),
(".unianimate.nodes", "UniAnimate"),
(".MTV.nodes", "MTV"),
(".HuMo.nodes", "HuMo"),
(".lynx.nodes", "Lynx"),
(".Ovi.nodes_ovi", "Ovi"),
(".steadydancer.nodes", "SteadyDancer"),
(".onetoall.nodes", "OneToAll"),
(".WanMove.nodes", "WanMove"),
(".SCAIL.nodes", "SCAIL"),
(".LongCat.nodes", "LongCat"),
(".LongVie2.nodes", "LongVie2"),
]
try: def register_nodes(module_path: str, name: str, optional: bool) -> None:
from .unianimate.nodes import NODE_CLASS_MAPPINGS as UNIANIMATE_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS """Import and register nodes from a module."""
except Exception as e: try:
log.warning(f"WanVideoWrapper WARNING: UniAnimate nodes not available due to error in importing them: {e}") import importlib
UNIANIMATE_NODE_CLASS_MAPPINGS = {} module = importlib.import_module(module_path, package=__package__)
UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS = {} NODE_CLASS_MAPPINGS.update(getattr(module, "NODE_CLASS_MAPPINGS", {}))
NODE_DISPLAY_NAME_MAPPINGS.update(getattr(module, "NODE_DISPLAY_NAME_MAPPINGS", {}))
except Exception as e:
if optional:
log.warning(f"WanVideoWrapper WARNING: {name} nodes not available: {e}")
else:
raise
try: # Register all node modules
from .MTV.nodes import NODE_CLASS_MAPPINGS as MTV_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MTV_NODE_DISPLAY_NAME_MAPPINGS for module_path, name in REQUIRED_MODULES:
except Exception as e: register_nodes(module_path, name, optional=False)
log.warning(f"WanVideoWrapper WARNING: MTV nodes not available due to error in importing them: {e}")
MTV_NODE_CLASS_MAPPINGS = {}
MTV_NODE_DISPLAY_NAME_MAPPINGS = {}
try: for module_path, name in OPTIONAL_MODULES:
from .HuMo.nodes import NODE_CLASS_MAPPINGS as HUMO_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as HUMO_NODE_DISPLAY_NAME_MAPPINGS register_nodes(module_path, name, optional=True)
except Exception as e:
log.warning(f"WanVideoWrapper WARNING: HuMo nodes not available due to error in importing them: {e}")
HUMO_NODE_CLASS_MAPPINGS = {}
HUMO_NODE_DISPLAY_NAME_MAPPINGS = {}
try: __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
from .lynx.nodes import NODE_CLASS_MAPPINGS as LYNX_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as LYNX_NODE_DISPLAY_NAME_MAPPINGS
except Exception as e:
log.warning(f"WanVideoWrapper WARNING: Lynx nodes not available due to error in importing them: {e}")
LYNX_NODE_CLASS_MAPPINGS = {}
LYNX_NODE_DISPLAY_NAME_MAPPINGS = {}
try:
from .Ovi.nodes_ovi import NODE_CLASS_MAPPINGS as OVI_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as OVI_NODE_DISPLAY_NAME_MAPPINGS
except Exception as e:
log.warning(f"WanVideoWrapper WARNING: Ovi nodes not available due to error in importing them: {e}")
OVI_NODE_CLASS_MAPPINGS = {}
OVI_NODE_DISPLAY_NAME_MAPPINGS = {}
NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(FANTASYTALKING_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(FANTASYPORTRAIT_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(FUN_CAMERA_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(UNI3C_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(CONTROLNET_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(ATI_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(MULTITALK_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(MODEL_LOADING_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(UTILITY_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(NODE_CACHE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(DEPRECATED_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(QWEN_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(MTV_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(S2V_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(HUMO_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(SAMPLER_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(LYNX_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(OVI_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(FLASHVSR_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(MOCHA_NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(SKYREELS_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(FANTASYTALKING_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(FANTASYPORTRAIT_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(FUN_CAMERA_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(UNI3C_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(CONTROLNET_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(ATI_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(MULTITALK_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(MODEL_LOADING_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(UTILITY_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(NODE_CACHE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(QWEN_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(MTV_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(S2V_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(HUMO_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(SAMPLER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(LYNX_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(OVI_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(FLASHVSR_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(MOCHA_NODE_DISPLAY_NAME_MAPPINGS)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+1 -13
View File
@@ -28,19 +28,7 @@ aggressive values this can happen and the motion suffers. Starting later can hel
When NOT using coefficients, the threshold value should be When NOT using coefficients, the threshold value should be
about 10 times smaller than the value used with coefficients. about 10 times smaller than the value used with coefficients.
Official recommended values https://github.com/ali-vilab/TeaCache/tree/main/TeaCache4Wan2.1: Official recommended values https://github.com/ali-vilab/TeaCache/tree/main/TeaCache4Wan2.1
<pre style='font-family:monospace'>
+-------------------+--------+---------+--------+
| Model | Low | Medium | High |
+-------------------+--------+---------+--------+
| Wan2.1 t2v 1.3B | 0.05 | 0.07 | 0.08 |
| Wan2.1 t2v 14B | 0.14 | 0.15 | 0.20 |
| Wan2.1 i2v 480P | 0.13 | 0.19 | 0.26 |
| Wan2.1 i2v 720P | 0.18 | 0.20 | 0.30 |
+-------------------+--------+---------+--------+
</pre>
""" """
def process(self, rel_l1_thresh, start_step, end_step, cache_device, use_coefficients, mode="e"): def process(self, rel_l1_thresh, start_step, end_step, cache_device, use_coefficients, mode="e"):
+12 -57
View File
@@ -10,18 +10,13 @@ from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscal
from diffusers.models.modeling_outputs import Transformer2DModelOutput from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.modeling_utils import ModelMixin from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.transformers.transformer_wan import ( from diffusers.models.transformers.transformer_wan import (
WanTimeTextImageEmbedding, WanTimeTextImageEmbedding,
WanRotaryPosEmbed, WanRotaryPosEmbed,
WanTransformerBlock WanTransformerBlock
) )
logger = logging.get_logger(__name__) # pylint: disable=invalid-name logger = logging.get_logger(__name__) # pylint: disable=invalid-name
def zero_module(module):
for p in module.parameters():
nn.init.zeros_(p)
return module
class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin):
r""" r"""
@@ -69,7 +64,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
_no_split_modules = ["WanTransformerBlock"] _no_split_modules = ["WanTransformerBlock"]
_keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"] _keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"]
_keys_to_ignore_on_load_unexpected = ["norm_added_q"] _keys_to_ignore_on_load_unexpected = ["norm_added_q"]
@register_to_config @register_to_config
def __init__( def __init__(
self, self,
@@ -100,10 +95,10 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
## Spatial compression with time awareness ## Spatial compression with time awareness
nn.Sequential( nn.Sequential(
nn.Conv3d( nn.Conv3d(
in_channels, in_channels,
input_channels[0], input_channels[0],
kernel_size=(3, downscale_coef + 1, downscale_coef + 1), kernel_size=(3, downscale_coef + 1, downscale_coef + 1),
stride=(1, downscale_coef, downscale_coef), stride=(1, downscale_coef, downscale_coef),
padding=(1, downscale_coef // 2, downscale_coef // 2) padding=(1, downscale_coef // 2, downscale_coef // 2)
), ),
nn.GELU(approximate="tanh"), nn.GELU(approximate="tanh"),
@@ -122,9 +117,9 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
nn.GroupNorm(2, input_channels[2]), nn.GroupNorm(2, input_channels[2]),
) )
]) ])
inner_dim = num_attention_heads * attention_head_dim inner_dim = num_attention_heads * attention_head_dim
# 1. Patch & position embedding # 1. Patch & position embedding
self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len) self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len)
self.patch_embedding = nn.Conv3d(vae_channels + input_channels[2], inner_dim, kernel_size=patch_size, stride=patch_size) self.patch_embedding = nn.Conv3d(vae_channels + input_channels[2], inner_dim, kernel_size=patch_size, stride=patch_size)
@@ -153,11 +148,10 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
for _ in range(len(self.blocks)): for _ in range(len(self.blocks)):
controlnet_block = nn.Linear(inner_dim, out_proj_dim) controlnet_block = nn.Linear(inner_dim, out_proj_dim)
controlnet_block = zero_module(controlnet_block)
self.controlnet_blocks.append(controlnet_block) self.controlnet_blocks.append(controlnet_block)
self.gradient_checkpointing = False self.gradient_checkpointing = False
def forward( def forward(
self, self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
@@ -187,7 +181,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
# 0. Controlnet encoder # 0. Controlnet encoder
for control_encoder_block in self.control_encoder: for control_encoder_block in self.control_encoder:
controlnet_states = control_encoder_block(controlnet_states) controlnet_states = control_encoder_block(controlnet_states)
hidden_states = torch.cat([hidden_states, controlnet_states], dim=1) hidden_states = torch.cat([hidden_states, controlnet_states], dim=1)
## 1. Patch embedding and stack ## 1. Patch embedding and stack
@@ -216,7 +210,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
if encoder_hidden_states_image is not None: if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1) encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
# 4. Transformer blocks # 4. Transformer blocks
controlnet_hidden_states = () controlnet_hidden_states = ()
if torch.is_grad_enabled() and self.gradient_checkpointing: if torch.is_grad_enabled() and self.gradient_checkpointing:
@@ -239,43 +233,4 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
return (controlnet_hidden_states,) return (controlnet_hidden_states,)
return Transformer2DModelOutput(sample=controlnet_hidden_states) return Transformer2DModelOutput(sample=controlnet_hidden_states)
if __name__ == "__main__":
parameters = {
"added_kv_proj_dim": None,
"attention_head_dim": 128,
"cross_attn_norm": True,
"eps": 1e-06,
"ffn_dim": 8960,
"freq_dim": 256,
"image_dim": None,
"in_channels": 3,
"num_attention_heads": 12,
"num_layers": 2,
"patch_size": [1, 2, 2],
"qk_norm": "rms_norm_across_heads",
"rope_max_seq_len": 1024,
"text_dim": 4096,
"downscale_coef": 8,
"out_proj_dim": 12 * 128,
"vae_channels": 16
}
controlnet = WanControlnet(**parameters)
hidden_states = torch.rand(1, 16, 13, 60, 90)
timestep = torch.tensor([1000]).repeat(17550).unsqueeze(0) #torch.randint(low=0, high=1000, size=(1,), dtype=torch.long)
encoder_hidden_states = torch.rand(1, 512, 4096)
controlnet_states = torch.rand(1, 3, 49, 480, 720)
controlnet_hidden_states = controlnet(
hidden_states=hidden_states,
timestep=timestep,
encoder_hidden_states=encoder_hidden_states,
controlnet_states=controlnet_states,
return_dict=False
)
print("Output states count", len(controlnet_hidden_states[0]))
for out_hidden_states in controlnet_hidden_states[0]:
print(out_hidden_states.shape)
+153 -38
View File
@@ -1,14 +1,52 @@
import torch import torch
import torch.nn as nn import torch.nn as nn
from accelerate import init_empty_weights from accelerate import init_empty_weights
from .gguf.gguf_utils import GGUFParameter, dequantize_gguf_tensor
@torch.library.custom_op("wanvideo::apply_lora", mutates_args=())
def apply_lora(weight: torch.Tensor, lora_diff_0: torch.Tensor, lora_diff_1: torch.Tensor, lora_diff_2: float, lora_strength: torch.Tensor) -> torch.Tensor:
patch_diff = torch.mm(
lora_diff_0.flatten(start_dim=1),
lora_diff_1.flatten(start_dim=1)
).reshape(weight.shape)
alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 != 0.0 else 1.0
scale = lora_strength * alpha
return weight + patch_diff * scale
@apply_lora.register_fake
def _(weight, lora_diff_0, lora_diff_1, lora_diff_2, lora_strength):
# Return weight with same metadata
return weight.clone()
@torch.library.custom_op("wanvideo::apply_single_lora", mutates_args=())
def apply_single_lora(weight: torch.Tensor, lora_diff: torch.Tensor, lora_strength: torch.Tensor) -> torch.Tensor:
return weight + lora_diff * lora_strength
@apply_single_lora.register_fake
def _(weight, lora_diff, lora_strength):
# Return weight with same metadata
return weight.clone()
@torch.library.custom_op("wanvideo::linear_forward", mutates_args=())
def linear_forward(input: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor | None) -> torch.Tensor:
return torch.nn.functional.linear(input, weight, bias)
@linear_forward.register_fake
def _(input, weight, bias):
# Calculate output shape: (..., out_features)
out_features = weight.shape[0]
output_shape = list(input.shape[:-1]) + [out_features]
return input.new_empty(output_shape)
#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py #based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py
def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None, compile_args=None): def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None, compile_args=None, modules_to_not_convert=[]):
has_children = list(model.children()) has_children = list(model.children())
if not has_children: if not has_children:
return return
allow_compile = False allow_compile = False
for name, module in model.named_children(): for name, module in model.named_children():
@@ -16,13 +54,22 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
allow_compile = compile_args.get("allow_unmerged_lora_compile", False) allow_compile = compile_args.get("allow_unmerged_lora_compile", False)
module_prefix = prefix + name + "." module_prefix = prefix + name + "."
module_prefix = module_prefix.replace("_orig_mod.", "") module_prefix = module_prefix.replace("_orig_mod.", "")
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args) _replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args, modules_to_not_convert)
if isinstance(module, nn.Linear) and "loras" not in module_prefix: if isinstance(module, nn.Linear) and "loras" not in module_prefix and "dual_controller" not in module_prefix and name not in modules_to_not_convert:
in_features = state_dict[module_prefix + "weight"].shape[1] weight_key = module_prefix + "weight"
out_features = state_dict[module_prefix + "weight"].shape[0] if weight_key not in state_dict:
if scale_weights is not None: continue
in_features = state_dict[weight_key].shape[1]
out_features = state_dict[weight_key].shape[0]
is_gguf = isinstance(state_dict[weight_key], GGUFParameter)
scale_weight = None
if not is_gguf and scale_weights is not None:
scale_key = f"{module_prefix}scale_weight" scale_key = f"{module_prefix}scale_weight"
scale_weight = scale_weights.get(scale_key)
with init_empty_weights(): with init_empty_weights():
model._modules[name] = CustomLinear( model._modules[name] = CustomLinear(
@@ -30,8 +77,9 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
out_features, out_features,
module.bias is not None, module.bias is not None,
compute_dtype=compute_dtype, compute_dtype=compute_dtype,
scale_weight=scale_weights.get(scale_key) if scale_weights else None, scale_weight=scale_weight,
allow_compile=allow_compile allow_compile=allow_compile,
is_gguf=is_gguf
) )
model._modules[name].source_cls = type(module) model._modules[name].source_cls = type(module)
model._modules[name].requires_grad_(False) model._modules[name].requires_grad_(False)
@@ -71,8 +119,8 @@ def set_lora_params(module, patches, module_prefix="", device=torch.device("cpu"
continue continue
lora_strengths = [p[0] for p in patch] lora_strengths = [p[0] for p in patch]
module.set_lora_diffs(lora_diffs, device=device) module.set_lora_diffs(lora_diffs, device=device)
module.lora_strengths = lora_strengths module.set_lora_strengths(lora_strengths, device=device)
module.step = 0 # Initialize step for LoRA scheduling module._step.fill_(0) # Initialize step for LoRA scheduling
class CustomLinear(nn.Linear): class CustomLinear(nn.Linear):
@@ -84,19 +132,56 @@ class CustomLinear(nn.Linear):
compute_dtype=None, compute_dtype=None,
device=None, device=None,
scale_weight=None, scale_weight=None,
allow_compile=False allow_compile=False,
is_gguf=False
) -> None: ) -> None:
super().__init__(in_features, out_features, bias, device) super().__init__(in_features, out_features, bias, device)
self.compute_dtype = compute_dtype self.compute_dtype = compute_dtype
self.lora_diffs = [] self.lora_diffs = []
self.step = 0 self.register_buffer("_step", torch.zeros((), dtype=torch.long))
self.scale_weight = scale_weight self.scale_weight = scale_weight
self.lora_strengths = [] self.lora_strengths = []
self.allow_compile = allow_compile self.allow_compile = allow_compile
self.is_gguf = is_gguf
if not allow_compile: if not allow_compile:
self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora) self._apply_lora_impl = self._apply_lora_custom_op
self._apply_single_lora_impl = self._apply_single_lora_custom_op
self._linear_forward_impl = self._linear_forward_custom_op
else:
self._apply_lora_impl = self._apply_lora_direct
self._apply_single_lora_impl = self._apply_single_lora_direct
self._linear_forward_impl = self._linear_forward_direct
# Direct implementations (no custom ops)
def _apply_lora_direct(self, weight, lora_diff_0, lora_diff_1, lora_diff_2, lora_strength):
patch_diff = torch.mm(
lora_diff_0.flatten(start_dim=1),
lora_diff_1.flatten(start_dim=1)
).reshape(weight.shape) + 0
alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 != 0.0 else 1.0
scale = lora_strength * alpha
return weight + patch_diff * scale
def _apply_single_lora_direct(self, weight, lora_diff, lora_strength):
return weight + lora_diff * lora_strength
def _linear_forward_direct(self, input, weight, bias):
return torch.nn.functional.linear(input, weight, bias)
# Custom op implementations
def _apply_lora_custom_op(self, weight, lora_diff_0, lora_diff_1, lora_diff_2, lora_strength):
return torch.ops.wanvideo.apply_lora(weight, lora_diff_0, lora_diff_1,
float(lora_diff_2) if lora_diff_2 is not None else 0.0, lora_strength
)
def _apply_single_lora_custom_op(self, weight, lora_diff, lora_strength):
return torch.ops.wanvideo.apply_single_lora(weight, lora_diff, lora_strength)
def _linear_forward_custom_op(self, input, weight, bias):
return torch.ops.wanvideo.linear_forward(input, weight, bias)
def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")): def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")):
self.lora_diffs = [] self.lora_diffs = []
for i, diff in enumerate(lora_diffs): for i, diff in enumerate(lora_diffs):
@@ -109,51 +194,81 @@ class CustomLinear(nn.Linear):
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype)) self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
self.lora_diffs.append(f"lora_diff_{i}_0") self.lora_diffs.append(f"lora_diff_{i}_0")
def set_lora_strengths(self, lora_strengths, device=torch.device("cpu")):
self._lora_strength_tensors = []
self._lora_strength_is_scheduled = []
self._step = self._step.to(device)
for i, strength in enumerate(lora_strengths):
if isinstance(strength, list):
tensor = torch.tensor(strength, dtype=self.compute_dtype, device=device)
self.register_buffer(f"_lora_strength_{i}", tensor)
self._lora_strength_is_scheduled.append(True)
else:
tensor = torch.tensor([strength], dtype=self.compute_dtype, device=device)
self.register_buffer(f"_lora_strength_{i}", tensor)
self._lora_strength_is_scheduled.append(False)
def _get_lora_strength(self, idx):
strength_tensor = getattr(self, f"_lora_strength_{idx}")
if self._lora_strength_is_scheduled[idx]:
return strength_tensor.index_select(0, self._step).squeeze(0)
return strength_tensor[0]
def _get_weight_with_lora(self, weight): def _get_weight_with_lora(self, weight):
"""Apply LoRA outside compiled region""" """Apply LoRA using custom ops to avoid graph breaks"""
if not hasattr(self, "lora_diff_0_0"): if not hasattr(self, "lora_diff_0_0"):
return weight return weight
for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths): for idx, lora_diff_names in enumerate(self.lora_diffs):
if isinstance(lora_strength, list): lora_strength = self._get_lora_strength(idx)
lora_strength = lora_strength[self.step]
if lora_strength == 0.0:
continue
elif lora_strength == 0.0:
continue
if isinstance(lora_diff_names, tuple): if isinstance(lora_diff_names, tuple):
lora_diff_0 = getattr(self, lora_diff_names[0]) lora_diff_0 = getattr(self, lora_diff_names[0])
lora_diff_1 = getattr(self, lora_diff_names[1]) lora_diff_1 = getattr(self, lora_diff_names[1])
lora_diff_2 = getattr(self, lora_diff_names[2]) lora_diff_2 = getattr(self, lora_diff_names[2])
patch_diff = torch.mm(
lora_diff_0.flatten(start_dim=1), weight = self._apply_lora_impl(
lora_diff_1.flatten(start_dim=1) weight, lora_diff_0, lora_diff_1,
).reshape(weight.shape) + 0 float(lora_diff_2) if lora_diff_2 is not None else 0.0, lora_strength
alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 is not None else 1.0 )
scale = lora_strength * alpha
weight = weight.add(patch_diff, alpha=scale)
else: else:
lora_diff = getattr(self, lora_diff_names) lora_diff = getattr(self, lora_diff_names)
weight = weight.add(lora_diff, alpha=lora_strength) weight = self._apply_single_lora_impl(weight, lora_diff, lora_strength)
return weight
def _prepare_weight(self, input):
"""Prepare weight tensor - handles both regular and GGUF weights"""
if self.is_gguf:
weight = dequantize_gguf_tensor(self.weight).to(self.compute_dtype)
else:
weight = self.weight.to(input)
return weight return weight
def forward(self, input): def forward(self, input):
weight = self._prepare_weight(input)
if self.bias is not None: if self.bias is not None:
bias = self.bias.to(input) bias = self.bias.to(input if not self.is_gguf else self.compute_dtype)
else: else:
bias = None bias = None
weight = self.weight.to(input)
if self.scale_weight is not None: # Only apply scale_weight for non-GGUF models
if not self.is_gguf and self.scale_weight is not None:
if weight.numel() < input.numel(): if weight.numel() < input.numel():
weight = weight * self.scale_weight weight = weight * self.scale_weight
else: else:
input = input * self.scale_weight input = input * self.scale_weight
weight = self._get_weight_with_lora(weight) weight = self._get_weight_with_lora(weight)
out = self._linear_forward_impl(input, weight, bias)
del weight, input, bias
return out
def update_lora_step(module, step):
for name, submodule in module.named_modules():
if isinstance(submodule, CustomLinear) and hasattr(submodule, "_step"):
submodule._step.fill_(step)
return torch.nn.functional.linear(input, weight, bias)
def remove_lora_from_module(module): def remove_lora_from_module(module):
for name, submodule in module.named_modules(): for name, submodule in module.named_modules():
if hasattr(submodule, "lora_diffs"): if hasattr(submodule, "lora_diffs"):
File diff suppressed because it is too large Load Diff
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -1,780 +0,0 @@
{
"id": "8b7a9a57-2303-4ef5-9fc2-bf41713bd1fc",
"revision": 0,
"last_node_id": 46,
"last_link_id": 58,
"nodes": [
{
"id": 33,
"type": "Note",
"pos": [
227.3764190673828,
-205.28524780273438
],
"size": [
351.70458984375,
88
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {},
"widgets_values": [
"Models:\nhttps://huggingface.co/Kijai/WanVideo_comfy/tree/main"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 11,
"type": "LoadWanVideoT5TextEncoder",
"pos": [
224.15325927734375,
-34.481563568115234
],
"size": [
377.1661376953125,
130
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "wan_t5_model",
"type": "WANTEXTENCODER",
"slot_index": 0,
"links": [
15
]
}
],
"properties": {
"Node name for S&R": "LoadWanVideoT5TextEncoder"
},
"widgets_values": [
"umt5-xxl-enc-bf16.safetensors",
"bf16",
"offload_device",
"disabled"
]
},
{
"id": 28,
"type": "WanVideoDecode",
"pos": [
1692.973876953125,
-404.8614501953125
],
"size": [
315,
174
],
"flags": {},
"order": 12,
"mode": 0,
"inputs": [
{
"name": "vae",
"type": "WANVAE",
"link": 43
},
{
"name": "samples",
"type": "LATENT",
"link": 33
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"slot_index": 0,
"links": [
48
]
}
],
"properties": {
"Node name for S&R": "WanVideoDecode"
},
"widgets_values": [
true,
272,
272,
144,
128
]
},
{
"id": 38,
"type": "WanVideoVAELoader",
"pos": [
1687.4093017578125,
-582.2750854492188
],
"size": [
416.25482177734375,
82
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "vae",
"type": "WANVAE",
"slot_index": 0,
"links": [
43
]
}
],
"properties": {
"Node name for S&R": "WanVideoVAELoader"
},
"widgets_values": [
"wanvideo\\Wan2_1_VAE_bf16.safetensors",
"bf16"
]
},
{
"id": 42,
"type": "GetImageSizeAndCount",
"pos": [
1708.7301025390625,
-140.99705505371094
],
"size": [
277.20001220703125,
86
],
"flags": {},
"order": 13,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 48
}
],
"outputs": [
{
"name": "image",
"type": "IMAGE",
"slot_index": 0,
"links": [
56
]
},
{
"label": "832 width",
"name": "width",
"type": "INT",
"links": null
},
{
"label": "480 height",
"name": "height",
"type": "INT",
"links": null
},
{
"label": "257 count",
"name": "count",
"type": "INT",
"links": null
}
],
"properties": {
"Node name for S&R": "GetImageSizeAndCount"
},
"widgets_values": []
},
{
"id": 16,
"type": "WanVideoTextEncode",
"pos": [
675.8850708007812,
-36.032100677490234
],
"size": [
420.30511474609375,
261.5306701660156
],
"flags": {},
"order": 10,
"mode": 0,
"inputs": [
{
"name": "t5",
"type": "WANTEXTENCODER",
"link": 15
},
{
"name": "model_to_offload",
"shape": 7,
"type": "WANVIDEOMODEL",
"link": null
}
],
"outputs": [
{
"name": "text_embeds",
"type": "WANVIDEOTEXTEMBEDS",
"slot_index": 0,
"links": [
30
]
}
],
"properties": {
"Node name for S&R": "WanVideoTextEncode"
},
"widgets_values": [
"high quality nature video featuring a red panda balancing on a bamboo stem while a bird lands on it's head, on the background there is a waterfall",
"色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
true
]
},
{
"id": 30,
"type": "VHS_VideoCombine",
"pos": [
2127.120849609375,
-511.9014587402344
],
"size": [
873.2135620117188,
840.2385864257812
],
"flags": {},
"order": 14,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 56
},
{
"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": 16,
"loop_count": 0,
"filename_prefix": "WanVideo2_1_T2V",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 19,
"save_metadata": true,
"trim_to_audio": false,
"pingpong": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "WanVideo2_1_T2V_00412.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
"frame_rate": 16,
"workflow": "WanVideo2_1_T2V_00412.png",
"fullpath": "N:\\AI\\ComfyUI\\output\\WanVideo2_1_T2V_00412.mp4"
}
}
}
},
{
"id": 37,
"type": "WanVideoEmptyEmbeds",
"pos": [
1305.26708984375,
-571.7843627929688
],
"size": [
315,
106
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "image_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"links": [
42
]
}
],
"properties": {
"Node name for S&R": "WanVideoEmptyEmbeds"
},
"widgets_values": [
832,
480,
257
]
},
{
"id": 35,
"type": "WanVideoTorchCompileSettings",
"pos": [
193.47103881835938,
-614.6900024414062
],
"size": [
390.5999755859375,
178
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "torch_compile_args",
"type": "WANCOMPILEARGS",
"slot_index": 0,
"links": []
}
],
"properties": {
"Node name for S&R": "WanVideoTorchCompileSettings"
},
"widgets_values": [
"inductor",
false,
"default",
false,
64,
true
]
},
{
"id": 45,
"type": "WanVideoTeaCache",
"pos": [
931.4036865234375,
-792.5159912109375
],
"size": [
315,
154
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "cache_args",
"type": "CACHEARGS",
"links": [
58
]
}
],
"properties": {
"Node name for S&R": "WanVideoTeaCache"
},
"widgets_values": [
0.10000000000000002,
1,
-1,
"offload_device",
true
]
},
{
"id": 36,
"type": "Note",
"pos": [
796.0189208984375,
-521.5020751953125
],
"size": [
298.2554016113281,
108.62744140625
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {},
"widgets_values": [
"sdpa should work too, haven't tested flaash\n\nfp8_fast seems to cause huge quality degradation"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 46,
"type": "Note",
"pos": [
937.9556274414062,
-940.750244140625
],
"size": [
297.4364013671875,
88
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {},
"widgets_values": [
"TeaCache with context windows is VERY experimental and lower values than normal should be used."
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 27,
"type": "WanVideoSampler",
"pos": [
1315.2401123046875,
-401.48028564453125
],
"size": [
315,
574.1923217773438
],
"flags": {},
"order": 11,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "WANVIDEOMODEL",
"link": 29
},
{
"name": "text_embeds",
"type": "WANVIDEOTEXTEMBEDS",
"link": 30
},
{
"name": "image_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"link": 42
},
{
"name": "samples",
"shape": 7,
"type": "LATENT",
"link": null
},
{
"name": "feta_args",
"shape": 7,
"type": "FETAARGS",
"link": null
},
{
"name": "context_options",
"shape": 7,
"type": "WANVIDCONTEXT",
"link": 57
},
{
"name": "cache_args",
"shape": 7,
"type": "CACHEARGS",
"link": 58
},
{
"name": "flowedit_args",
"shape": 7,
"type": "FLOWEDITARGS",
"link": null
},
{
"name": "slg_args",
"shape": 7,
"type": "SLGARGS",
"link": null
},
{
"name": "loop_args",
"shape": 7,
"type": "LOOPARGS",
"link": null
}
],
"outputs": [
{
"name": "samples",
"type": "LATENT",
"slot_index": 0,
"links": [
33
]
}
],
"properties": {
"Node name for S&R": "WanVideoSampler"
},
"widgets_values": [
30,
6,
5,
1057359483639288,
"fixed",
true,
"unipc",
0,
1,
"",
"comfy"
]
},
{
"id": 43,
"type": "WanVideoContextOptions",
"pos": [
1307.9542236328125,
-855.8865356445312
],
"size": [
315,
226
],
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "vae",
"shape": 7,
"type": "WANVAE",
"link": null
}
],
"outputs": [
{
"name": "context_options",
"type": "WANVIDCONTEXT",
"slot_index": 0,
"links": [
57
]
}
],
"properties": {
"Node name for S&R": "WanVideoContextOptions"
},
"widgets_values": [
"uniform_standard",
81,
4,
16,
true,
false,
6,
2
]
},
{
"id": 22,
"type": "WanVideoModelLoader",
"pos": [
620.3950805664062,
-357.8426818847656
],
"size": [
477.4410095214844,
226.43276977539062
],
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "compile_args",
"shape": 7,
"type": "WANCOMPILEARGS",
"link": null
},
{
"name": "block_swap_args",
"shape": 7,
"type": "BLOCKSWAPARGS",
"link": null
},
{
"name": "lora",
"shape": 7,
"type": "WANVIDLORA",
"link": null
},
{
"name": "vram_management_args",
"shape": 7,
"type": "VRAM_MANAGEMENTARGS",
"link": null
}
],
"outputs": [
{
"name": "model",
"type": "WANVIDEOMODEL",
"slot_index": 0,
"links": [
29
]
}
],
"properties": {
"Node name for S&R": "WanVideoModelLoader"
},
"widgets_values": [
"WanVideo\\wan2.1_t2v_1.3B_fp16.safetensors",
"fp16",
"disabled",
"offload_device",
"sdpa"
]
}
],
"links": [
[
15,
11,
0,
16,
0,
"WANTEXTENCODER"
],
[
29,
22,
0,
27,
0,
"WANVIDEOMODEL"
],
[
30,
16,
0,
27,
1,
"WANVIDEOTEXTEMBEDS"
],
[
33,
27,
0,
28,
1,
"LATENT"
],
[
42,
37,
0,
27,
2,
"WANVIDIMAGE_EMBEDS"
],
[
43,
38,
0,
28,
0,
"VAE"
],
[
48,
28,
0,
42,
0,
"IMAGE"
],
[
56,
42,
0,
30,
0,
"IMAGE"
],
[
57,
43,
0,
27,
5,
"WANVIDCONTEXT"
],
[
58,
45,
0,
27,
6,
"TEACACHEARGS"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.8140274938684471,
"offset": [
-122.25834160503663,
993.5739491626379
]
},
"node_versions": {
"ComfyUI-WanVideoWrapper": "5a2383621a05825d0d0437781afcb8552d9590fd",
"ComfyUI-KJNodes": "a5bd3c86c8ed6b83c55c2d0e7a59515b15a0137f",
"ComfyUI-VideoHelperSuite": "0a75c7958fe320efcb052f1d9f8451fd20c730a8"
},
"VHS_latentpreview": true,
"VHS_latentpreviewrate": 0,
"VHS_MetadataImage": true,
"VHS_KeepIntermediate": true
},
"version": 0.4
}
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+3 -1
View File
@@ -220,6 +220,7 @@ class WanVideoAddFantasyPortrait:
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the portrait embedding"}), "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the portrait embedding"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}), "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}), "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
"portrait_cfg": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 20.0, "step": 0.01, "tooltip": "CFG scale for the portrait embedding"}),
} }
} }
@@ -228,12 +229,13 @@ class WanVideoAddFantasyPortrait:
FUNCTION = "add" FUNCTION = "add"
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
def add(self, embeds, portrait_embeds, strength, start_percent=0.0, end_percent=1.0): def add(self, embeds, portrait_embeds, strength, start_percent=0.0, end_percent=1.0, portrait_cfg=1.0):
new_entry = { new_entry = {
"adapter_proj": portrait_embeds, "adapter_proj": portrait_embeds,
"strength": strength, "strength": strength,
"start_percent": start_percent, "start_percent": start_percent,
"end_percent": end_percent, "end_percent": end_percent,
"cfg_scale": portrait_cfg,
} }
updated = dict(embeds) updated = dict(embeds)
+1 -1
View File
@@ -167,7 +167,7 @@ class FantasyTalkingWav2VecEmbeds:
try: try:
audio_segment = audio_input[start_sample:end_sample] audio_segment = audio_input[start_sample:end_sample]
except: except Exception:
audio_segment = audio_input audio_segment = audio_input
print("audio_segment.shape", audio_segment.shape) print("audio_segment.shape", audio_segment.shape)
+7 -143
View File
@@ -1,15 +1,11 @@
import torch import torch
import torch.nn as nn
import numpy as np import numpy as np
import gguf import gguf
from accelerate import init_empty_weights
from .gguf_utils import GGUFParameter, dequantize_gguf_tensor from .gguf_utils import GGUFParameter
from ..utils import log
def load_gguf(model_path): def load_gguf(model_path):
from gguf import GGUFReader reader = gguf.GGUFReader(model_path)
reader = GGUFReader(model_path)
parsed_parameters = {} parsed_parameters = {}
for tensor in reader.tensors: for tensor in reader.tensors:
# if the tensor is a torch supported dtype do not use GGUFParameter # if the tensor is a torch supported dtype do not use GGUFParameter
@@ -18,144 +14,12 @@ def load_gguf(model_path):
parsed_parameters[tensor.name] = GGUFParameter(meta_tensor, quant_type=tensor.tensor_type) if is_gguf_quant else meta_tensor parsed_parameters[tensor.name] = GGUFParameter(meta_tensor, quant_type=tensor.tensor_type) if is_gguf_quant else meta_tensor
return parsed_parameters, reader return parsed_parameters, reader
#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py from ..custom_linear import _replace_linear, set_lora_params, CustomLinear
def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modules_to_not_convert=[], patches=None, compile_args=None): def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modules_to_not_convert=[], patches=None, compile_args=None):
def _should_convert_to_gguf(state_dict, prefix): return _replace_linear(model, compute_dtype, state_dict, prefix, patches, None, compile_args, modules_to_not_convert)
weight_key = prefix + "weight"
return weight_key in state_dict and isinstance(state_dict[weight_key], GGUFParameter)
has_children = list(model.children())
if not has_children:
return
allow_compile = False
for name, module in model.named_children():
if compile_args is not None:
allow_compile = compile_args.get("allow_unmerged_lora_compile", False)
module_prefix = prefix + name + "."
_replace_with_gguf_linear(module, compute_dtype, state_dict, module_prefix, modules_to_not_convert, patches, compile_args)
if (
isinstance(module, nn.Linear)
and not isinstance(module, GGUFLinear)
and _should_convert_to_gguf(state_dict, module_prefix)
and name not in modules_to_not_convert
):
in_features = state_dict[module_prefix + "weight"].shape[1]
out_features = state_dict[module_prefix + "weight"].shape[0]
with init_empty_weights():
model._modules[name] = GGUFLinear(
in_features,
out_features,
module.bias is not None,
compute_dtype=compute_dtype,
allow_compile=allow_compile
)
model._modules[name].source_cls = type(module)
model._modules[name].requires_grad_(False)
return model
def set_lora_params_gguf(module, patches, module_prefix="", device=torch.device("cpu")): def set_lora_params_gguf(module, patches, module_prefix="", device=torch.device("cpu")):
# Recursively set lora_diffs and lora_strengths for all GGUFLinear layers return set_lora_params(module, patches, module_prefix, device)
for name, child in module.named_children():
params = list(child.parameters())
if params:
device = params[0].device
else:
device = torch.device("cpu")
child_prefix = (f"{module_prefix}{name}.")
set_lora_params_gguf(child, patches, child_prefix, device)
if isinstance(module, GGUFLinear):
key = f"diffusion_model.{module_prefix}weight"
patch = patches.get(key, [])
#print(f"Processing LoRA patches for {key}: {len(patch)} patches found")
if len(patch) == 0:
key = key.replace("_orig_mod.", "")
patch = patches.get(key, [])
if len(patch) != 0:
lora_diffs = []
for p in patch:
lora_obj = p[1]
if "head" in key:
continue # For now skip LoRA for head layers
elif hasattr(lora_obj, "weights"):
lora_diffs.append(lora_obj.weights)
elif isinstance(lora_obj, tuple) and lora_obj[0] == "diff":
lora_diffs.append(lora_obj[1])
else:
continue
module.lora_strengths = [p[0] for p in patch]
module.set_lora_diffs(lora_diffs, device=device)
module.step = 0 # Initialize step for LoRA scheduling
GGUFLinear = CustomLinear
class GGUFLinear(nn.Linear):
def __init__(
self,
in_features,
out_features,
bias=False,
compute_dtype=None,
device=None,
allow_compile=False
) -> None:
super().__init__(in_features, out_features, bias, device)
self.compute_dtype = compute_dtype
self.lora_diffs = []
self.lora_strengths = []
self.step = 0
self.allow_compile = allow_compile
if not allow_compile:
self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora)
def forward(self, inputs):
weight = dequantize_gguf_tensor(self.weight).to(self.compute_dtype)
bias = self.bias.to(self.compute_dtype) if self.bias is not None else None
weight = self._get_weight_with_lora(weight)#.to(self.compute_dtype)
return torch.nn.functional.linear(inputs, weight, bias)
def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")):
self.lora_diffs = []
for i, diff in enumerate(lora_diffs):
if len(diff) > 1:
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device, self.compute_dtype))
setattr(self, f"lora_diff_{i}_2", diff[2])
self.lora_diffs.append((f"lora_diff_{i}_0", f"lora_diff_{i}_1", f"lora_diff_{i}_2"))
else:
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
self.lora_diffs.append(f"lora_diff_{i}_0")
def _get_weight_with_lora(self, weight):
"""Apply LoRA outside compiled region"""
if not hasattr(self, "lora_diff_0_0"):
return weight
for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths):
if isinstance(lora_strength, list):
lora_strength = lora_strength[self.step]
if lora_strength == 0.0:
continue
elif lora_strength == 0.0:
continue
if isinstance(lora_diff_names, tuple):
lora_diff_0 = getattr(self, lora_diff_names[0])
lora_diff_1 = getattr(self, lora_diff_names[1])
lora_diff_2 = getattr(self, lora_diff_names[2])
patch_diff = torch.mm(
lora_diff_0.flatten(start_dim=1),
lora_diff_1.flatten(start_dim=1)
).reshape(weight.shape) + 0
alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 is not None else 1.0
scale = lora_strength * alpha
weight = weight.add(patch_diff, alpha=scale)
else:
lora_diff = getattr(self, lora_diff_names)
weight = weight.add(lora_diff, alpha=lora_strength)
return weight
+1 -1
View File
@@ -85,7 +85,7 @@ def get_previewer(device, latent_format):
taesd = TAEHV(comfy.utils.load_torch_file(taehv_path)).to(device) taesd = TAEHV(comfy.utils.load_torch_file(taehv_path)).to(device)
previewer = TAESDPreviewerImpl(taesd) previewer = TAESDPreviewerImpl(taesd)
previewer = WrappedPreviewer(previewer, rate=16) previewer = WrappedPreviewer(previewer, rate=16)
except: except Exception:
log.info("Could not find TAEW model file 'taew2_1.safetensors' from models/vae_approx. You can download it from https://huggingface.co/Kijai/WanVideo_comfy/blob/main/taew2_1.safetensors") log.info("Could not find TAEW model file 'taew2_1.safetensors' from models/vae_approx. You can download it from https://huggingface.co/Kijai/WanVideo_comfy/blob/main/taew2_1.safetensors")
log.info("Using Latent2RGB previewer instead.") log.info("Using Latent2RGB previewer instead.")
method = LatentPreviewMethod.Latent2RGB method = LatentPreviewMethod.Latent2RGB
+46 -41
View File
@@ -42,41 +42,43 @@ def rotate_half(x):
x = torch.stack((-x2, x1), dim=-1) x = torch.stack((-x2, x1), dim=-1)
return rearrange(x, "... d r -> ... (d r)") return rearrange(x, "... d r -> ... (d r)")
def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, mode='mean', attn_bias=None): def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, split_num=4):
ref_k = ref_k.to(visual_q.dtype).to(visual_q.device)
scale = 1.0 / visual_q.shape[-1] ** 0.5 scale = 1.0 / visual_q.shape[-1] ** 0.5
visual_q = visual_q * scale visual_q = visual_q.transpose(1, 2) * scale
visual_q = visual_q.transpose(1, 2)
ref_k = ref_k.transpose(1, 2)
attn = visual_q @ ref_k.transpose(-2, -1)
if attn_bias is not None:
attn = attn + attn_bias
x_ref_attn_map_source = attn.softmax(-1) # B, H, x_seqlens, ref_seqlens
B, H, x_seqlens, K = visual_q.shape
x_ref_attn_maps = [] x_ref_attn_maps = []
ref_target_masks = ref_target_masks.to(visual_q.dtype)
x_ref_attn_map_source = x_ref_attn_map_source.to(visual_q.dtype)
for class_idx, ref_target_mask in enumerate(ref_target_masks): for class_idx, ref_target_mask in enumerate(ref_target_masks):
ref_target_mask = ref_target_mask[None, None, None, ...] ref_target_mask = ref_target_mask.view(1, 1, 1, -1)
x_ref_attnmap = x_ref_attn_map_source * ref_target_mask
x_ref_attnmap = x_ref_attnmap.sum(-1) / ref_target_mask.sum() # B, H, x_seqlens, ref_seqlens --> B, H, x_seqlens
x_ref_attnmap = x_ref_attnmap.permute(0, 2, 1) # B, x_seqlens, H
if mode == 'mean':
x_ref_attnmap = x_ref_attnmap.mean(-1) # B, x_seqlens
elif mode == 'max':
x_ref_attnmap = x_ref_attnmap.max(-1) # B, x_seqlens
x_ref_attn_maps.append(x_ref_attnmap)
del attn, x_ref_attn_map_source
return torch.concat(x_ref_attn_maps, dim=0) x_ref_attnmap = torch.zeros(B, H, x_seqlens, device=visual_q.device, dtype=visual_q.dtype)
chunk_size = min(max(x_seqlens // split_num, 1), x_seqlens)
for i in range(0, x_seqlens, chunk_size):
end_i = min(i + chunk_size, x_seqlens)
attn_chunk = visual_q[:, :, i:end_i] @ ref_k.permute(0, 2, 3, 1) # B, H, chunk, ref_seqlens
# Apply softmax
attn_max = attn_chunk.max(dim=-1, keepdim=True).values
attn_chunk = (attn_chunk - attn_max).exp()
attn_sum = attn_chunk.sum(dim=-1, keepdim=True)
attn_chunk = attn_chunk / (attn_sum + 1e-8)
# Apply mask and sum
masked_attn = attn_chunk * ref_target_mask
x_ref_attnmap[:, :, i:end_i] = masked_attn.sum(-1) / (ref_target_mask.sum() + 1e-8)
del attn_chunk, masked_attn
# Average across heads
x_ref_attnmap = x_ref_attnmap.mean(dim=1) # B, x_seqlens
x_ref_attn_maps.append(x_ref_attnmap)
del visual_q, ref_k
return torch.cat(x_ref_attn_maps, dim=0)
def get_attn_map_with_target(visual_q, ref_k, shape, ref_target_masks=None, split_num=2): def get_attn_map_with_target(visual_q, ref_k, shape, ref_target_masks=None, split_num=2):
"""Args: """Args:
@@ -129,27 +131,30 @@ class RotaryPositionalEmbedding1D(nn.Module):
query with the same shape as input. query with the same shape as input.
""" """
freqs_cis = self.precompute_freqs_cis_1d(pos_indices) freqs_cis = self.precompute_freqs_cis_1d(pos_indices)
in_dtype = x.dtype
x_ = x.float() x = x.float()
freqs_cis = freqs_cis.float().to(x.device) freqs_cis = freqs_cis.float().to(x.device)
cos, sin = freqs_cis.cos(), freqs_cis.sin() cos = rearrange(freqs_cis.cos(), 'n d -> 1 1 n d')
cos, sin = rearrange(cos, 'n d -> 1 1 n d'), rearrange(sin, 'n d -> 1 1 n d') sin = rearrange(freqs_cis.sin(), 'n d -> 1 1 n d')
x_ = (x_ * cos) + (rotate_half(x_) * sin)
return x_.type_as(x) # In-place rotation to save memory
x_rotated = rotate_half(x)
x.mul_(cos).add_(x_rotated * sin)
return x.to(in_dtype)
class AudioProjModel(nn.Module): class AudioProjModel(nn.Module):
def __init__( def __init__(
self, self,
seq_len=5, seq_len=5,
seq_len_vf=12, seq_len_vf=8,
blocks=12, blocks=12,
channels=768, channels=768,
intermediate_dim=512, intermediate_dim=512,
output_dim=768, output_dim=768,
context_tokens=32, context_tokens=32,
norm_output_audio=False, norm_output_audio=True,
): ):
super().__init__() super().__init__()
@@ -273,9 +278,9 @@ class SingleStreamMultiAttention(SingleStreamAttention):
def __init__( def __init__(
self, self,
dim: int, dim: int,
encoder_hidden_states_dim: int,
num_heads: int, num_heads: int,
qkv_bias: bool, qkv_bias: bool = True,
encoder_hidden_states_dim: int = 768,
class_range: int = 24, class_range: int = 24,
class_interval: int = 4, class_interval: int = 4,
attention_mode: str = 'sdpa', attention_mode: str = 'sdpa',
+569
View File
@@ -0,0 +1,569 @@
import torch
import os
import gc
from PIL import Image
import numpy as np
from ..latent_preview import prepare_callback
from ..wanvideo.schedulers import get_scheduler
from .multitalk import timestep_transform, add_noise
from ..utils import log, print_memory, temporal_score_rescaling, offload_transformer, init_blockswap, match_and_blend_colors
from comfy.utils import load_torch_file
from ..nodes_model_loading import load_weights
from ..HuMo.nodes import get_audio_emb_window
import comfy.model_management as mm
from tqdm import tqdm
import copy
VAE_STRIDE = (4, 8, 8)
PATCH_SIZE = (1, 2, 2)
vae_upscale_factor = 8
script_directory = os.path.dirname(os.path.abspath(__file__))
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
def multitalk_loop(self, **kwargs):
# Unpack kwargs into local variables
(latent, total_steps, steps, start_step, end_step, shift, cfg, denoise_strength,
sigmas, weight_dtype, transformer, patcher, block_swap_args, model, vae, dtype,
scheduler, scheduler_step_args, text_embeds, image_embeds, multitalk_embeds,
multitalk_audio_embeds, unianim_data, dwpose_data, unianimate_poses, uni3c_embeds,
humo_image_cond, humo_image_cond_neg, humo_audio, humo_reference_count,
add_noise_to_samples, audio_stride, use_tsr, tsr_k, tsr_sigma, fantasy_portrait_input,
noise, timesteps, force_offload, add_cond, control_latents, audio_proj,
control_camera_latents, samples, masks, seed_g, gguf_reader, predict_func
) = (kwargs.get(k) for k in (
'latent', 'total_steps', 'steps', 'start_step', 'end_step', 'shift', 'cfg',
'denoise_strength', 'sigmas', 'weight_dtype', 'transformer', 'patcher',
'block_swap_args', 'model', 'vae', 'dtype', 'scheduler', 'scheduler_step_args',
'text_embeds', 'image_embeds', 'multitalk_embeds', 'multitalk_audio_embeds',
'unianim_data', 'dwpose_data', 'unianimate_poses', 'uni3c_embeds',
'humo_image_cond', 'humo_image_cond_neg', 'humo_audio', 'humo_reference_count',
'add_noise_to_samples', 'audio_stride', 'use_tsr', 'tsr_k', 'tsr_sigma',
'fantasy_portrait_input', 'noise', 'timesteps', 'force_offload', 'add_cond',
'control_latents', 'audio_proj', 'control_camera_latents', 'samples', 'masks',
'seed_g', 'gguf_reader', 'predict_with_cfg'
))
mode = image_embeds.get("multitalk_mode", "multitalk")
if mode == "auto":
mode = transformer.multitalk_model_type.lower()
elif mode == "skyreelsv3":
num_pseudo_frames = 5
pseudo_frames = reference_keyframes = None
keyframe_index = 0
reference_video = image_embeds.get("reference_video", None)
log.info(f"Multitalk mode: {mode}")
drop_frames = image_embeds.get("drop_frames", 0)
cond_frame = None
offload = image_embeds.get("force_offload", False)
offloaded = False
tiled_vae = image_embeds.get("tiled_vae", False)
frame_num = clip_length = image_embeds.get("frame_window_size", 81)
clip_embeds = image_embeds.get("clip_context", None)
if clip_embeds is not None:
clip_embeds = clip_embeds.to(dtype)
colormatch = image_embeds.get("colormatch", "disabled")
motion_frame = image_embeds.get("motion_frame", 25)
target_w = image_embeds.get("target_w", None)
target_h = image_embeds.get("target_h", None)
original_images = image_embeds.get("multitalk_start_image", None)
cond_image = original_images.clone() if original_images is not None else None
original_color_reference = cond_image.clone() if cond_image is not None else None
if original_images is None:
original_images = torch.zeros([noise.shape[0], 1, target_h, target_w], device=device)
output_path = image_embeds.get("output_path", "")
img_counter = 0
if len(multitalk_embeds['audio_features'])==2 and (multitalk_embeds['ref_target_masks'] is None):
face_scale = 0.1
x_min, x_max = int(target_h * face_scale), int(target_h * (1 - face_scale))
lefty_min, lefty_max = int((target_w//2) * face_scale), int((target_w//2) * (1 - face_scale))
righty_min, righty_max = int((target_w//2) * face_scale + (target_w//2)), int((target_w//2) * (1 - face_scale) + (target_w//2))
human_mask1, human_mask2 = (torch.zeros([target_h, target_w]) for _ in range(2))
human_mask1[x_min:x_max, lefty_min:lefty_max] = 1
human_mask2[x_min:x_max, righty_min:righty_max] = 1
background_mask = torch.where((human_mask1 + human_mask2) > 0, torch.tensor(0), torch.tensor(1))
human_masks = [human_mask1, human_mask2, background_mask]
ref_target_masks = torch.stack(human_masks, dim=0)
multitalk_embeds['ref_target_masks'] = ref_target_masks
gen_video_list = []
is_first_clip = True
arrive_last_frame = False
cur_motion_frames_num = 1
audio_start_idx = iteration_count = step_iteration_count = 0
audio_end_idx = (audio_start_idx + clip_length) * audio_stride
indices = (torch.arange(4 + 1) - 2) * 1
current_condframe_index = 0
audio_embedding = multitalk_audio_embeds
human_num = len(audio_embedding)
audio_embs = None
uni3c_data = None
if uni3c_embeds is not None:
transformer.controlnet = uni3c_embeds["controlnet"]
uni3c_data = uni3c_embeds.copy()
encoded_silence = None
try:
silence_path = os.path.join(script_directory, "encoded_silence.safetensors")
encoded_silence = load_torch_file(silence_path)["audio_emb"].to(dtype)
except Exception:
log.warning("No encoded silence file found, padding with end of audio embedding instead.")
total_frames = len(audio_embedding[0])
estimated_iterations = total_frames // (frame_num - motion_frame - drop_frames) + 1
callback = prepare_callback(patcher, estimated_iterations)
# If reference_video is provided, extract keyframes from it
if mode == "skyreelsv3" and reference_video is not None:
ref_video_length = reference_video.shape[1] # (C, T, H, W)
if colormatch == "reinhard_torch":
reference_video = match_and_blend_colors(reference_video, original_color_reference, 1.0)
if ref_video_length >= total_frames:
# Reference is long enough - extract keyframes at the expected positions
segment_interval = frame_num - motion_frame - drop_frames
generate_idx = []
current_idx = frame_num - 1
while current_idx < total_frames:
generate_idx.append(min(current_idx, ref_video_length - 1))
current_idx += segment_interval
else:
# Calculate target indices then map to reference video
audio_length = total_frames
generate_idx_target = [0]
segment_interval = frame_num - motion_frame - drop_frames
current_idx = frame_num - 1
while current_idx < audio_length - 1:
generate_idx_target.append(current_idx)
current_idx += segment_interval
if generate_idx_target[-1] != audio_length - 1:
generate_idx_target.append(audio_length - 1)
# Map target indices to reference video
generate_idx_target = np.array(generate_idx_target, dtype=np.int16)
original_max = generate_idx_target[-1]
original_min = generate_idx_target[0]
if original_max > original_min:
generate_idx_float = (generate_idx_target.astype(np.float64) - original_min) * (ref_video_length - 1) / (original_max - original_min)
generate_idx = np.clip(np.round(generate_idx_float), 0, ref_video_length - 1).astype(np.int32).tolist()
else:
generate_idx = [0]
generate_idx = generate_idx[1:]
log.info(f"Reference video ({ref_video_length} frames) mapped to target ({total_frames} frames). Keyframe indices: {generate_idx}")
# Extract keyframes from reference video
# reference_video shape: (C, T, H, W) from nodes.py processing
# Select keyframes and add batch dimension: (C, num_keyframes, H, W) -> (1, C, num_keyframes, H, W)
selected_keyframes = reference_video[:, generate_idx] # (C, num_keyframes, H, W)
reference_keyframes = selected_keyframes.unsqueeze(0).cpu() # (1, C, num_keyframes, H, W)
log.info(f"Extracted {len(generate_idx)} keyframes from provided reference video at indices {generate_idx}, shape: {reference_keyframes.shape}")
log.info(f"Reference video total frames: {reference_video.shape[1]}, will generate {total_frames} total frames with {estimated_iterations} windows")
if frame_num >= total_frames:
arrive_last_frame = True
estimated_iterations = 1
log.info(f"Sampling {total_frames} frames in {estimated_iterations} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
while True: # start video generation iteratively
self.cache_state = [None, None]
if mode == "skyreelsv3" and reference_keyframes is not None:
clamped_index = min(keyframe_index, reference_keyframes.shape[2] - 1) # Clamp keyframe_index to reuse last keyframe if we run out
pseudo_frames = reference_keyframes[:, :, clamped_index:clamped_index+1].repeat(1, 1, num_pseudo_frames, 1, 1) # Use one keyframe and repeat it 5 times
log.info(f"Window {iteration_count}: using keyframe {clamped_index}/{reference_keyframes.shape[2]-1} for pseudo frames.")
keyframe_index += 1
else:
pseudo_frames = None
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
if mode == "infinitetalk":
cond_image = original_images[:, :, current_condframe_index:current_condframe_index+1] if cond_image is not None else None
if multitalk_embeds is not None:
audio_embs = []
# split audio with window size
for human_idx in range(human_num):
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
center_indices = torch.clamp(center_indices, min=0, max=audio_embedding[human_idx].shape[0]-1)
audio_emb = audio_embedding[human_idx][center_indices].unsqueeze(0).to(device)
audio_embs.append(audio_emb)
audio_embs = torch.cat(audio_embs, dim=0).to(dtype)
h, w = (cond_image.shape[-2], cond_image.shape[-1]) if cond_image is not None else (target_h, target_w)
lat_h, lat_w = h // VAE_STRIDE[1], w // VAE_STRIDE[2]
latent_frame_num = (frame_num - 1) // 4 + 1
noise = torch.randn(16, latent_frame_num, lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
# Calculate the correct latent slice based on current iteration
if is_first_clip:
latent_start_idx = 0
latent_end_idx = noise.shape[1]
else:
new_frames_per_iteration = frame_num - motion_frame
new_latent_frames_per_iteration = ((new_frames_per_iteration - 1) // 4 + 1)
latent_start_idx = iteration_count * new_latent_frames_per_iteration
latent_end_idx = latent_start_idx + noise.shape[1]
if samples is not None:
noise_mask = samples.get("noise_mask", None)
input_samples = samples["samples"]
if input_samples is not None:
input_samples = input_samples.squeeze(0).to(noise)
# Check if we have enough frames in input_samples
if latent_end_idx > input_samples.shape[1]:
# We need more frames than available - pad the input_samples at the end
pad_length = latent_end_idx - input_samples.shape[1]
last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
input_samples = torch.cat([input_samples, last_frame], dim=1)
input_samples = input_samples[:, latent_start_idx:latent_end_idx]
if noise_mask is not None:
original_image = input_samples.to(device)
assert input_samples.shape[1] == noise.shape[1], f"Slice mismatch: {input_samples.shape[1]} vs {noise.shape[1]}"
if add_noise_to_samples:
latent_timestep = timesteps[0]
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
else:
noise = input_samples
# diff diff prep
if noise_mask is not None:
if len(noise_mask.shape) == 4:
noise_mask = noise_mask.squeeze(1)
if audio_end_idx > noise_mask.shape[0]:
noise_mask = noise_mask.repeat(audio_end_idx // noise_mask.shape[0], 1, 1)
noise_mask = noise_mask[audio_start_idx:audio_end_idx]
noise_mask = torch.nn.functional.interpolate(
noise_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
size=(noise.shape[1], noise.shape[2], noise.shape[3]),
mode='trilinear',
align_corners=False
).repeat(1, noise.shape[0], 1, 1, 1)
thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps)
thresholds = thresholds.reshape(-1, 1, 1, 1, 1).to(device)
masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds
# zero padding and vae encode for img cond
if cond_image is not None or cond_frame is not None:
cond_ = cond_image if (is_first_clip or humo_image_cond is None) else cond_frame
cond_frame_num = cond_.shape[2]
# Prepare pseudo frames if enabled and available from reference_video
if mode == "skyreelsv3" and pseudo_frames is not None:
video_frames = torch.zeros(1, 3, frame_num-cond_frame_num-num_pseudo_frames, target_h, target_w, device=device, dtype=vae.dtype)
padding_frames_pixels_values = torch.cat([cond_.to(device, vae.dtype), video_frames, pseudo_frames.to(device, vae.dtype)], dim=2)
else:
video_frames = torch.zeros(1, 3, frame_num-cond_frame_num, target_h, target_w, device=device, dtype=vae.dtype)
padding_frames_pixels_values = torch.cat([cond_.to(device, vae.dtype), video_frames], dim=2)
# encode
vae.to(device)
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
if mode == "infinitetalk":
cond_ = cond_image if is_first_clip else cond_frame
latent_motion_frames = vae.encode(cond_.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
else:
latent_motion_frames = y[:, :cur_motion_frames_latent_num] # C T H W
vae.to(offload_device)
#motion_frame_index = cur_motion_frames_latent_num if mode == "infinitetalk" else 1
if mode == "skyreelsv3" and pseudo_frames is not None:
# create mask in pixel space, then transform
msk_pixel = torch.ones(1, frame_num, lat_h, lat_w, device=device)
msk_pixel[:, cur_motion_frames_num : -num_pseudo_frames] = 0
msk_pixel = torch.cat([
torch.repeat_interleave(msk_pixel[:, 0:1], repeats=4, dim=1),
msk_pixel[:, 1:],
], dim=1)
msk_pixel = msk_pixel.view(1, msk_pixel.shape[1] // 4, 4, lat_h, lat_w)
msk = msk_pixel.transpose(1, 2).squeeze(0).to(dtype) # 4 T H W
else:
msk = torch.zeros(4, latent_frame_num, lat_h, lat_w, device=device, dtype=dtype)
msk[:, :1] = 1
y = torch.cat([msk, y]) # 4+C T H W
mm.soft_empty_cache()
else:
y = None
latent_motion_frames = noise[:, :1]
partial_humo_cond_input = partial_humo_cond_neg_input = partial_humo_audio = partial_humo_audio_neg = None
if humo_image_cond is not None:
partial_humo_cond_input = humo_image_cond[:, :latent_frame_num]
partial_humo_cond_neg_input = humo_image_cond_neg[:, :latent_frame_num]
if y is not None:
partial_humo_cond_input[:, :1] = y[:, :1]
if humo_reference_count > 0:
partial_humo_cond_input[:, -humo_reference_count:] = humo_image_cond[:, -humo_reference_count:]
partial_humo_cond_neg_input[:, -humo_reference_count:] = humo_image_cond_neg[:, -humo_reference_count:]
if humo_audio is not None:
if is_first_clip:
audio_embs = None
partial_humo_audio, _ = get_audio_emb_window(humo_audio, frame_num, frame0_idx=audio_start_idx)
#zero_audio_pad = torch.zeros(humo_reference_count, *partial_humo_audio.shape[1:], device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
partial_humo_audio[-humo_reference_count:] = 0
partial_humo_audio_neg = torch.zeros_like(partial_humo_audio, device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
if scheduler == "multitalk":
timesteps = list(np.linspace(1000, 1, steps, dtype=np.float32))
timesteps.append(0.)
timesteps = [torch.tensor([t], device=device) for t in timesteps]
timesteps = [timestep_transform(t, shift=shift, num_timesteps=1000) for t in timesteps]
else:
if isinstance(scheduler, dict):
sample_scheduler = copy.deepcopy(scheduler["sample_scheduler"])
timesteps = scheduler["timesteps"]
else:
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, denoise_strength, sigmas=sigmas)
timesteps = [torch.tensor([float(t)], device=device) for t in timesteps] + [torch.tensor([0.], device=device)]
# sample videos
latent = noise
# injecting motion frames
if not is_first_clip and mode != "infinitetalk":
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[0])
latent[:, :add_latent.shape[1]] = add_latent
del motion_add_noise, add_latent
if offloaded:
# Load weights
if transformer.patched_linear and gguf_reader is None:
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args)
elif gguf_reader is not None: #handle GGUF
load_weights(transformer, patcher.model["sd"], base_dtype=dtype, transformer_load_device=device, patcher=patcher, gguf=True, reader=gguf_reader, block_swap_args=block_swap_args)
#blockswap init
init_blockswap(transformer, block_swap_args, model)
# Use the appropriate prompt for this section
if len(text_embeds["prompt_embeds"]) > 1:
prompt_index = min(iteration_count, len(text_embeds["prompt_embeds"]) - 1)
positive = [text_embeds["prompt_embeds"][prompt_index]]
log.info(f"Using prompt index: {prompt_index}")
else:
positive = text_embeds["prompt_embeds"]
# uni3c slices
if uni3c_embeds is not None:
vae.to(device)
# Pad original_images if needed
num_frames = original_images.shape[2]
if audio_end_idx > num_frames:
pad_len = audio_end_idx - num_frames
last_frame = original_images[:, :, -1:].repeat(1, 1, pad_len, 1, 1)
padded_images = torch.cat([original_images, last_frame], dim=2)
else:
padded_images = original_images
render_latent = vae.encode(
padded_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype),
device=device, tiled=tiled_vae
).to(dtype)
vae.to(offload_device)
uni3c_data['render_latent'] = render_latent
# unianimate slices
partial_unianim_data = None
if unianim_data is not None:
partial_dwpose = dwpose_data[:, :, latent_start_idx:latent_end_idx]
partial_unianim_data = {
"dwpose": partial_dwpose,
"random_ref": unianim_data["random_ref"],
"strength": unianimate_poses["strength"],
"start_percent": unianimate_poses["start_percent"],
"end_percent": unianimate_poses["end_percent"]
}
# fantasy portrait slices
partial_fantasy_portrait_input = None
if fantasy_portrait_input is not None:
adapter_proj = fantasy_portrait_input["adapter_proj"]
if latent_end_idx > adapter_proj.shape[1]:
pad_len = latent_end_idx - adapter_proj.shape[1]
last_frame = adapter_proj[:, -1:, :, :].repeat(1, pad_len, 1, 1)
padded_proj = torch.cat([adapter_proj, last_frame], dim=1)
else:
padded_proj = adapter_proj
partial_fantasy_portrait_input = fantasy_portrait_input.copy()
partial_fantasy_portrait_input["adapter_proj"] = padded_proj[:, latent_start_idx:latent_end_idx]
mm.soft_empty_cache()
gc.collect()
# sampling loop
sampling_pbar = tqdm(total=len(timesteps)-1, desc=f"Sampling audio indices {audio_start_idx}-{audio_end_idx}", position=0, leave=True)
for i in range(len(timesteps)-1):
timestep = timesteps[i]
latent_model_input = latent.to(device)
if mode == "infinitetalk":
if humo_image_cond is None or not is_first_clip:
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
noise_pred, _, self.cache_state = predict_func(
latent_model_input, cfg[min(i, len(timesteps)-1)], positive, text_embeds["negative_prompt_embeds"],
timestep, i, y, clip_embeds, control_latents, None, partial_unianim_data, audio_proj, control_camera_latents, add_cond,
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs, fantasy_portrait_input=partial_fantasy_portrait_input,
humo_image_cond=partial_humo_cond_input, humo_image_cond_neg=partial_humo_cond_neg_input, humo_audio=partial_humo_audio, humo_audio_neg=partial_humo_audio_neg,
uni3c_data = uni3c_data)
if callback is not None:
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * timestep.to(device) / 1000).detach().permute(1,0,2,3)
callback(step_iteration_count, callback_latent, None, estimated_iterations*(len(timesteps)-1))
del callback_latent
sampling_pbar.update(1)
step_iteration_count += 1
# update latent
if use_tsr:
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
if scheduler == "multitalk":
noise_pred = -noise_pred
dt = (timesteps[i] - timesteps[i + 1]) / 1000
latent = latent + noise_pred * dt[:, None, None, None]
else:
latent = sample_scheduler.step(noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0).to(noise_pred.device), **scheduler_step_args)[0].squeeze(0)
del noise_pred, latent_model_input, timestep
# differential diffusion inpaint
if masks is not None:
if i < len(timesteps) - 1:
image_latent = add_noise(original_image.to(device), noise.to(device), timesteps[i+1])
mask = masks[i].to(latent)
latent = image_latent * mask + latent * (1-mask)
# injecting motion frames
if not is_first_clip and mode != "infinitetalk":
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1])
latent[:, :add_latent.shape[1]] = add_latent
del motion_add_noise, add_latent
elif mode == "infinitetalk":
if humo_image_cond is None or not is_first_clip:
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
del noise, latent_motion_frames
if offload:
offload_transformer(transformer, remove_lora=False)
offloaded = True
if humo_image_cond is not None and humo_reference_count > 0:
latent = latent[:,:-humo_reference_count]
vae.to(device)
videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
vae.to(offload_device)
sampling_pbar.close()
# crop drop_frames from end if enabled
if mode == "skyreelsv3" and drop_frames > 0 and not arrive_last_frame:
videos = videos[:, :-drop_frames]
# optional color correction (less relevant for InfiniteTalk)
if colormatch != "disabled":
if colormatch == "reinhard_torch":
videos = match_and_blend_colors(videos, original_color_reference, 1.0)
else:
videos = videos.permute(1, 2, 3, 0).float().numpy()
from color_matcher import ColorMatcher
cm = ColorMatcher()
cm_result_list = []
for img in videos:
if mode == "infinitetalk":
cm_result = cm.transfer(src=img, ref=cond_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
else:
cm_result = cm.transfer(src=img, ref=original_images[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
cm_result_list.append(torch.from_numpy(cm_result).to(vae.dtype))
videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2)
# optionally save generated samples to disk
if output_path:
video_np = videos.clamp(-1.0, 1.0).add(1.0).div(2.0).mul(255).cpu().float().numpy().transpose(1, 2, 3, 0).astype('uint8')
num_frames_to_save = video_np.shape[0] if is_first_clip else video_np.shape[0] - cur_motion_frames_num
log.info(f"Saving {num_frames_to_save} generated frames to {output_path}")
start_idx = 0 if is_first_clip else cur_motion_frames_num
for i in range(start_idx, video_np.shape[0]):
im = Image.fromarray(video_np[i])
im.save(os.path.join(output_path, f"frame_{img_counter:05d}.png"))
img_counter += 1
else:
gen_video_list.append(videos if is_first_clip else videos[:, cur_motion_frames_num:])
current_condframe_index += 1
iteration_count += 1
# decide whether is done
if arrive_last_frame:
break
# update next condition frames
is_first_clip = False
cur_motion_frames_num = motion_frame
cond_ = videos[:, -cur_motion_frames_num:].unsqueeze(0)
if mode == "infinitetalk":
cond_frame = cond_
else:
cond_image = cond_
del videos, latent
# Repeat audio emb
if multitalk_embeds is not None:
audio_start_idx += (frame_num - cur_motion_frames_num - humo_reference_count - drop_frames)
audio_end_idx = audio_start_idx + clip_length
if audio_end_idx >= len(audio_embedding[0]):
arrive_last_frame = True
miss_lengths = []
source_frames = []
for human_inx in range(human_num):
source_frame = len(audio_embedding[human_inx])
source_frames.append(source_frame)
if audio_end_idx >= len(audio_embedding[human_inx]):
log.warning(f"Audio embedding for subject {human_inx} not long enough: {len(audio_embedding[human_inx])}, need {audio_end_idx}, padding...")
miss_length = audio_end_idx - len(audio_embedding[human_inx]) + 3
log.warning(f"Padding length: {miss_length}")
if encoded_silence is not None:
add_audio_emb = encoded_silence[-1*miss_length:]
else:
add_audio_emb = torch.flip(audio_embedding[human_inx][-1*miss_length:], dims=[0])
audio_embedding[human_inx] = torch.cat([audio_embedding[human_inx], add_audio_emb.to(device, dtype)], dim=0)
miss_lengths.append(miss_length)
else:
miss_lengths.append(0)
if mode == "infinitetalk" and current_condframe_index >= original_images.shape[2]:
last_frame = original_images[:, :, -1:, :, :]
miss_length = 1
original_images = torch.cat([original_images, last_frame.repeat(1, 1, miss_length, 1, 1)], dim=2)
if not output_path:
gen_video_samples = torch.cat(gen_video_list, dim=1)
else:
gen_video_samples = torch.zeros(3, 1, 64, 64) # dummy output
if force_offload:
if not model["auto_cpu_offload"]:
offload_transformer(transformer)
try:
print_memory(device)
torch.cuda.reset_peak_memory_stats(device)
except Exception:
pass
return {"video": gen_video_samples.permute(1, 2, 3, 0), "output_path": output_path},
+109 -3
View File
@@ -7,6 +7,8 @@ from ..utils import log, set_module_tensor_to_device
import os import os
import json import json
import datetime import datetime
import scipy.signal as ss
import numpy as np
script_directory = os.path.dirname(os.path.abspath(__file__)) script_directory = os.path.dirname(os.path.abspath(__file__))
folder_paths.add_model_folder_path("wav2vec2", os.path.join(folder_paths.models_dir, "wav2vec2")) folder_paths.add_model_folder_path("wav2vec2", os.path.join(folder_paths.models_dir, "wav2vec2"))
@@ -126,7 +128,7 @@ class MultiTalkModelLoader:
def loudness_norm(audio_array, sr=16000, lufs=-23): def loudness_norm(audio_array, sr=16000, lufs=-23):
try: try:
import pyloudnorm import pyloudnorm
except: except Exception:
raise ImportError("pyloudnorm package is not installed") raise ImportError("pyloudnorm package is not installed")
meter = pyloudnorm.Meter(sr) meter = pyloudnorm.Meter(sr)
loudness = meter.integrated_loudness(audio_array) loudness = meter.integrated_loudness(audio_array)
@@ -134,6 +136,15 @@ def loudness_norm(audio_array, sr=16000, lufs=-23):
return audio_array return audio_array
normalized_audio = pyloudnorm.normalize.loudness(audio_array, loudness, lufs) normalized_audio = pyloudnorm.normalize.loudness(audio_array, loudness, lufs)
return normalized_audio return normalized_audio
def _add_noise_floor(audio, noise_db=-45):
noise_amp = 10 ** (noise_db / 20)
noise = np.random.randn(len(audio)) * noise_amp
return audio + noise
def _smooth_transients(audio, sr=16000):
b, a = ss.butter(3, 3000 / (sr/2))
return ss.lfilter(b, a, audio)
class MultiTalkWav2VecEmbeds: class MultiTalkWav2VecEmbeds:
@classmethod @classmethod
@@ -153,6 +164,8 @@ class MultiTalkWav2VecEmbeds:
"audio_3": ("AUDIO",), "audio_3": ("AUDIO",),
"audio_4": ("AUDIO",), "audio_4": ("AUDIO",),
"ref_target_masks": ("MASK", {"tooltip": "Per-speaker semantic mask(s) in pixel space. Supply one mask per speaker (plus optional background) to guide mouth assignment"}), "ref_target_masks": ("MASK", {"tooltip": "Per-speaker semantic mask(s) in pixel space. Supply one mask per speaker (plus optional background) to guide mouth assignment"}),
"add_noise_floor": ("BOOLEAN", {"default": False, "tooltip": "Add a low-level noise floor to the audio to reduce silent gaps"}),
"smooth_transients": ("BOOLEAN", {"default": False, "tooltip": "Apply a low-pass filter to the audio to smooth out transients"}),
} }
} }
@@ -161,7 +174,8 @@ class MultiTalkWav2VecEmbeds:
FUNCTION = "process" FUNCTION = "process"
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
def process(self, wav2vec_model, normalize_loudness, fps, num_frames, audio_1, audio_scale, audio_cfg_scale, multi_audio_type, audio_2=None, audio_3=None, audio_4=None, ref_target_masks=None): def process(self, wav2vec_model, normalize_loudness, fps, num_frames, audio_1, audio_scale, audio_cfg_scale, multi_audio_type, audio_2=None, audio_3=None, audio_4=None,
ref_target_masks=None, add_noise_floor=False, smooth_transients=False):
model_type = wav2vec_model["model_type"] model_type = wav2vec_model["model_type"]
if not "tencent" in model_type.lower(): if not "tencent" in model_type.lower():
raise ValueError("Only tencent wav2vec2 models supported by MultiTalk") raise ValueError("Only tencent wav2vec2 models supported by MultiTalk")
@@ -207,6 +221,10 @@ class MultiTalkWav2VecEmbeds:
if normalize_loudness: if normalize_loudness:
audio_segment = loudness_norm(audio_segment, sr=sr) audio_segment = loudness_norm(audio_segment, sr=sr)
if add_noise_floor:
audio_segment = _add_noise_floor(audio_segment, noise_db=-45)
if smooth_transients:
audio_segment = _smooth_transients(audio_segment, sr=sr)
audio_feature = np.squeeze( audio_feature = np.squeeze(
wav2vec2_feature_extractor(audio_segment, sampling_rate=sr).input_values wav2vec2_feature_extractor(audio_segment, sampling_rate=sr).input_values
@@ -443,13 +461,100 @@ class WanVideoImageToVideoMultiTalk:
} }
return (image_embeds, output_path) return (image_embeds, output_path)
class WanVideoImageToVideoSkyreelsv3_audio:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"vae": ("WANVAE",),
"width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the generation"}),
"height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the generation"}),
"frame_window_size": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "The number of frames to process at once, should be a value the model is generally good at."}),
"motion_frame": ("INT", {"default": 5, "min": 1, "max": 10000, "step": 1, "tooltip": "Driven frame length used in the long video generation. Basically the overlap length."}),
"drop_frames": ("INT", {"default": 12, "min": 0, "max": 10000, "step": 1, "tooltip": "Additional frames to drop when advancing the audio window. Higher values = less overlap = faster generation but potentially less smooth transitions."}),
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
"force_offload": ("BOOLEAN", {"default": False, "tooltip": "Whether to force offload the model within the loop for VAE operations, enable if you encounter memory issues."}),
"colormatch": (
[
'disabled',
'reinhard_torch',
'mkl',
'hm',
'reinhard',
'mvgd',
'hm-mvgd-hm',
'hm-mkl-hm',
], {
"default": 'disabled', "tooltip": "Color matching method to use between the windows"
},),
},
"optional": {
"start_image": ("IMAGE", {"tooltip": "Images to encode"}),
"reference_video": ("IMAGE", {"tooltip": "Optional: Pre-generated reference video to use for keyframes instead of extracting from first generation. Should be color-matched to source image."}),
"clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}),
"output_path": ("STRING", {"default": "", "tooltip": "If set, will save each window's resulting frames to this folder, also DISABLES returning the final video tensor to save memory"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "STRING",)
RETURN_NAMES = ("image_embeds", "output_path")
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Enables Multi/InfiniteTalk long video generation sampling method, the video is created in windows with overlapping frames. Not compatible or necessary to be used with context windows and many other features besides Multi/InfiniteTalk."
def process(self, vae, width, height, frame_window_size, motion_frame, drop_frames, force_offload, colormatch, start_image=None,
tiled_vae=False, clip_embeds=None, mode="multitalk", output_path="", reference_video=None):
H, W = height, width
num_frames = ((frame_window_size - 1) // 4) * 4 + 1
# Resize and rearrange the input image dimensions
if start_image is not None:
resized_start_image = common_upscale(start_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
resized_start_image = resized_start_image * 2 - 1
resized_start_image = resized_start_image.unsqueeze(0)
target_shape = (16, (num_frames - 1) // 4 + 1, height // 8, width // 8)
if output_path:
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
output_path = os.path.join(output_path, f"{timestamp}_{mode}_output")
os.makedirs(output_path, exist_ok=True)
processed_reference_video = None
if reference_video is not None:
processed_reference_video = common_upscale(reference_video.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
processed_reference_video = processed_reference_video * 2 - 1
image_embeds = {
"multitalk_sampling": True,
"multitalk_start_image": resized_start_image if start_image is not None else None,
"frame_window_size": num_frames,
"motion_frame": motion_frame,
"drop_frames": drop_frames,
"use_pseudo_frames": True,
"reference_video": processed_reference_video,
"target_h": H,
"target_w": W,
"tiled_vae": tiled_vae,
"force_offload": force_offload,
"vae": vae,
"target_shape": target_shape,
"clip_context": clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None,
"colormatch": colormatch,
"multitalk_mode": "skyreelsv3",
"output_path": output_path
}
return (image_embeds, output_path)
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"MultiTalkModelLoader": MultiTalkModelLoader, "MultiTalkModelLoader": MultiTalkModelLoader,
"MultiTalkWav2VecEmbeds": MultiTalkWav2VecEmbeds, "MultiTalkWav2VecEmbeds": MultiTalkWav2VecEmbeds,
"WanVideoImageToVideoMultiTalk": WanVideoImageToVideoMultiTalk, "WanVideoImageToVideoMultiTalk": WanVideoImageToVideoMultiTalk,
"Wav2VecModelLoader": Wav2VecModelLoader, "Wav2VecModelLoader": Wav2VecModelLoader,
"MultiTalkSilentEmbeds": MultiTalkSilentEmbeds, "MultiTalkSilentEmbeds": MultiTalkSilentEmbeds,
"WanVideoImageToVideoSkyreelsv3_audio": WanVideoImageToVideoSkyreelsv3_audio,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
@@ -458,4 +563,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoImageToVideoMultiTalk": "WanVideo Long I2V Multi/InfiniteTalk", "WanVideoImageToVideoMultiTalk": "WanVideo Long I2V Multi/InfiniteTalk",
"Wav2VecModelLoader": "Wav2vec2 Model Loader", "Wav2VecModelLoader": "Wav2vec2 Model Loader",
"MultiTalkSilentEmbeds": "MultiTalk Silent Embeds", "MultiTalkSilentEmbeds": "MultiTalk Silent Embeds",
"WanVideoImageToVideoSkyreelsv3_audio": "WanVideo Long SkyReelsV3 A2V",
} }
+335 -249
View File
@@ -1,10 +1,8 @@
import os, gc, math import os, gc, math
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
import numpy as np
import hashlib import hashlib
from tqdm import tqdm
from .wanvideo.schedulers import get_scheduler, scheduler_list
from .utils import(log, clip_encode_image_tiled, add_noise_to_reference_video, set_module_tensor_to_device) from .utils import(log, clip_encode_image_tiled, add_noise_to_reference_video, set_module_tensor_to_device)
from .taehv import TAEHV from .taehv import TAEHV
@@ -329,7 +327,7 @@ class WanVideoTextEncode:
try: try:
log.info(f"Moving video model to {offload_device}") log.info(f"Moving video model to {offload_device}")
model_to_offload.model.to(offload_device) model_to_offload.model.to(offload_device)
except: except Exception:
pass pass
encoder = t5["model"] encoder = t5["model"]
@@ -369,10 +367,18 @@ class WanVideoTextEncode:
cast_dtype = encoder.dtype cast_dtype = encoder.dtype
params_to_keep = {'norm', 'pos_embedding', 'token_embedding'} params_to_keep = {'norm', 'pos_embedding', 'token_embedding'}
for name, param in encoder.model.named_parameters(): if hasattr(encoder, 'state_dict'):
model_state_dict = encoder.state_dict
else:
model_state_dict = encoder.model.state_dict()
params_list = list(encoder.model.named_parameters())
pbar = tqdm(params_list, desc="Loading T5 parameters", leave=True)
for name, param in pbar:
dtype_to_use = dtype if any(keyword in name for keyword in params_to_keep) else cast_dtype dtype_to_use = dtype if any(keyword in name for keyword in params_to_keep) else cast_dtype
value = encoder.state_dict[name] if hasattr(encoder, 'state_dict') else encoder.model.state_dict()[name] value = model_state_dict[name]
set_module_tensor_to_device(encoder.model, name, device=device_to, dtype=dtype_to_use, value=value) set_module_tensor_to_device(encoder.model, name, device=device_to, dtype=dtype_to_use, value=value)
del model_state_dict
if hasattr(encoder, 'state_dict'): if hasattr(encoder, 'state_dict'):
del encoder.state_dict del encoder.state_dict
mm.soft_empty_cache() mm.soft_empty_cache()
@@ -496,7 +502,7 @@ class WanVideoTextEncodeSingle:
log.info(f"Moving video model to {offload_device}") log.info(f"Moving video model to {offload_device}")
model_to_offload.model.to(offload_device) model_to_offload.model.to(offload_device)
mm.soft_empty_cache() mm.soft_empty_cache()
except: except Exception:
pass pass
encoder = t5["model"] encoder = t5["model"]
@@ -553,6 +559,9 @@ class WanVideoApplyNAG:
"nag_tau": ("FLOAT", {"default": 2.5, "min": 0.0, "max": 10.0, "step": 0.1}), "nag_tau": ("FLOAT", {"default": 2.5, "min": 0.0, "max": 10.0, "step": 0.1}),
"nag_alpha": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.01}), "nag_alpha": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.01}),
}, },
"optional": {
"inplace": ("BOOLEAN", {"default": True, "tooltip": "If true, modifies tensors in place to save memory. Leads to different numerical results which may change the output slightly."}),
}
} }
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", ) RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", )
@@ -561,7 +570,7 @@ class WanVideoApplyNAG:
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Adds NAG prompt embeds to original prompt embeds: 'https://github.com/ChenDarYen/Normalized-Attention-Guidance'" DESCRIPTION = "Adds NAG prompt embeds to original prompt embeds: 'https://github.com/ChenDarYen/Normalized-Attention-Guidance'"
def process(self, original_text_embeds, nag_text_embeds, nag_scale, nag_tau, nag_alpha): def process(self, original_text_embeds, nag_text_embeds, nag_scale, nag_tau, nag_alpha, inplace=True):
prompt_embeds_dict_copy = original_text_embeds.copy() prompt_embeds_dict_copy = original_text_embeds.copy()
prompt_embeds_dict_copy.update({ prompt_embeds_dict_copy.update({
"nag_prompt_embeds": nag_text_embeds["prompt_embeds"], "nag_prompt_embeds": nag_text_embeds["prompt_embeds"],
@@ -569,6 +578,7 @@ class WanVideoApplyNAG:
"nag_scale": nag_scale, "nag_scale": nag_scale,
"nag_tau": nag_tau, "nag_tau": nag_tau,
"nag_alpha": nag_alpha, "nag_alpha": nag_alpha,
"inplace": inplace,
} }
}) })
return (prompt_embeds_dict_copy,) return (prompt_embeds_dict_copy,)
@@ -765,6 +775,98 @@ class WanVideoAddStandInLatent:
updated = dict(embeds) updated = dict(embeds)
updated["standin_input"] = new_entry updated["standin_input"] = new_entry
return (updated,) return (updated,)
class WanVideoAddBindweaveEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"reference_latents": ("LATENT", {"tooltip": "Reference image to encode"}),
},
"optional": {
"ref_masks": ("MASK", {"tooltip": "Reference mask to encode"}),
"qwenvl_embeds_pos": ("QWENVL_EMBEDS", {"tooltip": "Qwen-VL image embeddings for the reference image"}),
"qwenvl_embeds_neg": ("QWENVL_EMBEDS", {"tooltip": "Qwen-VL image embeddings for the reference image"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "LATENT", "MASK",)
RETURN_NAMES = ("image_embeds", "image_embed_preview", "mask_preview",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, reference_latents, ref_masks=None, qwenvl_embeds_pos=None, qwenvl_embeds_neg=None):
updated = dict(embeds)
image_embeds = embeds["image_embeds"]
max_refs = 4
num_refs = reference_latents["samples"].shape[0]
pad = torch.zeros(image_embeds.shape[0], max_refs-num_refs, image_embeds.shape[2], image_embeds.shape[3], device=image_embeds.device, dtype=image_embeds.dtype)
if num_refs < max_refs:
image_embeds = torch.cat([pad, image_embeds], dim=1)
ref_latents = [ref_latent for ref_latent in reference_latents["samples"]]
image_embeds = torch.cat([*ref_latents, image_embeds], dim=1)
mask = embeds.get("mask", None)
if mask is not None:
mask_pad = torch.zeros(mask.shape[0], max_refs-num_refs, mask.shape[2], mask.shape[3], device=mask.device, dtype=mask.dtype)
if num_refs < max_refs:
mask = torch.cat([mask_pad, mask], dim=1)
if ref_masks is not None:
ref_mask_ = common_upscale(ref_masks.unsqueeze(1), mask.shape[3], mask.shape[2], "nearest", "disabled").movedim(0,1)
ref_mask_ = torch.cat([ref_mask_, torch.zeros(3, ref_mask_.shape[1], ref_mask_.shape[2], ref_mask_.shape[3], device=ref_mask_.device, dtype=ref_mask_.dtype)])
mask = torch.cat([ref_mask_, mask], dim=1)
else:
mask = torch.cat([torch.ones(mask.shape[0], num_refs, mask.shape[2], mask.shape[3], device=mask.device, dtype=mask.dtype), mask], dim=1)
updated["mask"] = mask
clip_embeds = updated.get("clip_context", None)
if clip_embeds is not None:
B, T, C = clip_embeds.shape
target_len = max_refs * 257 # 4 * 257 = 1028
if T < target_len:
pad = torch.zeros(B, target_len - T, C, device=clip_embeds.device, dtype=clip_embeds.dtype)
padded_embeds = torch.cat([clip_embeds, pad], dim=1)
log.info(f"Padded clip embeds from {clip_embeds.shape} to {padded_embeds.shape} for Bindweave")
updated["clip_context"] = padded_embeds
else:
updated["clip_context"] = clip_embeds
updated["image_embeds"] = image_embeds
updated["qwenvl_embeds_pos"] = qwenvl_embeds_pos
updated["qwenvl_embeds_neg"] = qwenvl_embeds_neg
return (updated, {"samples": image_embeds.unsqueeze(0)}, mask[0].float())
class TextImageEncodeQwenVL():
@classmethod
def INPUT_TYPES(s):
return {"required": {
"clip": ("CLIP",),
"prompt": ("STRING", {"default": "", "multiline": True}),
},
"optional": {
"image": ("IMAGE", ),
}
}
RETURN_TYPES = ("QWENVL_EMBEDS",)
RETURN_NAMES = ("qwenvl_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(cls, clip, prompt, image=None):
if image is None:
input_images = []
llama_template = None
else:
input_images = [image[:, :, :, :3]]
llama_template = "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>{}<|im_end|>\n<|im_start|>assistant\n"
tokens = clip.tokenize(prompt, images=input_images, llama_template=llama_template)
conditioning = clip.encode_from_tokens_scheduled(tokens)
print("Qwen-VL embeds shape:", conditioning[0][0].shape)
return (conditioning[0][0],)
class WanVideoAddMTVMotion: class WanVideoAddMTVMotion:
@classmethod @classmethod
@@ -799,6 +901,86 @@ class WanVideoAddMTVMotion:
updated["mtv_crafter_motion"] = new_entry updated["mtv_crafter_motion"] = new_entry
return (updated,) return (updated,)
class WanVideoAddStoryMemLatents:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"vae": ("WANVAE",),
"embeds": ("WANVIDIMAGE_EMBEDS",),
"memory_images": ("IMAGE",),
"rope_negative_offset": ("BOOLEAN", {"default": False, "tooltip": "Use positive RoPE frequency offset for the memory latents"}),
"rope_negative_offset_frames": ("INT", {"default": 5, "min": 0, "max": 100, "step": 1, "tooltip": "RoPE frequency offset for the memory latents"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, vae, embeds, memory_images, rope_negative_offset, rope_negative_offset_frames):
updated = dict(embeds)
story_mem_latents, = WanVideoEncodeLatentBatch().encode(vae, memory_images)
updated["story_mem_latents"] = story_mem_latents["samples"].squeeze(2).permute(1, 0, 2, 3) # [C, T, H, W]
updated["rope_negative_offset_frames"] = rope_negative_offset_frames if rope_negative_offset else 0
return (updated,)
class WanVideoSVIProEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"anchor_samples": ("LATENT", {"tooltip": "Initial start image encoded"}),
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}),
},
"optional": {
"prev_samples": ("LATENT", {"tooltip": "Last latent from previous generation"}),
"motion_latent_count": ("INT", {"default": 1, "min": 0, "max": 100, "step": 1, "tooltip": "Number of latents used to continue"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, anchor_samples, num_frames, prev_samples=None, motion_latent_count=1):
anchor_latent = anchor_samples["samples"][0].clone()
C, T, H, W = anchor_latent.shape
total_latents = (num_frames - 1) // 4 + 1
device = anchor_latent.device
dtype = anchor_latent.dtype
if prev_samples is None or motion_latent_count == 0:
padding_size = total_latents - anchor_latent.shape[1]
padding = torch.zeros(C, padding_size, H, W, dtype=dtype, device=device)
y = torch.concat([anchor_latent, padding], dim=1)
else:
prev_latent = prev_samples["samples"][0].clone()
motion_latent = prev_latent[:, -motion_latent_count:]
padding_size = total_latents - anchor_latent.shape[1] - motion_latent.shape[1]
padding = torch.zeros(C, padding_size, H, W, dtype=dtype, device=device)
y = torch.concat([anchor_latent, motion_latent, padding], dim=1)
msk = torch.ones(1, num_frames, H, W, device=device, dtype=dtype)
msk[:, 1:] = 0
msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
msk = msk.view(1, msk.shape[1] // 4, 4, H, W)
msk = msk.transpose(1, 2)[0]
image_embeds = {
"image_embeds": y,
"num_frames": num_frames,
"lat_h": H,
"lat_w": W,
"mask": msk
}
return (image_embeds,)
#region I2V encode #region I2V encode
class WanVideoImageToVideoEncode: class WanVideoImageToVideoEncode:
@classmethod @classmethod
@@ -823,6 +1005,8 @@ class WanVideoImageToVideoEncode:
"extra_latents": ("LATENT", {"tooltip": "Extra latents to add to the input front, used for Skyreels A2 reference images"}), "extra_latents": ("LATENT", {"tooltip": "Extra latents to add to the input front, used for Skyreels A2 reference images"}),
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}), "tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
"add_cond_latents": ("ADD_COND_LATENTS", {"advanced": True, "tooltip": "Additional cond latents WIP"}), "add_cond_latents": ("ADD_COND_LATENTS", {"advanced": True, "tooltip": "Additional cond latents WIP"}),
"augment_empty_frames": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "EXPERIMENTAL: Augment empty frames with the difference to the start image to force more motion"}),
"empty_frame_pad_image": ("IMAGE", {"tooltip": "Use this image to pad empty frames instead of gray, used with SVI-shot and SVI 2.0 LoRAs"}),
} }
} }
@@ -832,18 +1016,14 @@ class WanVideoImageToVideoEncode:
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
def process(self, width, height, num_frames, force_offload, noise_aug_strength, def process(self, width, height, num_frames, force_offload, noise_aug_strength,
start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False, start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False,
temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None, vae=None): temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None, vae=None, augment_empty_frames=0.0, empty_frame_pad_image=None):
if start_image is None and end_image is None and add_cond_latents is None:
return WanVideoEmptyEmbeds().process(
num_frames, width, height, control_embeds=control_embeds, extra_latents=extra_latents,
)
if vae is None: if vae is None:
raise ValueError("VAE is required for image encoding.") raise ValueError("VAE is required for image encoding.")
H = height H = height
W = width W = width
lat_h = H // vae.upsampling_factor lat_h = H // vae.upsampling_factor
lat_w = W // vae.upsampling_factor lat_w = W // vae.upsampling_factor
@@ -868,6 +1048,8 @@ class WanVideoImageToVideoEncode:
mask = torch.cat([mask, torch.zeros(base_frames - mask.shape[0], lat_h, lat_w, device=device)]) mask = torch.cat([mask, torch.zeros(base_frames - mask.shape[0], lat_h, lat_w, device=device)])
mask = mask.unsqueeze(0).to(device, vae.dtype) mask = mask.unsqueeze(0).to(device, vae.dtype)
pixel_mask = mask.clone()
# Repeat first frame and optionally end frame # Repeat first frame and optionally end frame
start_mask_repeated = torch.repeat_interleave(mask[:, 0:1], repeats=4, dim=1) # T, C, H, W start_mask_repeated = torch.repeat_interleave(mask[:, 0:1], repeats=4, dim=1) # T, C, H, W
if end_image is not None and not fun_or_fl2v_model: if end_image is not None and not fun_or_fl2v_model:
@@ -890,7 +1072,7 @@ class WanVideoImageToVideoEncode:
resized_start_image = resized_start_image * 2 - 1 resized_start_image = resized_start_image * 2 - 1
if noise_aug_strength > 0.0: if noise_aug_strength > 0.0:
resized_start_image = add_noise_to_reference_video(resized_start_image, ratio=noise_aug_strength) resized_start_image = add_noise_to_reference_video(resized_start_image, ratio=noise_aug_strength)
if end_image is not None: if end_image is not None:
end_image = end_image[..., :3] end_image = end_image[..., :3]
if end_image.shape[1] != H or end_image.shape[2] != W: if end_image.shape[1] != H or end_image.shape[2] != W:
@@ -900,30 +1082,46 @@ class WanVideoImageToVideoEncode:
resized_end_image = resized_end_image * 2 - 1 resized_end_image = resized_end_image * 2 - 1
if noise_aug_strength > 0.0: if noise_aug_strength > 0.0:
resized_end_image = add_noise_to_reference_video(resized_end_image, ratio=noise_aug_strength) resized_end_image = add_noise_to_reference_video(resized_end_image, ratio=noise_aug_strength)
# Concatenate image with zero frames and encode # Concatenate image with zero frames and encode
if temporal_mask is None: if start_image is not None and end_image is None:
if start_image is not None and end_image is None: zero_frames = torch.zeros(3, num_frames-start_image.shape[0], H, W, device=device, dtype=vae.dtype)
zero_frames = torch.zeros(3, num_frames-start_image.shape[0], H, W, device=device, dtype=vae.dtype) concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames], dim=1)
concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames], dim=1) del resized_start_image, zero_frames
del resized_start_image, zero_frames elif start_image is None and end_image is not None:
elif start_image is None and end_image is not None: zero_frames = torch.zeros(3, num_frames-end_image.shape[0], H, W, device=device, dtype=vae.dtype)
zero_frames = torch.zeros(3, num_frames-end_image.shape[0], H, W, device=device, dtype=vae.dtype) concatenated = torch.cat([zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1)
concatenated = torch.cat([zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1) del zero_frames
del zero_frames elif start_image is None and end_image is None:
elif start_image is None and end_image is None: concatenated = torch.zeros(3, num_frames, H, W, device=device, dtype=vae.dtype)
concatenated = torch.zeros(3, num_frames, H, W, device=device, dtype=vae.dtype)
else:
if fun_or_fl2v_model:
zero_frames = torch.zeros(3, num_frames-(start_image.shape[0]+end_image.shape[0]), H, W, device=device, dtype=vae.dtype)
else:
zero_frames = torch.zeros(3, num_frames-1, H, W, device=device, dtype=vae.dtype)
concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1)
del resized_start_image, zero_frames
else: else:
temporal_mask = common_upscale(temporal_mask.unsqueeze(1), W, H, "nearest", "disabled").squeeze(1) if fun_or_fl2v_model:
concatenated = resized_start_image[:,:num_frames].to(vae.dtype)# * temporal_mask[:num_frames].unsqueeze(0).to(vae.dtype) zero_frames = torch.zeros(3, num_frames-(start_image.shape[0]+end_image.shape[0]), H, W, device=device, dtype=vae.dtype)
del resized_start_image, temporal_mask else:
zero_frames = torch.zeros(3, num_frames-1, H, W, device=device, dtype=vae.dtype)
concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1)
del resized_start_image, zero_frames
if empty_frame_pad_image is not None:
pad_img = empty_frame_pad_image.clone()[..., :3]
if pad_img.shape[1] != H or pad_img.shape[2] != W:
pad_img = common_upscale(pad_img.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1)
pad_img = (pad_img.movedim(-1, 0) * 2 - 1).to(device, dtype=vae.dtype)
num_pad_frames = pad_img.shape[1]
num_target_frames = concatenated.shape[1]
if num_pad_frames < num_target_frames:
pad_img = torch.cat([pad_img, pad_img[:, -1:].expand(-1, num_target_frames - num_pad_frames, -1, -1)], dim=1)
else:
pad_img = pad_img[:, :num_target_frames]
frame_is_empty = (pixel_mask[0].mean(dim=(-2, -1)) < 0.5)[:concatenated.shape[1]].clone()
if start_image is not None:
frame_is_empty[:start_image.shape[0]] = False
if end_image is not None:
frame_is_empty[-end_image.shape[0]:] = False
concatenated[:, frame_is_empty] = pad_img[:, frame_is_empty]
mm.soft_empty_cache() mm.soft_empty_cache()
gc.collect() gc.collect()
@@ -941,6 +1139,9 @@ class WanVideoImageToVideoEncode:
has_ref = True has_ref = True
y[:, :1] *= start_latent_strength y[:, :1] *= start_latent_strength
y[:, -1:] *= end_latent_strength y[:, -1:] *= end_latent_strength
if augment_empty_frames > 0.0:
frame_is_empty = (mask[0].mean(dim=(-2, -1)) < 0.5).view(1, -1, 1, 1)
y = y[:, :1] + (y - y[:, :1]) * ((augment_empty_frames+1) * frame_is_empty + ~frame_is_empty)
# Calculate maximum sequence length # Calculate maximum sequence length
patches_per_frame = lat_h * lat_w // (PATCH_SIZE[1] * PATCH_SIZE[2]) patches_per_frame = lat_h * lat_w // (PATCH_SIZE[1] * PATCH_SIZE[2])
@@ -949,14 +1150,14 @@ class WanVideoImageToVideoEncode:
if add_cond_latents is not None: if add_cond_latents is not None:
add_cond_latents["ref_latent_neg"] = vae.encode(torch.zeros(1, 3, 1, H, W, device=device, dtype=vae.dtype), device) add_cond_latents["ref_latent_neg"] = vae.encode(torch.zeros(1, 3, 1, H, W, device=device, dtype=vae.dtype), device)
if force_offload: if force_offload:
vae.model.to(offload_device) vae.model.to(offload_device)
mm.soft_empty_cache() mm.soft_empty_cache()
gc.collect() gc.collect()
image_embeds = { image_embeds = {
"image_embeds": y, "image_embeds": y.cpu(),
"clip_context": clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None, "clip_context": clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None,
"negative_clip_context": clip_embeds.get("negative_clip_embeds", None) if clip_embeds is not None else None, "negative_clip_context": clip_embeds.get("negative_clip_embeds", None) if clip_embeds is not None else None,
"max_seq_len": max_seq_len, "max_seq_len": max_seq_len,
@@ -968,11 +1169,11 @@ class WanVideoImageToVideoEncode:
"fun_or_fl2v_model": fun_or_fl2v_model, "fun_or_fl2v_model": fun_or_fl2v_model,
"has_ref": has_ref, "has_ref": has_ref,
"add_cond_latents": add_cond_latents, "add_cond_latents": add_cond_latents,
"mask": mask "mask": mask.cpu()
} }
return (image_embeds,) return (image_embeds,)
# region WanAnimate # region WanAnimate
class WanVideoAnimateEmbeds: class WanVideoAnimateEmbeds:
@classmethod @classmethod
@@ -983,15 +1184,15 @@ class WanVideoAnimateEmbeds:
"height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}), "height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}),
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}), "num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}),
"force_offload": ("BOOLEAN", {"default": True}), "force_offload": ("BOOLEAN", {"default": True}),
"frame_window_size": ("INT", {"default": 77, "min": 1, "max": 1000, "step": 1, "tooltip": "Number of frames to use for temporal attention window"}), "frame_window_size": ("INT", {"default": 77, "min": 1, "max": 10000, "step": 1, "tooltip": "Number of frames to use for temporal attention window"}),
"colormatch": ( "colormatch": (
[ [
'disabled', 'disabled',
'mkl', 'mkl',
'hm', 'hm',
'reinhard', 'reinhard',
'mvgd', 'mvgd',
'hm-mvgd-hm', 'hm-mvgd-hm',
'hm-mkl-hm', 'hm-mkl-hm',
], { ], {
"default": 'disabled', "tooltip": "Color matching method to use between the windows" "default": 'disabled', "tooltip": "Color matching method to use between the windows"
@@ -1006,6 +1207,7 @@ class WanVideoAnimateEmbeds:
"face_images": ("IMAGE", {"tooltip": "end frame"}), "face_images": ("IMAGE", {"tooltip": "end frame"}),
"bg_images": ("IMAGE", {"tooltip": "background images"}), "bg_images": ("IMAGE", {"tooltip": "background images"}),
"mask": ("MASK", {"tooltip": "mask"}), "mask": ("MASK", {"tooltip": "mask"}),
"start_ref_image": ("IMAGE", {"tooltip": "start ref image"}),
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}), "tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
} }
} }
@@ -1016,7 +1218,7 @@ class WanVideoAnimateEmbeds:
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
def process(self, vae, width, height, num_frames, force_offload, frame_window_size, colormatch, pose_strength, face_strength, def process(self, vae, width, height, num_frames, force_offload, frame_window_size, colormatch, pose_strength, face_strength,
ref_images=None, pose_images=None, face_images=None, clip_embeds=None, tiled_vae=False, bg_images=None, mask=None): ref_images=None, pose_images=None, face_images=None, clip_embeds=None, tiled_vae=False, bg_images=None, mask=None, start_ref_image=None):
W = (width // 16) * 16 W = (width // 16) * 16
H = (height // 16) * 16 H = (height // 16) * 16
@@ -1027,7 +1229,7 @@ class WanVideoAnimateEmbeds:
num_refs = ref_images.shape[0] if ref_images is not None else 0 num_refs = ref_images.shape[0] if ref_images is not None else 0
num_frames = ((num_frames - 1) // 4) * 4 + 1 num_frames = ((num_frames - 1) // 4) * 4 + 1
looping = num_frames > frame_window_size looping = num_frames > frame_window_size or start_ref_image is not None
if num_frames < frame_window_size: if num_frames < frame_window_size:
frame_window_size = num_frames frame_window_size = num_frames
@@ -1125,6 +1327,12 @@ class WanVideoAnimateEmbeds:
resized_face_images = (resized_face_images * 2 - 1).unsqueeze(0) resized_face_images = (resized_face_images * 2 - 1).unsqueeze(0)
resized_face_images = resized_face_images.to(offload_device, dtype=vae.dtype) resized_face_images = resized_face_images.to(offload_device, dtype=vae.dtype)
if start_ref_image is not None:
if start_ref_image.shape[1] != H or start_ref_image.shape[2] != W:
resized_start_ref_image = common_upscale(start_ref_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
else:
resized_start_ref_image = start_ref_image.permute(3, 0, 1, 2) # C, T, H, W
resized_start_ref_image = resized_start_ref_image[:3] * 2 - 1
seq_len = math.ceil((target_shape[2] * target_shape[3]) / 4 * target_shape[1]) seq_len = math.ceil((target_shape[2] * target_shape[3]) / 4 * target_shape[1])
@@ -1144,6 +1352,7 @@ class WanVideoAnimateEmbeds:
"is_masked": mask is not None, "is_masked": mask is not None,
"ref_latent": ref_latent, "ref_latent": ref_latent,
"ref_image": resized_ref_images if ref_images is not None else None, "ref_image": resized_ref_images if ref_images is not None else None,
"start_ref_image": resized_start_ref_image if start_ref_image is not None else None,
"face_pixels": resized_face_images if face_images is not None else None, "face_pixels": resized_face_images if face_images is not None else None,
"num_frames": num_frames, "num_frames": num_frames,
"target_shape": target_shape, "target_shape": target_shape,
@@ -1718,33 +1927,7 @@ class WanVideoContextOptions:
} }
return (context_options,) return (context_options,)
class WanVideoFlowEdit:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"source_embeds": ("WANVIDEOTEXTEMBEDS", ),
"skip_steps": ("INT", {"default": 4, "min": 0}),
"drift_steps": ("INT", {"default": 0, "min": 0}),
"drift_flow_shift": ("FLOAT", {"default": 3.0, "min": 1.0, "max": 30.0, "step": 0.01}),
"source_cfg": ("FLOAT", {"default": 6.0, "min": 0.0, "max": 30.0, "step": 0.01}),
"drift_cfg": ("FLOAT", {"default": 6.0, "min": 0.0, "max": 30.0, "step": 0.01}),
},
"optional": {
"source_image_embeds": ("WANVIDIMAGE_EMBEDS", ),
}
}
RETURN_TYPES = ("FLOWEDITARGS", )
RETURN_NAMES = ("flowedit_args",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Flowedit options for WanVideo"
def process(self, **kwargs):
return (kwargs,)
class WanVideoLoopArgs: class WanVideoLoopArgs:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -1816,155 +1999,6 @@ class WanVideoFreeInitArgs:
def process(self, **kwargs): def process(self, **kwargs):
return (kwargs,) return (kwargs,)
class WanVideoScheduler: #WIP
@classmethod
def INPUT_TYPES(s):
return {"required": {
"scheduler": (scheduler_list, {"default": "unipc"}),
"steps": ("INT", {"default": 30, "min": 1, "tooltip": "Number of steps for the scheduler"}),
"shift": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 1000.0, "step": 0.01}),
"start_step": ("INT", {"default": 0, "min": 0, "tooltip": "Starting step for the scheduler"}),
"end_step": ("INT", {"default": -1, "min": -1, "tooltip": "Ending step for the scheduler"})
},
"optional": {
"sigmas": ("SIGMAS", ),
},
"hidden": {
"unique_id": "UNIQUE_ID",
},
}
RETURN_TYPES = ("SIGMAS", "INT", "FLOAT", scheduler_list, "INT", "INT",)
RETURN_NAMES = ("sigmas", "steps", "shift", "scheduler", "start_step", "end_step")
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
def process(self, scheduler, steps, start_step, end_step, shift, unique_id, sigmas=None):
sample_scheduler, timesteps, start_idx, end_idx = get_scheduler(
scheduler,
steps,
start_step, end_step, shift,
device,
sigmas=sigmas,
log_timesteps=True)
scheduler_dict = {
"sample_scheduler": sample_scheduler,
"timesteps": timesteps,
}
try:
from server import PromptServer
import io
import base64
import matplotlib.pyplot as plt
except:
PromptServer = None
if unique_id and PromptServer is not None:
try:
# Plot sigmas and save to a buffer
sigmas_np = sample_scheduler.full_sigmas.cpu().numpy()
if not np.isclose(sigmas_np[-1], 0.0, atol=1e-6):
sigmas_np = np.append(sigmas_np, 0.0)
buf = io.BytesIO()
fig = plt.figure(facecolor='#353535')
ax = fig.add_subplot(111)
ax.set_facecolor('#353535') # Set axes background color
x_values = range(0, len(sigmas_np))
ax.plot(x_values, sigmas_np)
# Annotate each sigma value
ax.scatter(x_values, sigmas_np, color='white', s=20, zorder=3) # Small dots at each sigma
for x, y in zip(x_values, sigmas_np):
# Show all annotations if few steps, or just show split step annotations
show_annotation = len(sigmas_np) <= 10
is_split_step = (start_idx > 0 and x == start_idx) or (end_idx != -1 and x == end_idx + 1)
if show_annotation or is_split_step:
color = 'orange'
if is_split_step:
color = 'yellow'
ax.annotate(f"{y:.3f}", (x, y), textcoords="offset points", xytext=(10, 1), ha='center', color=color, fontsize=12)
ax.set_xticks(x_values)
ax.set_title("Sigmas", color='white') # Title font color
ax.set_xlabel("Step", color='white') # X label font color
ax.set_ylabel("Sigma Value", color='white') # Y label font color
ax.tick_params(axis='x', colors='white', labelsize=10) # X tick color
ax.tick_params(axis='y', colors='white', labelsize=10) # Y tick color
# Add split point if end_step is defined
end_idx += 1
if end_idx != -1 and 0 <= end_idx < len(sigmas_np) - 1:
ax.axvline(end_idx, color='red', linestyle='--', linewidth=2, label='end_step split')
# Add split point if start_step is defined
if start_idx > 0 and 0 <= start_idx < len(sigmas_np):
ax.axvline(start_idx, color='green', linestyle='--', linewidth=2, label='start_step split')
if (end_idx != -1 and 0 <= end_idx < len(sigmas_np)) or (start_idx > 0 and 0 <= start_idx < len(sigmas_np)):
handles, labels = ax.get_legend_handles_labels()
if labels:
ax.legend()
if start_idx < end_idx and 0 <= start_idx < len(sigmas_np) and 0 < end_idx < len(sigmas_np):
ax.axvspan(start_idx, end_idx, color='lightblue', alpha=0.1, label='Sampled Range')
plt.tight_layout()
plt.savefig(buf, format='png')
plt.close(fig)
buf.seek(0)
img_base64 = base64.b64encode(buf.read()).decode('utf-8')
buf.close()
# Send as HTML img tag with base64 data
html_img = f"<img src='data:image/png;base64,{img_base64}' alt='Sigmas Plot' style='max-width:100%; height:100%; overflow:hidden; display:block;'>"
PromptServer.instance.send_progress_text(html_img, unique_id)
except Exception as e:
print("Failed to send sigmas plot:", e)
pass
return (sigmas, steps, shift, scheduler_dict, start_step, end_step)
class WanVideoSchedulerSA_ODE:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"use_adaptive_order": ("BOOLEAN", {"default": False, "tooltip": "Use adaptive order"}),
"use_velocity_smoothing": ("BOOLEAN", {"default": True, "tooltip": "Use velocity smoothing"}),
"convergence_threshold": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.001, "tooltip": "Convergence threshold for velocity smoothing"}),
"smoothing_factor": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 1.0, "step": 0.001, "tooltip": "Smoothing factor for velocity smoothing"}),
"steps": ("INT", {"default": 30, "min": 1, "tooltip": "Number of steps for the scheduler"}),
"shift": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 1000.0, "step": 0.01}),
"start_step": ("INT", {"default": 0, "min": 0, "tooltip": "Starting step for the scheduler"}),
"end_step": ("INT", {"default": -1, "min": -1, "tooltip": "Ending step for the scheduler"})
},
"optional": {
"sigmas": ("SIGMAS", ),
},
}
RETURN_TYPES = ("SIGMAS", "INT", "FLOAT", scheduler_list, "INT", "INT",)
RETURN_NAMES = ("sigmas", "steps", "shift", "scheduler", "start_step", "end_step")
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
def process(self, steps, start_step, end_step, shift, use_adaptive_order, use_velocity_smoothing, convergence_threshold, smoothing_factor, sigmas=None):
sample_scheduler, timesteps, _, _ = get_scheduler(
scheduler="sa_ode_stable/lowstep",
steps=steps,
start_step=start_step, end_step=end_step, shift=shift,
device=device,
sigmas=sigmas,
log_timesteps=True,
use_adaptive_order=use_adaptive_order,
use_velocity_smoothing=use_velocity_smoothing,
convergence_threshold=convergence_threshold,
smoothing_factor=smoothing_factor
)
scheduler_dict = {
"sample_scheduler": sample_scheduler,
"timesteps": timesteps,
}
return (sigmas, steps, shift, scheduler_dict, start_step, end_step)
rope_functions = ["default", "comfy", "comfy_chunked"] rope_functions = ["default", "comfy", "comfy_chunked"]
class WanVideoRoPEFunction: class WanVideoRoPEFunction:
@@ -1995,6 +2029,53 @@ class WanVideoRoPEFunction:
return (rope_func_dict,) return (rope_func_dict,)
return (rope_function,) return (rope_function,)
#region TTM
class WanVideoAddTTMLatents:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"reference_latents": ("LATENT", {"tooltip": "Latents used as reference for TTM"}),
"mask": ("MASK", {"tooltip": "Mask used for TTM"}),
"start_step": ("INT", {"default": 0, "min": -1, "max": 1000, "step": 1, "tooltip": "Start step for whole denoising process"}),
"end_step": ("INT", {"default": 1, "min": 1, "max": 1000, "step": 1, "tooltip": "The step to stop applying TTM"}),
},
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", )
RETURN_NAMES = ("image_embeds", )
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "https://github.com/time-to-move/TTM"
def add(self, embeds, reference_latents, mask, start_step, end_step):
if end_step < max(0, start_step):
raise ValueError(f"`end_step` ({end_step}) must be >= `start_step` ({start_step}).")
mask_sampled = mask[::4]
mask_sampled = mask_sampled.unsqueeze(1).unsqueeze(0) # [1, T, 1, H, W]
vae_upscale_factor = 8
if reference_latents["samples"].shape[1] == 48:
vae_upscale_factor = 16
# Upsample spatially to latent resolution
H_latent = mask_sampled.shape[-2] // vae_upscale_factor
W_latent = mask_sampled.shape[-1] // vae_upscale_factor
mask_latent = F.interpolate(
mask_sampled.float(),
size=(mask_sampled.shape[2], H_latent, W_latent),
mode="nearest"
)
updated = dict(embeds)
updated["ttm_reference_latents"] = reference_latents["samples"].squeeze(0)
updated["ttm_mask"] = mask_latent.squeeze(0).movedim(1, 0) # [T, 1, H, W]
updated["ttm_start_step"] = start_step
updated["ttm_end_step"] = end_step
return (updated,)
#region VideoDecode #region VideoDecode
class WanVideoDecode: class WanVideoDecode:
@@ -2016,7 +2097,7 @@ class WanVideoDecode:
"tile_stride_y": ("INT", {"default": 128, "min": 32, "max": 2040, "step": 8, "tooltip": "Tile stride height in pixels. Smaller values use less VRAM but will introduce more seams."}), "tile_stride_y": ("INT", {"default": 128, "min": 32, "max": 2040, "step": 8, "tooltip": "Tile stride height in pixels. Smaller values use less VRAM but will introduce more seams."}),
}, },
"optional": { "optional": {
"normalization": (["default", "minmax"], {"advanced": True}), "normalization": (["default", "minmax", "none"], {"advanced": True}),
} }
} }
@@ -2040,7 +2121,7 @@ class WanVideoDecode:
video.clamp_(-1.0, 1.0) video.clamp_(-1.0, 1.0)
video.add_(1.0).div_(2.0) video.add_(1.0).div_(2.0)
return video.cpu().float(), return video.cpu().float(),
latents = samples["samples"] latents = samples["samples"].clone()
end_image = samples.get("end_image", None) end_image = samples.get("end_image", None)
has_ref = samples.get("has_ref", False) has_ref = samples.get("has_ref", False)
drop_last = samples.get("drop_last", False) drop_last = samples.get("drop_last", False)
@@ -2059,25 +2140,24 @@ class WanVideoDecode:
if drop_last: if drop_last:
latents = latents[:, :, :-1] latents = latents[:, :, :-1]
if type(vae).__name__ == "TAEHV": if type(vae).__name__ == "TAEHV":
images = vae.decode_video(latents.permute(0, 2, 1, 3, 4), cond=flashvsr_LQ_images.to(vae.dtype) if flashvsr_LQ_images is not None else None)[0].permute(1, 0, 2, 3) images = vae.decode_video(latents.permute(0, 2, 1, 3, 4), cond=flashvsr_LQ_images.to(vae.dtype) if flashvsr_LQ_images is not None else None)[0].permute(1, 0, 2, 3)
images = torch.clamp(images, 0.0, 1.0) images = torch.clamp(images, 0.0, 1.0)
images = images.permute(1, 2, 3, 0).cpu().float() images = images.permute(1, 2, 3, 0).cpu().float()
return (images,) return (images,)
else: else:
if end_image is not None:
enable_vae_tiling = False
images = vae.decode(latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//8, tile_y//8), tile_stride=(tile_stride_x//8, tile_stride_y//8))[0] images = vae.decode(latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//8, tile_y//8), tile_stride=(tile_stride_x//8, tile_stride_y//8))[0]
images = images.cpu().float() images = images.cpu().float()
if normalization == "minmax": if normalization != "none":
images.sub_(images.min()).div_(images.max() - images.min()) if normalization == "minmax":
else: images.sub_(images.min()).div_(images.max() - images.min())
images.clamp_(-1.0, 1.0) else:
images.add_(1.0).div_(2.0) images.clamp_(-1.0, 1.0)
images.add_(1.0).div_(2.0)
if is_looped: if is_looped:
temp_latents = torch.cat([latents[:, :, -3:]] + [latents[:, :, :2]], dim=2) temp_latents = torch.cat([latents[:, :, -3:]] + [latents[:, :, :2]], dim=2)
temp_images = vae.decode(temp_latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))[0] temp_images = vae.decode(temp_latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))[0]
@@ -2088,7 +2168,7 @@ class WanVideoDecode:
if end_image is not None: if end_image is not None:
images = images[:, 0:-1] images = images[:, 0:-1]
vae.to(offload_device) vae.to(offload_device)
mm.soft_empty_cache() mm.soft_empty_cache()
@@ -2117,7 +2197,7 @@ class WanVideoEncodeLatentBatch:
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Encodes a batch of images individually to create a latent video batch where each video is a single frame, useful for I2V init purposes, for example as multiple context window inits" DESCRIPTION = "Encodes a batch of images individually to create a latent video batch where each video is a single frame, useful for I2V init purposes, for example as multiple context window inits"
def encode(self, vae, images, enable_vae_tiling, tile_x, tile_y, tile_stride_x, tile_stride_y, latent_strength=1.0): def encode(self, vae, images, enable_vae_tiling=False, tile_x=272, tile_y=272, tile_stride_x=144, tile_stride_y=128, latent_strength=1.0):
vae.to(device) vae.to(device)
images = images.clone() images = images.clone()
@@ -2139,7 +2219,7 @@ class WanVideoEncodeLatentBatch:
latent = vae.encode(img.unsqueeze(0).unsqueeze(0).permute(0, 4, 1, 2, 3), device=device, tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor)) latent = vae.encode(img.unsqueeze(0).unsqueeze(0).permute(0, 4, 1, 2, 3), device=device, tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))
else: else:
latent = vae.encode(img.unsqueeze(0).unsqueeze(0).permute(0, 4, 1, 2, 3), device=device, tiled=enable_vae_tiling) latent = vae.encode(img.unsqueeze(0).unsqueeze(0).permute(0, 4, 1, 2, 3), device=device, tiled=enable_vae_tiling)
if latent_strength != 1.0: if latent_strength != 1.0:
latent *= latent_strength latent *= latent_strength
latent_list.append(latent.squeeze(0).cpu()) latent_list.append(latent.squeeze(0).cpu())
@@ -2199,14 +2279,16 @@ class WanVideoEncode:
latents = latents.permute(0, 2, 1, 3, 4) latents = latents.permute(0, 2, 1, 3, 4)
else: else:
latents = vae.encode(image * 2.0 - 1.0, device=device, tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor)) latents = vae.encode(image * 2.0 - 1.0, device=device, tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))
vae.to(offload_device) vae.to(offload_device)
if latent_strength != 1.0: if latent_strength != 1.0:
latents *= latent_strength latents *= latent_strength
latents = latents.cpu()
log.info(f"WanVideoEncode: Encoded latents shape {latents.shape}") log.info(f"WanVideoEncode: Encoded latents shape {latents.shape}")
mm.soft_empty_cache() mm.soft_empty_cache()
return ({"samples": latents, "noise_mask": mask},) return ({"samples": latents, "noise_mask": mask},)
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
@@ -2221,7 +2303,6 @@ NODE_CLASS_MAPPINGS = {
"WanVideoEnhanceAVideo": WanVideoEnhanceAVideo, "WanVideoEnhanceAVideo": WanVideoEnhanceAVideo,
"WanVideoContextOptions": WanVideoContextOptions, "WanVideoContextOptions": WanVideoContextOptions,
"WanVideoTextEmbedBridge": WanVideoTextEmbedBridge, "WanVideoTextEmbedBridge": WanVideoTextEmbedBridge,
"WanVideoFlowEdit": WanVideoFlowEdit,
"WanVideoControlEmbeds": WanVideoControlEmbeds, "WanVideoControlEmbeds": WanVideoControlEmbeds,
"WanVideoSLG": WanVideoSLG, "WanVideoSLG": WanVideoSLG,
"WanVideoLoopArgs": WanVideoLoopArgs, "WanVideoLoopArgs": WanVideoLoopArgs,
@@ -2237,7 +2318,6 @@ NODE_CLASS_MAPPINGS = {
"WanVideoBlockList": WanVideoBlockList, "WanVideoBlockList": WanVideoBlockList,
"WanVideoTextEncodeCached": WanVideoTextEncodeCached, "WanVideoTextEncodeCached": WanVideoTextEncodeCached,
"WanVideoAddExtraLatent": WanVideoAddExtraLatent, "WanVideoAddExtraLatent": WanVideoAddExtraLatent,
"WanVideoScheduler": WanVideoScheduler,
"WanVideoAddStandInLatent": WanVideoAddStandInLatent, "WanVideoAddStandInLatent": WanVideoAddStandInLatent,
"WanVideoAddControlEmbeds": WanVideoAddControlEmbeds, "WanVideoAddControlEmbeds": WanVideoAddControlEmbeds,
"WanVideoAddMTVMotion": WanVideoAddMTVMotion, "WanVideoAddMTVMotion": WanVideoAddMTVMotion,
@@ -2245,8 +2325,12 @@ NODE_CLASS_MAPPINGS = {
"WanVideoAddPusaNoise": WanVideoAddPusaNoise, "WanVideoAddPusaNoise": WanVideoAddPusaNoise,
"WanVideoAnimateEmbeds": WanVideoAnimateEmbeds, "WanVideoAnimateEmbeds": WanVideoAnimateEmbeds,
"WanVideoAddLucyEditLatents": WanVideoAddLucyEditLatents, "WanVideoAddLucyEditLatents": WanVideoAddLucyEditLatents,
"WanVideoSchedulerSA_ODE": WanVideoSchedulerSA_ODE, "WanVideoAddBindweaveEmbeds": WanVideoAddBindweaveEmbeds,
"TextImageEncodeQwenVL": TextImageEncodeQwenVL,
"WanVideoUniLumosEmbeds": WanVideoUniLumosEmbeds, "WanVideoUniLumosEmbeds": WanVideoUniLumosEmbeds,
"WanVideoAddTTMLatents": WanVideoAddTTMLatents,
"WanVideoAddStoryMemLatents": WanVideoAddStoryMemLatents,
"WanVideoSVIProEmbeds": WanVideoSVIProEmbeds,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
@@ -2262,7 +2346,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoEnhanceAVideo": "WanVideo Enhance-A-Video", "WanVideoEnhanceAVideo": "WanVideo Enhance-A-Video",
"WanVideoContextOptions": "WanVideo Context Options", "WanVideoContextOptions": "WanVideo Context Options",
"WanVideoTextEmbedBridge": "WanVideo TextEmbed Bridge", "WanVideoTextEmbedBridge": "WanVideo TextEmbed Bridge",
"WanVideoFlowEdit": "WanVideo FlowEdit",
"WanVideoControlEmbeds": "WanVideo Control Embeds", "WanVideoControlEmbeds": "WanVideo Control Embeds",
"WanVideoSLG": "WanVideo SLG", "WanVideoSLG": "WanVideo SLG",
"WanVideoLoopArgs": "WanVideo Loop Args", "WanVideoLoopArgs": "WanVideo Loop Args",
@@ -2285,6 +2368,9 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddPusaNoise": "WanVideo Add Pusa Noise", "WanVideoAddPusaNoise": "WanVideo Add Pusa Noise",
"WanVideoAnimateEmbeds": "WanVideo Animate Embeds", "WanVideoAnimateEmbeds": "WanVideo Animate Embeds",
"WanVideoAddLucyEditLatents": "WanVideo Add LucyEdit Latents", "WanVideoAddLucyEditLatents": "WanVideo Add LucyEdit Latents",
"WanVideoSchedulerSA_ODE": "WanVideo Scheduler SA-ODE", "WanVideoAddBindweaveEmbeds": "WanVideo Add Bindweave Embeds",
"WanVideoUniLumosEmbeds": "WanVideo UniLumos Embeds", "WanVideoUniLumosEmbeds": "WanVideo UniLumos Embeds",
"WanVideoAddTTMLatents": "WanVideo Add TTMLatents",
"WanVideoAddStoryMemLatents": "WanVideo Add StoryMem Latents",
"WanVideoSVIProEmbeds": "WanVideo SVIPro Embeds",
} }
+390 -170
View File
File diff suppressed because it is too large Load Diff
+735 -994
View File
File diff suppressed because it is too large Load Diff
+116 -92
View File
@@ -1,4 +1,5 @@
import torch import torch
import torch.nn.functional as F
import numpy as np import numpy as np
from comfy.utils import common_upscale from comfy.utils import common_upscale
from comfy import model_management from comfy import model_management
@@ -8,7 +9,7 @@ from einops import rearrange
try: try:
from server import PromptServer from server import PromptServer
except: except Exception:
PromptServer = None PromptServer = None
VAE_STRIDE = (4, 8, 8) VAE_STRIDE = (4, 8, 8)
@@ -35,7 +36,7 @@ class WanVideoImageResizeToClosest:
DESCRIPTION = "Resizes image to the closest supported resolution based on aspect ratio and max pixels, according to the original code" DESCRIPTION = "Resizes image to the closest supported resolution based on aspect ratio and max pixels, according to the original code"
def process(self, image, generation_width, generation_height, aspect_ratio_preservation ): def process(self, image, generation_width, generation_height, aspect_ratio_preservation ):
H, W = image.shape[1], image.shape[2] H, W = image.shape[1], image.shape[2]
max_area = generation_width * generation_height max_area = generation_width * generation_height
@@ -47,7 +48,7 @@ class WanVideoImageResizeToClosest:
aspect_ratio = generation_height / generation_width aspect_ratio = generation_height / generation_width
if aspect_ratio_preservation == "crop_to_new": if aspect_ratio_preservation == "crop_to_new":
crop = "center" crop = "center"
lat_h = round( lat_h = round(
np.sqrt(max_area * aspect_ratio) // VAE_STRIDE[1] // np.sqrt(max_area * aspect_ratio) // VAE_STRIDE[1] //
PATCH_SIZE[1] * PATCH_SIZE[1]) PATCH_SIZE[1] * PATCH_SIZE[1])
@@ -135,27 +136,27 @@ class WanVideoVACEStartToEndFrame:
# Convert negative end_index to positive # Convert negative end_index to positive
if end_index < 0: if end_index < 0:
end_index = num_frames + end_index end_index = num_frames + end_index
# Create output batch with empty frames # Create output batch with empty frames
out_batch = torch.ones((num_frames, H, W, 3), device=device) * empty_frame_level out_batch = torch.ones((num_frames, H, W, 3), device=device) * empty_frame_level
# Create mask tensor with proper dimensions # Create mask tensor with proper dimensions
masks = torch.ones((num_frames, H, W), device=device) masks = torch.ones((num_frames, H, W), device=device)
# Pre-process all images at once to avoid redundant work # Pre-process all images at once to avoid redundant work
if end_image is not None and (end_image.shape[1] != H or end_image.shape[2] != W): if end_image is not None and (end_image.shape[1] != H or end_image.shape[2] != W):
end_image = common_upscale(end_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1) end_image = common_upscale(end_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1)
if control_images is not None and (control_images.shape[1] != H or control_images.shape[2] != W): if control_images is not None and (control_images.shape[1] != H or control_images.shape[2] != W):
control_images = common_upscale(control_images.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1) control_images = common_upscale(control_images.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1)
# Place start image at start_index # Place start image at start_index
if start_image is not None: if start_image is not None:
frames_to_copy = min(start_image.shape[0], num_frames - start_index) frames_to_copy = min(start_image.shape[0], num_frames - start_index)
if frames_to_copy > 0: if frames_to_copy > 0:
out_batch[start_index:start_index + frames_to_copy] = start_image[:frames_to_copy] out_batch[start_index:start_index + frames_to_copy] = start_image[:frames_to_copy]
masks[start_index:start_index + frames_to_copy] = 0 masks[start_index:start_index + frames_to_copy] = 0
# Place end image at end_index # Place end image at end_index
if end_image is not None: if end_image is not None:
# Calculate where to start placing end images # Calculate where to start placing end images
@@ -163,28 +164,28 @@ class WanVideoVACEStartToEndFrame:
if end_start < 0: # Handle case where end images won't all fit if end_start < 0: # Handle case where end images won't all fit
end_image = end_image[abs(end_start):] end_image = end_image[abs(end_start):]
end_start = 0 end_start = 0
frames_to_copy = min(end_image.shape[0], num_frames - end_start) frames_to_copy = min(end_image.shape[0], num_frames - end_start)
if frames_to_copy > 0: if frames_to_copy > 0:
out_batch[end_start:end_start + frames_to_copy] = end_image[:frames_to_copy] out_batch[end_start:end_start + frames_to_copy] = end_image[:frames_to_copy]
masks[end_start:end_start + frames_to_copy] = 0 masks[end_start:end_start + frames_to_copy] = 0
# Apply control images to remaining frames that don't have start or end images # Apply control images to remaining frames that don't have start or end images
if control_images is not None: if control_images is not None:
# Create a mask of frames that are still empty (mask == 1) # Create a mask of frames that are still empty (mask == 1)
empty_frames = masks.sum(dim=(1, 2)) > 0.5 * H * W empty_frames = masks.sum(dim=(1, 2)) > 0.5 * H * W
if empty_frames.any(): if empty_frames.any():
# Only apply control images where they exist # Only apply control images where they exist
control_length = control_images.shape[0] control_length = control_images.shape[0]
for frame_idx in range(num_frames): for frame_idx in range(num_frames):
if empty_frames[frame_idx] and frame_idx < control_length: if empty_frames[frame_idx] and frame_idx < control_length:
out_batch[frame_idx] = control_images[frame_idx] out_batch[frame_idx] = control_images[frame_idx]
# Apply inpaint mask if provided # Apply inpaint mask if provided
if inpaint_mask is not None: if inpaint_mask is not None:
inpaint_mask = common_upscale(inpaint_mask.unsqueeze(1), W, H, "nearest-exact", "disabled").squeeze(1).to(device) inpaint_mask = common_upscale(inpaint_mask.unsqueeze(1), W, H, "nearest-exact", "disabled").squeeze(1).to(device)
# Handle different mask lengths efficiently # Handle different mask lengths efficiently
if inpaint_mask.shape[0] > num_frames: if inpaint_mask.shape[0] > num_frames:
inpaint_mask = inpaint_mask[:num_frames] inpaint_mask = inpaint_mask[:num_frames]
@@ -226,40 +227,40 @@ class CreateCFGScheduleFloatList:
cfg_list = [1.0] * steps cfg_list = [1.0] * steps
start_idx = min(int(steps * start_percent), steps - 1) start_idx = min(int(steps * start_percent), steps - 1)
end_idx = min(int(steps * end_percent), steps - 1) end_idx = min(int(steps * end_percent), steps - 1)
for i in range(start_idx, end_idx + 1): for i in range(start_idx, end_idx + 1):
if i >= steps: if i >= steps:
break break
if end_idx == start_idx: if end_idx == start_idx:
t = 0 t = 0
else: else:
t = (i - start_idx) / (end_idx - start_idx) t = (i - start_idx) / (end_idx - start_idx)
if interpolation == "linear": if interpolation == "linear":
factor = t factor = t
elif interpolation == "ease_in": elif interpolation == "ease_in":
factor = t * t factor = t * t
elif interpolation == "ease_out": elif interpolation == "ease_out":
factor = t * (2 - t) factor = t * (2 - t)
cfg_list[i] = round(cfg_scale_start + factor * (cfg_scale_end - cfg_scale_start), 2) cfg_list[i] = round(cfg_scale_start + factor * (cfg_scale_end - cfg_scale_start), 2)
# If start_percent > 0, always include the first step # If start_percent > 0, always include the first step
if start_percent > 0: if start_percent > 0:
cfg_list[0] = 1.0 cfg_list[0] = 1.0
if unique_id and PromptServer is not None: if unique_id and PromptServer is not None:
try: try:
PromptServer.instance.send_progress_text( PromptServer.instance.send_progress_text(
f"{cfg_list}", f"{cfg_list}",
unique_id unique_id
) )
except: except Exception:
pass pass
return (cfg_list,) return (cfg_list,)
class CreateScheduleFloatList: class CreateScheduleFloatList:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -289,16 +290,16 @@ class CreateScheduleFloatList:
cfg_list = [default_value] * steps cfg_list = [default_value] * steps
start_idx = min(int(steps * start_percent), steps - 1) start_idx = min(int(steps * start_percent), steps - 1)
end_idx = min(int(steps * end_percent), steps - 1) end_idx = min(int(steps * end_percent), steps - 1)
for i in range(start_idx, end_idx + 1): for i in range(start_idx, end_idx + 1):
if i >= steps: if i >= steps:
break break
if end_idx == start_idx: if end_idx == start_idx:
t = 0 t = 0
else: else:
t = (i - start_idx) / (end_idx - start_idx) t = (i - start_idx) / (end_idx - start_idx)
if interpolation == "linear": if interpolation == "linear":
factor = t factor = t
elif interpolation == "ease_in": elif interpolation == "ease_in":
@@ -313,16 +314,16 @@ class CreateScheduleFloatList:
cfg_list[0] = default_value cfg_list[0] = default_value
if unique_id and PromptServer is not None: if unique_id and PromptServer is not None:
try: try:
PromptServer.instance.send_progress_text( PromptServer.instance.send_progress_text(
f"{cfg_list}", f"{cfg_list}",
unique_id unique_id
) )
except: except Exception:
pass pass
return (cfg_list,) return (cfg_list,)
class DummyComfyWanModelObject: class DummyComfyWanModelObject:
@classmethod @classmethod
@@ -348,7 +349,7 @@ class DummyComfyWanModelObject:
return model_sampling return model_sampling
return None return None
return (DummyModel(),) return (DummyModel(),)
class WanVideoLatentReScale: class WanVideoLatentReScale:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -405,7 +406,7 @@ class WanVideoLatentReScale:
samples["samples"] = latents samples["samples"] = latents
return (samples,) return (samples,)
class WanVideoSigmaToStep: class WanVideoSigmaToStep:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -422,7 +423,7 @@ class WanVideoSigmaToStep:
def convert(self, sigma): def convert(self, sigma):
return (sigma,) return (sigma,)
class NormalizeAudioLoudness: class NormalizeAudioLoudness:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -437,11 +438,11 @@ class NormalizeAudioLoudness:
FUNCTION = "normalize" FUNCTION = "normalize"
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
def normalize(self, audio, lufs): def normalize(self, audio, lufs):
audio_input = audio["waveform"] audio_input = audio["waveform"]
sample_rate = audio["sample_rate"] sample_rate = audio["sample_rate"]
if audio_input.dim() == 3: if audio_input.dim() == 3:
audio_input = audio_input.squeeze(0) audio_input = audio_input.squeeze(0)
audio_input_np = audio_input.detach().transpose(0, 1).numpy().astype(np.float32) audio_input_np = audio_input.detach().transpose(0, 1).numpy().astype(np.float32)
audio_input_np = np.ascontiguousarray(audio_input_np) audio_input_np = np.ascontiguousarray(audio_input_np)
normalized_audio = self.loudness_norm(audio_input_np, sr=sample_rate, lufs=lufs) normalized_audio = self.loudness_norm(audio_input_np, sr=sample_rate, lufs=lufs)
@@ -449,11 +450,11 @@ class NormalizeAudioLoudness:
out_audio = {"waveform": torch.from_numpy(normalized_audio).transpose(0, 1).unsqueeze(0).float(), "sample_rate": sample_rate} out_audio = {"waveform": torch.from_numpy(normalized_audio).transpose(0, 1).unsqueeze(0).float(), "sample_rate": sample_rate}
return (out_audio, ) return (out_audio, )
def loudness_norm(self, audio_array, sr=16000, lufs=-23): def loudness_norm(self, audio_array, sr=16000, lufs=-23):
try: try:
import pyloudnorm import pyloudnorm
except: except Exception:
raise ImportError("pyloudnorm package is not installed") raise ImportError("pyloudnorm package is not installed")
meter = pyloudnorm.Meter(sr) meter = pyloudnorm.Meter(sr)
loudness = meter.integrated_loudness(audio_array) loudness = meter.integrated_loudness(audio_array)
@@ -461,7 +462,7 @@ class NormalizeAudioLoudness:
return audio_array return audio_array
normalized_audio = pyloudnorm.normalize.loudness(audio_array, loudness, lufs) normalized_audio = pyloudnorm.normalize.loudness(audio_array, loudness, lufs)
return normalized_audio return normalized_audio
class WanVideoPassImagesFromSamples: class WanVideoPassImagesFromSamples:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -505,15 +506,15 @@ class FaceMaskFromPoseKeypoints:
for i, pose_frame in enumerate(pose_frames): for i, pose_frame in enumerate(pose_frames):
selected_idx, prev_center = self.select_closest_person(pose_frame, person_index if i == 0 else prev_center) selected_idx, prev_center = self.select_closest_person(pose_frame, person_index if i == 0 else prev_center)
np_frames.append(self.draw_kps(pose_frame, selected_idx)) np_frames.append(self.draw_kps(pose_frame, selected_idx))
if not np_frames: if not np_frames:
# Handle case where no frames were processed # Handle case where no frames were processed
log.warning("No valid pose frames found, returning empty mask") log.warning("No valid pose frames found, returning empty mask")
return (torch.zeros((1, 64, 64), dtype=torch.float32),) return (torch.zeros((1, 64, 64), dtype=torch.float32),)
np_frames = np.stack(np_frames, axis=0) np_frames = np.stack(np_frames, axis=0)
tensor = torch.from_numpy(np_frames).float() / 255. tensor = torch.from_numpy(np_frames).float() / 255.
print("tensor.shape:", tensor.shape) log.info(f"tensor.shape: {tensor.shape}")
tensor = tensor[:, :, :, 0] tensor = tensor[:, :, :, 0]
return (tensor,) return (tensor,)
@@ -521,41 +522,41 @@ class FaceMaskFromPoseKeypoints:
people = pose_frame["people"] people = pose_frame["people"]
if not people: if not people:
return -1, None return -1, None
centers = [] centers = []
valid_people_indices = [] valid_people_indices = []
for idx, person in enumerate(people): for idx, person in enumerate(people):
# Check if face keypoints exist and are valid # Check if face keypoints exist and are valid
if "face_keypoints_2d" not in person or not person["face_keypoints_2d"]: if "face_keypoints_2d" not in person or not person["face_keypoints_2d"]:
continue continue
kps = np.array(person["face_keypoints_2d"]) kps = np.array(person["face_keypoints_2d"])
if len(kps) == 0: if len(kps) == 0:
continue continue
n = len(kps) // 3 n = len(kps) // 3
if n == 0: if n == 0:
continue continue
facial_kps = rearrange(kps, "(n c) -> n c", n=n, c=3)[:, :2] facial_kps = rearrange(kps, "(n c) -> n c", n=n, c=3)[:, :2]
# Check if we have valid coordinates (not all zeros) # Check if we have valid coordinates (not all zeros)
if np.all(facial_kps == 0): if np.all(facial_kps == 0):
continue continue
center = facial_kps.mean(axis=0) center = facial_kps.mean(axis=0)
# Check if center is valid (not NaN or infinite) # Check if center is valid (not NaN or infinite)
if np.isnan(center).any() or np.isinf(center).any(): if np.isnan(center).any() or np.isinf(center).any():
continue continue
centers.append(center) centers.append(center)
valid_people_indices.append(idx) valid_people_indices.append(idx)
if not centers: if not centers:
return -1, None return -1, None
if isinstance(prev_center_or_index, (int, np.integer)): if isinstance(prev_center_or_index, (int, np.integer)):
# First frame: use person_index, but map to valid people # First frame: use person_index, but map to valid people
if 0 <= prev_center_or_index < len(valid_people_indices): if 0 <= prev_center_or_index < len(valid_people_indices):
@@ -587,58 +588,58 @@ class FaceMaskFromPoseKeypoints:
width, height = pose_frame["canvas_width"], pose_frame["canvas_height"] width, height = pose_frame["canvas_width"], pose_frame["canvas_height"]
canvas = np.zeros((height, width, 3), dtype=np.uint8) canvas = np.zeros((height, width, 3), dtype=np.uint8)
people = pose_frame["people"] people = pose_frame["people"]
if person_index < 0 or person_index >= len(people): if person_index < 0 or person_index >= len(people):
return canvas # Out of bounds, return blank return canvas # Out of bounds, return blank
person = people[person_index] person = people[person_index]
# Check if face keypoints exist and are valid # Check if face keypoints exist and are valid
if "face_keypoints_2d" not in person or not person["face_keypoints_2d"]: if "face_keypoints_2d" not in person or not person["face_keypoints_2d"]:
return canvas # No face keypoints, return blank return canvas # No face keypoints, return blank
face_kps_data = person["face_keypoints_2d"] face_kps_data = person["face_keypoints_2d"]
if len(face_kps_data) == 0: if len(face_kps_data) == 0:
return canvas # Empty keypoints, return blank return canvas # Empty keypoints, return blank
n = len(face_kps_data) // 3 n = len(face_kps_data) // 3
if n < 17: # Need at least 17 points for outer contour if n < 17: # Need at least 17 points for outer contour
return canvas # Not enough keypoints, return blank return canvas # Not enough keypoints, return blank
facial_kps = rearrange(np.array(face_kps_data), "(n c) -> n c", n=n, c=3)[:, :2] facial_kps = rearrange(np.array(face_kps_data), "(n c) -> n c", n=n, c=3)[:, :2]
# Check if we have valid coordinates (not all zeros) # Check if we have valid coordinates (not all zeros)
if np.all(facial_kps == 0): if np.all(facial_kps == 0):
return canvas # All keypoints are zero, return blank return canvas # All keypoints are zero, return blank
# Check for NaN or infinite values # Check for NaN or infinite values
if np.isnan(facial_kps).any() or np.isinf(facial_kps).any(): if np.isnan(facial_kps).any() or np.isinf(facial_kps).any():
return canvas # Invalid coordinates, return blank return canvas # Invalid coordinates, return blank
# Check for negative coordinates or coordinates that would create streaks # Check for negative coordinates or coordinates that would create streaks
if np.any(facial_kps < 0): if np.any(facial_kps < 0):
return canvas # Negative coordinates, likely bad detection return canvas # Negative coordinates, likely bad detection
# Check if coordinates are reasonable (not too close to edges which might indicate bad detection) # Check if coordinates are reasonable (not too close to edges which might indicate bad detection)
min_margin = 5 # Minimum distance from edges min_margin = 5 # Minimum distance from edges
if (np.any(facial_kps[:, 0] < min_margin) or if (np.any(facial_kps[:, 0] < min_margin) or
np.any(facial_kps[:, 1] < min_margin) or np.any(facial_kps[:, 1] < min_margin) or
np.any(facial_kps[:, 0] > width - min_margin) or np.any(facial_kps[:, 0] > width - min_margin) or
np.any(facial_kps[:, 1] > height - min_margin)): np.any(facial_kps[:, 1] > height - min_margin)):
# Check if this looks like a streak to corner (many points near 0,0) # Check if this looks like a streak to corner (many points near 0,0)
corner_points = np.sum((facial_kps[:, 0] < min_margin) & (facial_kps[:, 1] < min_margin)) corner_points = np.sum((facial_kps[:, 0] < min_margin) & (facial_kps[:, 1] < min_margin))
if corner_points > 3: # Too many points near corner, likely bad detection if corner_points > 3: # Too many points near corner, likely bad detection
return canvas return canvas
facial_kps = facial_kps.astype(np.int32) facial_kps = facial_kps.astype(np.int32)
# Ensure coordinates are within canvas bounds # Ensure coordinates are within canvas bounds
facial_kps[:, 0] = np.clip(facial_kps[:, 0], 0, width - 1) facial_kps[:, 0] = np.clip(facial_kps[:, 0], 0, width - 1)
facial_kps[:, 1] = np.clip(facial_kps[:, 1], 0, height - 1) facial_kps[:, 1] = np.clip(facial_kps[:, 1], 0, height - 1)
part_color = (255, 255, 255) part_color = (255, 255, 255)
outer_contour = facial_kps[:17] outer_contour = facial_kps[:17]
# Additional validation for the contour before drawing # Additional validation for the contour before drawing
# Check if contour points are too spread out (indicating bad detection) # Check if contour points are too spread out (indicating bad detection)
if len(outer_contour) >= 3: if len(outer_contour) >= 3:
@@ -647,11 +648,11 @@ class FaceMaskFromPoseKeypoints:
max_x, max_y = np.max(outer_contour, axis=0) max_x, max_y = np.max(outer_contour, axis=0)
contour_width = max_x - min_x contour_width = max_x - min_x
contour_height = max_y - min_y contour_height = max_y - min_y
# If contour spans more than 80% of canvas, likely bad detection # If contour spans more than 80% of canvas, likely bad detection
if (contour_width > 0.8 * width or contour_height > 0.8 * height): if (contour_width > 0.8 * width or contour_height > 0.8 * height):
return canvas return canvas
# Check if we have a valid contour (at least 3 unique points) # Check if we have a valid contour (at least 3 unique points)
unique_points = np.unique(outer_contour, axis=0) unique_points = np.unique(outer_contour, axis=0)
if len(unique_points) >= 3: if len(unique_points) >= 3:
@@ -659,11 +660,11 @@ class FaceMaskFromPoseKeypoints:
# Calculate area to see if it's too large or too small # Calculate area to see if it's too large or too small
contour_area = cv2.contourArea(outer_contour) contour_area = cv2.contourArea(outer_contour)
canvas_area = width * height canvas_area = width * height
# If contour is less than 0.1% or more than 50% of canvas, skip # If contour is less than 0.1% or more than 50% of canvas, skip
if 0.001 * canvas_area <= contour_area <= 0.5 * canvas_area: if 0.001 * canvas_area <= contour_area <= 0.5 * canvas_area:
cv2.fillPoly(canvas, pts=[outer_contour], color=part_color) cv2.fillPoly(canvas, pts=[outer_contour], color=part_color)
return canvas return canvas
@@ -679,7 +680,7 @@ class DrawGaussianNoiseOnImage:
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
} }
} }
RETURN_TYPES = ("IMAGE", ) RETURN_TYPES = ("IMAGE", )
RETURN_NAMES = ("images",) RETURN_NAMES = ("images",)
FUNCTION = "apply" FUNCTION = "apply"
@@ -694,68 +695,89 @@ class DrawGaussianNoiseOnImage:
in_masks = mask.clone().to(processing_device) in_masks = mask.clone().to(processing_device)
in_images = image.clone().to(processing_device) in_images = image.clone().to(processing_device)
# Resize mask to match image dimensions # Resize mask to match image dimensions
if HM != H or WM != W: if HM != H or WM != W:
in_masks = F.interpolate(mask.unsqueeze(1), size=(H, W), mode='nearest-exact').squeeze(1) in_masks = F.interpolate(mask.unsqueeze(1), size=(H, W), mode='nearest-exact').squeeze(1)
# Match batch sizes # Match batch sizes
if B > BM: if B > BM:
in_masks = in_masks.repeat((B + BM - 1) // BM, 1, 1)[:B] in_masks = in_masks.repeat((B + BM - 1) // BM, 1, 1)[:B]
elif BM > B: elif BM > B:
in_masks = in_masks[:B] in_masks = in_masks[:B]
output_images = [] output_images = []
# Set random seed for reproducibility # Set random seed for reproducibility
generator = torch.Generator(device=processing_device).manual_seed(seed) generator = torch.Generator(device=processing_device).manual_seed(seed)
for i in tqdm(range(B), desc="DrawGaussianNoiseOnImage batch"): for i in tqdm(range(B), desc="DrawGaussianNoiseOnImage batch"):
curr_mask = in_masks[i] curr_mask = in_masks[i]
img_idx = min(i, B - 1) img_idx = min(i, B - 1)
curr_image = in_images[img_idx] curr_image = in_images[img_idx]
# Expand mask to 3 channels # Expand mask to 3 channels
mask_expanded = curr_mask.unsqueeze(-1).expand(-1, -1, 3) mask_expanded = curr_mask.unsqueeze(-1).expand(-1, -1, 3)
# Calculate mean and std per channel from the subject region (where mask is 1) # Calculate mean and std per channel from the subject region (where mask is 1)
subject_mask = mask_expanded > 0.5 subject_mask = mask_expanded > 0.5
# Initialize noise tensor # Initialize noise tensor
noise = torch.zeros_like(curr_image) noise = torch.zeros_like(curr_image)
for c in range(C): for c in range(C):
channel = curr_image[:, :, c] channel = curr_image[:, :, c]
channel_mask = subject_mask[:, :, c] channel_mask = subject_mask[:, :, c]
if channel_mask.sum() > 0: if channel_mask.sum() > 0:
# Get subject pixels # Get subject pixels
subject_pixels = channel[channel_mask] subject_pixels = channel[channel_mask]
# Calculate statistics # Calculate statistics
mean = subject_pixels.mean() mean = subject_pixels.mean()
std = subject_pixels.std() std = subject_pixels.std()
# Generate Gaussian noise for this channel # Generate Gaussian noise for this channel
noise[:, :, c] = torch.normal(mean=mean.item(), std=std.item(), noise[:, :, c] = torch.normal(mean=mean.item(), std=std.item(),
size=(H, W), generator=generator, size=(H, W), generator=generator,
device=processing_device) device=processing_device)
# Clamp noise to valid range # Clamp noise to valid range
noise = torch.clamp(noise, 0.0, 1.0) noise = torch.clamp(noise, 0.0, 1.0)
# Apply: keep subject, fill background with noise # Apply: keep subject, fill background with noise
masked_image = curr_image * mask_expanded + noise * (1 - mask_expanded) masked_image = curr_image * mask_expanded + noise * (1 - mask_expanded)
output_images.append(masked_image) output_images.append(masked_image)
# If no masks were processed, return empty tensor # If no masks were processed, return empty tensor
if not output_images: if not output_images:
return (torch.zeros((0, H, W, 3), dtype=image.dtype),) return (torch.zeros((0, H, W, 3), dtype=image.dtype),)
out_rgb = torch.stack(output_images, dim=0).cpu() out_rgb = torch.stack(output_images, dim=0).cpu()
return (out_rgb, ) return (out_rgb, )
class WanVideoPreviewEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
}
}
RETURN_TYPES = ("LATENT", "MASK")
RETURN_NAMES = ("image_embeds", "mask",)
FUNCTION = "get"
CATEGORY = "WanVideoWrapper"
def get(self, embeds):
latents = embeds.get("image_embeds", None)
mask = embeds.get("mask", None)
if mask is not None:
mask = mask[0].float().cpu()
return ({"samples": latents.unsqueeze(0)}, mask)
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"WanVideoImageResizeToClosest": WanVideoImageResizeToClosest, "WanVideoImageResizeToClosest": WanVideoImageResizeToClosest,
"WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame, "WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame,
@@ -769,6 +791,7 @@ NODE_CLASS_MAPPINGS = {
"WanVideoPassImagesFromSamples": WanVideoPassImagesFromSamples, "WanVideoPassImagesFromSamples": WanVideoPassImagesFromSamples,
"FaceMaskFromPoseKeypoints": FaceMaskFromPoseKeypoints, "FaceMaskFromPoseKeypoints": FaceMaskFromPoseKeypoints,
"DrawGaussianNoiseOnImage": DrawGaussianNoiseOnImage, "DrawGaussianNoiseOnImage": DrawGaussianNoiseOnImage,
"WanVideoPreviewEmbeds": WanVideoPreviewEmbeds,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest", "WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest",
@@ -783,4 +806,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoPassImagesFromSamples": "WanVideo Pass Images From Samples", "WanVideoPassImagesFromSamples": "WanVideo Pass Images From Samples",
"FaceMaskFromPoseKeypoints": "Face Mask From Pose Keypoints", "FaceMaskFromPoseKeypoints": "Face Mask From Pose Keypoints",
"DrawGaussianNoiseOnImage": "Draw Gaussian Noise On Image", "DrawGaussianNoiseOnImage": "Draw Gaussian Noise On Image",
} "WanVideoPreviewEmbeds": "WanVideo Preview Embeds",
}
+440
View File
@@ -0,0 +1,440 @@
import torch
from torch import nn
from torch.nn import functional as F
from einops import rearrange
import numpy as np
from typing import Tuple
from .unet_causal_3d_blocks import get_down_block3d, CausalConv3d
class ControlNetCausalConditioningEmbedding(nn.Module):
def __init__(self, conditioning_embedding_channels: int, conditioning_channels: int = 3, block_out_channels: Tuple[int, ...] = (16, 32, 96, 256)):
super().__init__()
self.conv_in = CausalConv3d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1)
self.blocks = nn.ModuleList([])
for i in range(len(block_out_channels) - 1):
channel_in = block_out_channels[i]
channel_out = block_out_channels[i + 1]
self.blocks.append(nn.Conv2d(channel_in, channel_in, kernel_size=3, padding=1))
self.blocks.append(nn.Conv2d(channel_in, channel_out, kernel_size=3, padding=1, stride=2))
self.conv_out = nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1)
def forward(self, conditioning):
embedding = self.conv_in(conditioning)
embedding = F.silu(embedding)
for block in self.blocks:
embedding = block(embedding)
embedding = F.silu(embedding)
embedding = self.conv_out(embedding)
return embedding
class MiniHunyuanEncoder(nn.Module):
'''
a direct copy of hunyuan encoder
'''
def __init__(
self,
in_channels = 3,
out_channels = 3,
down_block_types = ['DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D'],
block_out_channels = [128, 256, 512, 512],
layers_per_block = 2,
norm_num_groups = 32,
act_fn: str = "silu",
time_compression_ratio: int = 4,
spatial_compression_ratio: int = 8,
):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(
in_channels, block_out_channels[0], kernel_size=3, stride=1)
self.mid_block = None
self.down_blocks = nn.ModuleList([])
# down
output_channel = block_out_channels[0]
for i, down_block_type in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_downsample_layers = int(
np.log2(spatial_compression_ratio))
num_time_downsample_layers = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i >= (
len(block_out_channels) - 1 - num_time_downsample_layers) and not is_final_block)
elif time_compression_ratio == 8:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i < num_time_downsample_layers)
else:
raise ValueError(
f"Unsupported time_compression_ratio: {time_compression_ratio}")
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
downsample_stride_T = (2, ) if add_time_downsample else (1, )
downsample_stride = tuple(
downsample_stride_T + downsample_stride_HW)
down_block = get_down_block3d(
down_block_type,
num_layers=self.layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=bool(
add_spatial_downsample or add_time_downsample),
downsample_stride=downsample_stride,
resnet_eps=1e-6,
downsample_padding=0,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
)
self.down_blocks.append(down_block)
self.conv_out = CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
def forward(self, sample):
assert len(sample.shape) == 5, "The input tensor should have 5 dimensions"
sample = self.conv_in(sample)
# down
for down_block in self.down_blocks:
sample = down_block(sample)
sample = self.conv_out(sample)
return sample
class ControlNetConditioningEmbedding(nn.Module):
"""
Quoting from https://arxiv.org/abs/2302.05543: "Stable Diffusion uses a pre-processing method similar to VQ-GAN
[11] to convert the entire dataset of 512 × 512 images into smaller 64 × 64 “latent images” for stabilized
training. This requires ControlNets to convert image-based conditions to 64 × 64 feature space to match the
convolution size. We use a tiny network E(·) of four convolution layers with 4 × 4 kernels and 2 × 2 strides
(activated by ReLU, channels are 16, 32, 64, 128, initialized with Gaussian weights, trained jointly with the full
model) to encode image-space conditions ... into feature maps ..."
"""
def __init__(
self,
conditioning_embedding_channels: int,
conditioning_channels: int = 3,
block_out_channels: Tuple[int, ...] = (16, 32, 96, 256),
):
super().__init__()
self.conv_in = nn.Conv2d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1)
self.blocks = nn.ModuleList([])
for i in range(len(block_out_channels) - 1):
channel_in = block_out_channels[i]
channel_out = block_out_channels[i + 1]
self.blocks.append(nn.Conv2d(channel_in, channel_in, kernel_size=3, padding=1))
self.blocks.append(nn.Conv2d(channel_in, channel_out, kernel_size=3, padding=1, stride=2))
self.conv_out = nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1)
def forward(self, conditioning):
embedding = self.conv_in(conditioning)
embedding = F.silu(embedding)
for block in self.blocks:
embedding = block(embedding)
embedding = F.silu(embedding)
embedding = self.conv_out(embedding)
return embedding
class InflatedGroupNorm(nn.GroupNorm):
def forward(self, x):
video_length = x.shape[2]
x = rearrange(x, "b c f h w -> (b f) c h w")
x = super().forward(x)
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
return x
class InflatedConv3d(nn.Conv2d):
def forward(self, x):
video_length = x.shape[2]
x = rearrange(x, "b c f h w -> (b f) c h w")
x = super().forward(x)
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
return x
class ResnetBlockInflated(nn.Module):
def __init__(self, *, in_channels, out_channels=None, dropout=0.0, groups=32, groups_out=None, pre_norm=True, eps=1e-6, non_linearity="swish", output_scale_factor=1.0):
super().__init__()
self.pre_norm = pre_norm
self.pre_norm = True
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.output_scale_factor = output_scale_factor
if groups_out is None:
groups_out = groups
self.norm1 = InflatedGroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
self.conv1 = InflatedConv3d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
self.norm2 = InflatedGroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
self.dropout = torch.nn.Dropout(dropout)
self.conv2 = InflatedConv3d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
if non_linearity == "swish":
self.nonlinearity = lambda x: F.silu(x)
elif non_linearity == "silu":
self.nonlinearity = nn.SiLU()
def forward(self, input_tensor, temb):
if temb is not None:
print("Warning: temb is None in ResnetBlockInflated")
hidden_states = input_tensor
hidden_states = self.norm1(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.conv1(hidden_states)
if temb is not None:
hidden_states = hidden_states + temb
hidden_states = self.norm2(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.dropout(hidden_states)
hidden_states = self.conv2(hidden_states)
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
return output_tensor
class DownEncoderBlockInflated(nn.Module):
def __init__(self, *, num_layers: int, in_channels: int, out_channels: int, add_downsample: bool, downsample_stride: tuple = (1, 2, 2),
resnet_eps: float = 1e-6, resnet_act_fn: str = "silu", resnet_groups: int = 32):
super().__init__()
self.resnets = nn.ModuleList([ResnetBlockInflated(
in_channels=in_channels if i == 0 else out_channels,
out_channels=out_channels,
eps=resnet_eps,
non_linearity=resnet_act_fn,
groups=resnet_groups,
) for i in range(num_layers)])
self.downsamplers = nn.ModuleList()
if add_downsample:
self.downsamplers.append(
InflatedConv3d(
out_channels,
out_channels,
kernel_size=3,
stride=2,
padding=1,
)
)
self.down_stride = downsample_stride
else:
self.down_stride = (1, 1, 1)
def forward(self, x, temb=None):
for resnet in self.resnets:
x = resnet(x, temb)
for down in self.downsamplers:
x = down(x)
return x
class SFT(nn.Module): # 2D SFT
def __init__(
self, in_channels, out_channels, intermediate_channels=128, groups=32, eps=1e-6):
super().__init__()
self.out_channels = out_channels
self.norm = InflatedGroupNorm(groups, out_channels, eps, affine=True)
self.mlp_shared = nn.Sequential(InflatedConv3d(in_channels, intermediate_channels, kernel_size=3, stride=1, padding=1), nn.SiLU())
self.mlp_gamma = InflatedConv3d(intermediate_channels, out_channels, kernel_size=3, stride=1, padding=1)
self.mlp_beta = InflatedConv3d(intermediate_channels, out_channels, kernel_size=3, stride=1, padding=1)
def forward(self, hidden_state, condition):
"""
hidden_state : (B, Cout, T, H, W)
condition : (B, Cin, 1, H, W)
"""
hidden_state = self.norm(hidden_state) #2D SFT 2D Norm
actv = self.mlp_shared(condition)
gamma = self.mlp_gamma(actv)
beta = self.mlp_beta(actv)
return torch.addcmul(beta, hidden_state, 1 + gamma)
class MiniEncoder2D(nn.Module):
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
down_block_types: list = (
"DownEncoderBlockInflated",
"DownEncoderBlockInflated",
"DownEncoderBlockInflated",
"DownEncoderBlockInflated",
),
block_out_channels: list = (128, 256, 512, 512),
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
spatial_compression_ratio: int = 8,
):
super().__init__()
# -------------------------------------------------------------------
# conv in
# -------------------------------------------------------------------
self.conv_in = InflatedConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1, padding=1)
self.down_blocks = nn.ModuleList()
output_channel = block_out_channels[0]
num_spatial_down_layers = int(np.log2(spatial_compression_ratio))
for i, block_type in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
# is_final_block = i == len(block_out_channels) - 1
add_spatial_downsample = bool(i < num_spatial_down_layers)
downsample_stride = (1, 2, 2) if add_spatial_downsample else (1, 1, 1)
down_block = DownEncoderBlockInflated(
num_layers=layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=add_spatial_downsample,
downsample_stride=downsample_stride,
resnet_eps=1e-6,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
)
self.down_blocks.append(down_block)
self.conv_out = InflatedConv3d(output_channel, out_channels, kernel_size=3, stride=1, padding=1)
def forward(self, x):
# (B,C,1,H,W)
x = self.conv_in(x)
for block in self.down_blocks:
x = block(x)
return self.conv_out(x)
class Driven_Ref_PoseEncoder(nn.Module):
def __init__(
self, in_channels = 3, out_channels = 3,
down_block_types = ['DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D'],
block_out_channels = [128, 256, 512, 512], layers_per_block = 2, norm_num_groups = 32,
act_fn: str = "silu", time_compression_ratio: int = 4, spatial_compression_ratio: int = 8,
):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1)
self.mid_block = None
self.down_blocks = nn.ModuleList([])
# down
output_channel = block_out_channels[0]
for i, down_block_type in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_downsample_layers = int(
np.log2(spatial_compression_ratio))
num_time_downsample_layers = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i >= (
len(block_out_channels) - 1 - num_time_downsample_layers) and not is_final_block)
elif time_compression_ratio == 8:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i < num_time_downsample_layers)
else:
raise ValueError(
f"Unsupported time_compression_ratio: {time_compression_ratio}")
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
downsample_stride_T = (2, ) if add_time_downsample else (1, )
downsample_stride = tuple(
downsample_stride_T + downsample_stride_HW)
down_block = get_down_block3d(
down_block_type,
num_layers=self.layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=bool(
add_spatial_downsample or add_time_downsample),
downsample_stride=downsample_stride,
resnet_eps=1e-6,
downsample_padding=0,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
attention_head_dim=output_channel,
)
self.down_blocks.append(down_block)
self.conv_out = CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
self.ref_pose_encoder = MiniEncoder2D(
in_channels = in_channels,
out_channels = out_channels,
block_out_channels = block_out_channels,
norm_num_groups = norm_num_groups,
layers_per_block = layers_per_block,
spatial_compression_ratio = spatial_compression_ratio,
)
self.sft_layers = nn.ModuleList()
for i, ch in enumerate(block_out_channels):
if i == 0: # 0 层 (H/2,W/2) 不做 SFT
self.sft_layers.append(None)
else: # H/4、H/8、H/16 做 SFT
self.sft_layers.append(
SFT(
in_channels=ch,
out_channels=ch,
intermediate_channels=max(8, ch // 2),
groups=norm_num_groups,
)
)
def forward(self, driven_pose, ref_pose):
# driven_pose b c t h w
# ref_pose b c 1 h w
ref_pose_cond, ref_feats = self.ref_pose_encoder(ref_pose)
x = self.conv_in(driven_pose)
for i, down_block in enumerate(self.down_blocks):
x = down_block(x)
if self.sft_layers[i] is not None:
cond_feat = ref_feats[i]
x = self.sft_layers[i](x, cond_feat)
driven_pose_cond = self.conv_out(x)
return driven_pose_cond, ref_pose_cond
+152
View File
@@ -0,0 +1,152 @@
import torch
from ..utils import log
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoAddOneToAllReferenceEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"vae": ("WANVAE", {"tooltip": "VAE model"}),
"ref_image": ("IMAGE",),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
},
"optional": {
"ref_mask": ("MASK",),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, vae, ref_image, strength, start_percent, end_percent, ref_mask=None):
updated = dict(embeds)
ref_latent = ref_latent_empty = None
vae.to(device)
ref_image_in = (ref_image[..., :3].permute(3, 0, 1, 2) * 2 - 1).to(device, vae.dtype)
ref_latent = vae.encode([ref_image_in], device, tiled=False)
ref_mask_in = None
if ref_mask is not None:
ref_mask_in = (ref_mask.unsqueeze(0).repeat(3, 1, 1, 1) * 2 - 1.).to(device, vae.dtype)
else:
ref_mask_in = torch.zeros_like(ref_image_in)-1
ref_mask_latent = vae.encode([ref_mask_in], device, tiled=False)
if ref_mask is not None and not torch.all(ref_mask == 0):
ref_latent_empty = vae.encode([torch.zeros_like(ref_image_in)-1], device, tiled=False)
else:
ref_latent_empty = ref_mask_latent
vae.to(offload_device)
updated.setdefault("one_to_all_embeds", {})
updated["one_to_all_embeds"]["ref_latent_pos"] = torch.cat([ref_latent, ref_latent_empty], dim=1)
updated["one_to_all_embeds"]["ref_latent_neg"] = torch.cat([ref_latent_empty, ref_latent_empty], dim=1)
updated["one_to_all_embeds"]["ref_strength"] = strength
updated["one_to_all_embeds"]["ref_start_percent"] = start_percent
updated["one_to_all_embeds"]["ref_end_percent"] = end_percent
return (updated,)
class WanVideoAddOneToAllPoseEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"pose_images": ("IMAGE", {"tooltip": "Pose images for the entire video"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the pose control"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the pose control application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the pose control application"}),
},
"optional": {
"pose_prefix_image": ("IMAGE",),
"pose_cfg_scale": ("FLOAT", {"default": 1.5, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "CFG scale for the pose control, has no effect if main cfg scale is 1.0"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, pose_images, strength, pose_prefix_image=None, start_percent=0.0, end_percent=1.0, pose_cfg_scale=1.5):
updated = dict(embeds)
updated.setdefault("one_to_all_embeds", {})
pose_images_in = pose_images[..., :3].unsqueeze(0).permute(0, 4, 1, 2, 3) * 2 - 1 # 1 B H W C -> B C 1 H W
updated["one_to_all_embeds"]["pose_images"] = pose_images_in
if pose_prefix_image is not None:
updated["one_to_all_embeds"]["pose_prefix_image"] = pose_prefix_image.unsqueeze(0).permute(0, 4, 1, 2, 3) * 2 - 1 # 1 B H W C -> B C 1 H W
else:
updated["one_to_all_embeds"]["pose_prefix_image"] = pose_images_in[:, :, :1]
updated["one_to_all_embeds"]["controlnet_strength"] = strength
updated["one_to_all_embeds"]["controlnet_start_percent"] = start_percent
updated["one_to_all_embeds"]["controlnet_end_percent"] = end_percent
updated["one_to_all_embeds"]["pose_cfg_scale"] = pose_cfg_scale
return (updated,)
class WanVideoAddOneToAllExtendEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"prev_latents": ("LATENT", {"tooltip": "Previous latents to be used to continue generation"}),
"window_size": ("INT", {"default": 81, "min": 1, "max": 256, "step": 1, "tooltip": "Number of new frames to generate" }),
"overlap": ("INT", {"default": 5, "min": 0, "max": 64, "step": 1, "tooltip": "Number of overlapping frames between previous and new frames" }),
"frames_processed": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1, "tooltip": "Number of frames already processed in the video" }),
"if_not_enough_frames": (["pad_with_last", "error"], {"default": "pad_with_last", "tooltip": "What to do if there are not enough frames in pose_images for the window"}),
},
"optional": {
"pose_images": ("IMAGE", {"tooltip": "Pose images for the entire video"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "IMAGE",)
RETURN_NAMES = ("image_embeds", "pose_slice",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, prev_latents, if_not_enough_frames, window_size=81, overlap=5, frames_processed=0, pose_images=None):
updated = dict(embeds)
updated.setdefault("one_to_all_embeds", {})
updated["one_to_all_embeds"]["prev_latents"] = prev_latents["samples"][0]
if pose_images is not None:
pose_images_in = pose_images.clone()[..., :3]
start = max(0, frames_processed - overlap)
end = start + window_size
log.info(f"Extracting pose images from {start} to {end}")
if start >= pose_images_in.shape[0]:
raise ValueError(f"start index {start} exceeds pose images length {pose_images_in.shape[0]}")
if end > pose_images_in.shape[0]:
if if_not_enough_frames == "pad_with_last":
padding_needed = end - pose_images_in.shape[0]
pose_images_in = torch.cat([pose_images_in, pose_images_in[-1:].repeat(padding_needed, 1, 1, 1)], dim=0)
log.info(f"Not enough frames, padding with {padding_needed} frames to reach {end} total frames")
else:
raise ValueError(f"end index {end} exceeds pose images length {pose_images.shape[0]}")
pose_slice = pose_images_in[start:end]
else:
pose_slice = torch.zeros((1, 64, 64, 3))
return (updated, pose_slice)
NODE_CLASS_MAPPINGS = {
"WanVideoAddOneToAllReferenceEmbeds": WanVideoAddOneToAllReferenceEmbeds,
"WanVideoAddOneToAllPoseEmbeds": WanVideoAddOneToAllPoseEmbeds,
"WanVideoAddOneToAllExtendEmbeds": WanVideoAddOneToAllExtendEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddOneToAllReferenceEmbeds": "WanVideo Add OneToAll Reference Embeds",
"WanVideoAddOneToAllPoseEmbeds": "WanVideo Add OneToAll Pose Embeds",
"WanVideoAddOneToAllExtendEmbeds": "WanVideo Add OneToAll Extend Embeds",
}
+220
View File
@@ -0,0 +1,220 @@
from typing import Dict, Union
import torch
import torch.nn as nn
from ..wanvideo.modules.model import WanLayerNorm, WanSelfAttention, EmbedND_RifleX, sinusoidal_embedding_1d, apply_rotary_emb_split, apply_rope_comfy1
class WanAttentionBlock(nn.Module):
def __init__(self, in_features, out_features, ffn_dim, ffn2_dim, num_heads, qk_norm=True, cross_attn_norm=False, eps=1e-6, attention_mode="sdpa", rope_func="comfy", rms_norm_function="default"):
super().__init__()
self.dim = out_features
self.ffn_dim = ffn_dim
self.num_heads = num_heads
self.head_dim = out_features // num_heads
self.qk_norm = qk_norm
self.cross_attn_norm = cross_attn_norm
self.eps = eps
self.attention_mode = attention_mode
self.rope_func = rope_func
# layers
self.norm1 = WanLayerNorm(self.dim, eps)
self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm, eps, self.attention_mode, rms_norm_function=rms_norm_function, head_norm=False)
self.norm2 = WanLayerNorm(self.dim, eps)
self.ffn = nn.Sequential(nn.Linear(in_features, ffn_dim), nn.GELU(approximate='tanh'), nn.Linear(ffn2_dim, out_features))
self.modulation = nn.Parameter(torch.randn(1, 6, out_features) / in_features**0.5)
def get_mod(self, e, modulation):
if e.dim() == 3:
if e.shape[-1] == 512:
e = self.modulation(e)
return e.unsqueeze(2).chunk(6, dim=-1)
return (modulation + e).chunk(6, dim=1) # 1, 6, dim
elif e.dim() == 4:
e_mod = modulation.unsqueeze(2) + e
return [ei.squeeze(1) for ei in e_mod.unbind(dim=1)]
def modulate(self, norm_x, shift_msa, scale_msa):
return torch.addcmul(shift_msa, norm_x, 1 + scale_msa)
def ffn_chunked(self, mod_x, num_chunks=4):
seq_len = mod_x.shape[1]
if seq_len <= 8192 or num_chunks <= 1:
return self.ffn(mod_x)
return torch.cat([self.ffn(chunk.contiguous()) for chunk in mod_x.chunk(num_chunks, dim=1)], dim=1)
#region attention forward
def forward(self, x, e, seq_lens, freqs, split_rope=True, e_tr=None, tr_start=0, tr_num=0):
use_token_replace = False
if e_tr is not None and tr_num > 0:
tr_shift_msa, tr_scale_msa, tr_gate_msa, tr_shift_mlp, tr_scale_mlp, tr_gate_mlp = self.get_mod(e_tr.to(x.device), self.modulation)
use_token_replace = True
tr_start = tr_start or 0
tr_end = tr_start + (tr_num or 0)
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device), self.modulation)
del e
input_dtype = x.dtype
if use_token_replace:
norm_x = self.norm1(x.to(shift_msa.dtype))
input_x = torch.cat([
torch.addcmul(shift_msa, norm_x[:, :tr_start], 1 + scale_msa), # before replace → T
torch.addcmul(tr_shift_msa, norm_x[:, tr_start:tr_end], 1 + tr_scale_msa), # replace segment → t=0
torch.addcmul(shift_msa, norm_x[:, tr_end:], 1 + scale_msa) # after replace → T
], dim=1).to(input_dtype)
else:
input_x = self.modulate(self.norm1(x.to(shift_msa.dtype)), shift_msa, scale_msa).to(input_dtype)
del shift_msa, scale_msa
b, s, n, d = *x.shape[:2], self.self_attn.num_heads, self.self_attn.head_dim
h_dim = w_dim = 2 * (self.head_dim // 6)
t_dim = self.head_dim - h_dim - w_dim
q = self.self_attn.norm_q(self.self_attn.q(input_x)).to(self.self_attn.norm_q.weight.dtype).view(b, s, n, d)
if split_rope:
q = apply_rotary_emb_split(q, freqs, t_dim) # Apply split rotary embedding (only to H/W dimensions, leaving T unchanged)
else:
q = apply_rope_comfy1(q, freqs)
k = self.self_attn.norm_k(self.self_attn.k(input_x).to(self.self_attn.norm_k.weight.dtype)).to(input_x.dtype).view(b, s, n, d)
if split_rope:
k = apply_rotary_emb_split(k, freqs, t_dim)
else:
k = apply_rope_comfy1(k, freqs)
v = self.self_attn.v(input_x).view(b, s, n, d)
del input_x
y = self.self_attn.forward(q, k, v, seq_lens)
del q, k, v
if use_token_replace:
x = x + torch.cat([
y[:, :tr_start] * gate_msa,
y[:, tr_start:tr_end] * tr_gate_msa,
y[:, tr_end:] * gate_msa
], dim=1).to(input_dtype)
else:
x = x.addcmul(y, gate_msa)
del y, gate_msa
# ffn
if use_token_replace:
norm2_x = self.norm2(x.to(shift_mlp.dtype))
mod_x = torch.cat([
torch.addcmul(shift_mlp, norm2_x[:, :tr_start], 1 + scale_mlp),
torch.addcmul(tr_shift_mlp, norm2_x[:, tr_start:tr_end], 1 + tr_scale_mlp),
torch.addcmul(shift_mlp, norm2_x[:, tr_end:], 1 + scale_mlp)
], dim=1)
else:
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
del shift_mlp, scale_mlp
x_ffn = self.ffn_chunked(mod_x.to(input_dtype), num_chunks=1)
del mod_x
# gate_mlp
if use_token_replace:
x = x + torch.cat([
x_ffn[:, :tr_start] * gate_mlp,
x_ffn[:, tr_start:tr_end] * tr_gate_mlp,
x_ffn[:, tr_end:] * gate_mlp
], dim=1).to(input_dtype)
else:
x = x.addcmul(x_ffn.to(gate_mlp.dtype), gate_mlp).to(input_dtype)
del gate_mlp
return x
class WanRefextractor(nn.Module):
def __init__(self, patch_size=(1, 2, 2), in_dim=16, dim=5120, in_features=5120, out_features=5120, ffn_dim=8192, ffn2_dim=8192,
freq_dim=256, num_heads=16, num_layers=32, eps=1e-6,
qk_norm=True, cross_attn_norm=True,
attention_mode='sdpa', rope_func='comfy', rms_norm_function='default',
main_device=torch.device('cuda'), offload_device=torch.device('cpu'), dtype=torch.float16):
super().__init__()
self.patch_size = patch_size
self.freq_dim = freq_dim
self.dim = dim
self.main_device = main_device
self.base_dtype = dtype
self.attention_mode = attention_mode
self.patch_embedding = nn.Conv3d(in_dim, dim, kernel_size=patch_size, stride=patch_size)
self.time_embedding = nn.Sequential(nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
self.time_projection = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6))
self.blocks = nn.ModuleList([
WanAttentionBlock(in_features, out_features, ffn_dim, ffn2_dim, num_heads,
qk_norm, cross_attn_norm, eps, attention_mode="sdpa", rope_func=rope_func, rms_norm_function=rms_norm_function)
for i in range(num_layers)
])
self.ref_blocks = nn.ModuleList([])
for _ in range(len(self.blocks)+1):
self.ref_blocks.append(nn.Linear(in_features, out_features))
d = dim // num_heads
self.rope_embedder = EmbedND_RifleX(d,10000.0, [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)], num_frames=1, k=0)
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None):
patch_size = self.patch_size
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
h_len = ((h + (patch_size[1] // 2)) // patch_size[1])
w_len = ((w + (patch_size[2] // 2)) // patch_size[2])
if steps_t is None:
steps_t = t_len
if steps_h is None:
steps_h = h_len
if steps_w is None:
steps_w = w_len
img_ids = torch.zeros((steps_t, steps_h, steps_w, 3), device=device, dtype=dtype)
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1)
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(freq_offset, h_len - 1, steps=steps_h, device=device, dtype=dtype).reshape(1, -1, 1)
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(freq_offset, w_len - 1, steps=steps_w, device=device, dtype=dtype).reshape(1, 1, -1)
img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
freqs = self.rope_embedder(img_ids, ntk_alphas).movedim(1, 2)
return freqs
def forward(
self,
x: torch.Tensor,
timestep: torch.LongTensor,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
B, C, F, H, W = x.shape
freqs = self.rope_encode_comfy(F, H, W, device=x.device, dtype=x.dtype)
self.patch_embedding.to(self.main_device)
x = self.patch_embedding(x.float()).to(x.dtype).flatten(2).transpose(1, 2).to(self.base_dtype)
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.int32)
time_embed_dtype = self.time_embedding[0].weight.dtype
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
time_embed_dtype = self.base_dtype
e = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, timestep.flatten()).to(time_embed_dtype)) # b, dim
e0 = self.time_projection(e).unflatten(1, (6, self.dim)).to(self.base_dtype) # b, 6, dim
del e
# 4. Transformer blocks
block_samples = ()
for block in self.blocks:
block_samples = block_samples + (x, )
x = block(x, e0, seq_lens, freqs)
block_samples = block_samples + (x, )
ref_block_samples = ()
for block_sample, ref_block in zip(block_samples, self.ref_blocks):
block_sample = ref_block(block_sample)
ref_block_samples = ref_block_samples + (block_sample, )
return ref_block_samples, freqs
+144
View File
@@ -0,0 +1,144 @@
# Copyright 2024 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
#
# Modified from diffusers==0.29.2
#
# ==============================================================================
from typing import Optional
import torch
import torch.nn.functional as F
from torch import nn
import comfy.ops
ops = comfy.ops.disable_weight_init
def prepare_causal_attention_mask(n_frame: int, n_hw: int, dtype, device, batch_size: int = None):
seq_len = n_frame * n_hw
mask = torch.full((seq_len, seq_len), float(
"-inf"), dtype=dtype, device=device)
for i in range(seq_len):
i_frame = i // n_hw
mask[i, : (i_frame + 1) * n_hw] = 0
if batch_size is not None:
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
return mask
class CausalConv3d(nn.Module):
def __init__(self, chan_in, chan_out, kernel_size, stride = 1, dilation = 1, pad_mode='replicate', **kwargs):
super().__init__()
self.pad_mode = pad_mode
padding = (kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size - 1, 0) # W, H, T
self.time_causal_padding = padding
self.conv = ops.Conv3d(chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs)
def forward(self, x):
x = F.pad(x, self.time_causal_padding, mode=self.pad_mode)
return self.conv(x)
class DownsampleCausal3D(nn.Module):
def __init__(self, channels, use_conv=False, out_channels=None, padding=1, name="conv", kernel_size=3, bias=True, stride=2):
super().__init__()
self.channels, self.out_channels, self.use_conv, self.padding, self.name = channels, out_channels or channels, use_conv, padding, name
self.conv = CausalConv3d(self.channels, self.out_channels, kernel_size=kernel_size, stride=stride, bias=bias)
def forward(self, x, scale=1.0):
return self.conv(x)
class ResnetBlockCausal3D(nn.Module):
def __init__(self, *, in_channels: int, out_channels: Optional[int] = None, groups: int = 32, eps: float = 1e-6, conv_3d_out_channels: Optional[int] = None):
super().__init__()
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
self.norm2 = torch.nn.GroupNorm(num_groups=groups, num_channels=out_channels, eps=eps, affine=True)
self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3, stride=1)
conv_3d_out_channels = conv_3d_out_channels or out_channels
self.conv2 = CausalConv3d(out_channels, conv_3d_out_channels, kernel_size=3, stride=1)
def forward(self, input_tensor: torch.FloatTensor, temb: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
hidden_states = input_tensor
hidden_states = self.conv1(nn.SiLU()(self.norm1(hidden_states)))
if temb is not None:
hidden_states = hidden_states + temb
hidden_states = self.conv2(nn.SiLU()(self.norm2(hidden_states)))
return input_tensor + hidden_states
def get_down_block3d(down_block_type: str, num_layers: int, in_channels: int, out_channels: int,
add_downsample: bool, downsample_stride: int, resnet_eps: float, resnet_act_fn: str, resnet_groups: Optional[int] = None,
downsample_padding: Optional[int] = None, **kwargs):
down_block_type = down_block_type[7:] if down_block_type.startswith(
"UNetRes") else down_block_type
if down_block_type == "DownEncoderBlockCausal3D":
return DownEncoderBlockCausal3D(
num_layers=num_layers,
in_channels=in_channels,
out_channels=out_channels,
add_downsample=add_downsample,
downsample_stride=downsample_stride,
resnet_eps=resnet_eps,
resnet_act_fn=resnet_act_fn,
resnet_groups=resnet_groups,
downsample_padding=downsample_padding,
)
raise ValueError(f"{down_block_type} does not exist.")
class DownEncoderBlockCausal3D(nn.Module):
def __init__(self, in_channels: int, out_channels: int, num_layers: int = 1, resnet_eps: float = 1e-6,
resnet_groups: int = 32, add_downsample: bool = True, downsample_stride: int = 2, downsample_padding: int = 1, **kwargs):
super().__init__()
resnets = []
for i in range(num_layers):
in_channels = in_channels if i == 0 else out_channels
resnets.append(
ResnetBlockCausal3D(
in_channels=in_channels,
out_channels=out_channels,
eps=resnet_eps,
groups=resnet_groups,
)
)
self.resnets = nn.ModuleList(resnets)
if add_downsample:
self.downsamplers = nn.ModuleList([DownsampleCausal3D(
out_channels,
use_conv=True,
out_channels=out_channels,
padding=downsample_padding,
name="op",
stride=downsample_stride,
)])
else:
self.downsamplers = None
def forward(self, hidden_states: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
for resnet in self.resnets:
hidden_states = resnet(hidden_states, temb=None, scale=scale)
if self.downsamplers is not None:
for downsampler in self.downsamplers:
hidden_states = downsampler(hidden_states, scale)
return hidden_states
+1 -1
View File
@@ -1,7 +1,7 @@
[project] [project]
name = "ComfyUI-WanVideoWrapper" name = "ComfyUI-WanVideoWrapper"
description = "ComfyUI wrapper nodes for WanVideo" description = "ComfyUI wrapper nodes for WanVideo"
version = "1.3.9" version = "1.4.7"
license = {file = "LICENSE"} license = {file = "LICENSE"}
dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.33.0", "peft >= 0.17.0", "ftfy", "gguf >= 0.17.1", "pyloudnorm"] dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.33.0", "peft >= 0.17.0", "ftfy", "gguf >= 0.17.1", "pyloudnorm"]
+22
View File
@@ -1,5 +1,11 @@
# ComfyUI wrapper nodes for [WanVideo](https://github.com/Wan-Video/Wan2.1) and related models. # ComfyUI wrapper nodes for [WanVideo](https://github.com/Wan-Video/Wan2.1) and related models.
## Memory use update (again)
I've made everythign less reliant on torch.compile for VRAM efficiency, so things should work better even without it. Also figured workaround for some issues when using compile that made first run use drastically more VRAM, issue I battled with myself a lot.
## Update notification that can affect memory use in old workflows ## Update notification that can affect memory use in old workflows
In a recent update I changed how unmerged LoRA weights are handled: In a recent update I changed how unmerged LoRA weights are handled:
@@ -98,6 +104,22 @@ WanAnimate: https://github.com/Wan-Video/Wan2.2/tree/main/wan/modules/animate
Lynx: https://github.com/bytedance/lynx Lynx: https://github.com/bytedance/lynx
MoCha: https://github.com/Orange-3DV-Team/MoCha
UniLumos: https://github.com/alibaba-damo-academy/Lumos-Custom
Bindweave: https://github.com/bytedance/BindWeave
Training free techniques:
TimeToMove: https://github.com/time-to-move/TTM
SteadyDancer: https://github.com/MCG-NJU/SteadyDancer
One-to-all-Animation: https://github.com/ssj9596/One-to-All-Animation
SCAIL: https://github.com/zai-org/SCAIL
Not exactly Wan model, but close enough to work with the code base: Not exactly Wan model, but close enough to work with the code base:
+2 -2
View File
@@ -548,7 +548,7 @@ class WanVideoDiffusionForcingSampler:
gc.collect() gc.collect()
try: try:
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
#region main loop start #region main loop start
@@ -615,7 +615,7 @@ class WanVideoDiffusionForcingSampler:
try: try:
print_memory(device) print_memory(device)
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return ({ return ({
+106
View File
@@ -0,0 +1,106 @@
# Modify from https://github.com/liyunsheng13/dcd/blob/main/models/imagenet/mobilenetv2_dcd.py
import torch
import torch.nn as nn
import torch.nn.functional as F
class Hsigmoid(nn.Module):
def __init__(self, inplace=True):
super(Hsigmoid, self).__init__()
self.inplace = inplace
def forward(self, x):
return F.relu6(x + 3., inplace=self.inplace) / 3.
class DYModule(nn.Module):
def __init__(self, inp, oup, fc_squeeze=8):
super(DYModule, self).__init__()
self.conv = nn.Conv2d(inp, oup, 1, 1, 0, bias=False)
if inp < oup:
self.mul = 4
reduction = 8
self.avg_pool = nn.AdaptiveAvgPool2d(2)
else:
self.mul = 1
reduction = 2
self.avg_pool = nn.AdaptiveAvgPool2d(1)
self.dim = min((inp * self.mul) // reduction, oup // reduction)
while self.dim ** 2 > inp * self.mul * 2:
reduction *= 2
self.dim = min((inp * self.mul) // reduction, oup // reduction)
if self.dim < 4:
self.dim = 4
squeeze = max(inp * self.mul, self.dim ** 2) // fc_squeeze
if squeeze < 4:
squeeze = 4
self.conv_q = nn.Conv2d(inp, self.dim, 1, 1, 0, bias=False)
self.fc = nn.Sequential(
nn.Linear(inp * self.mul, squeeze, bias=False),
SEModule_small(squeeze),
)
self.fc_phi = nn.Linear(squeeze, self.dim ** 2, bias=False)
self.fc_scale = nn.Linear(squeeze, oup, bias=False)
self.hs = Hsigmoid()
self.conv_p = nn.Conv2d(self.dim, oup, 1, 1, 0, bias=False)
# self.bn1 = nn.BatchNorm2d(self.dim)
self.bn1 = nn.GroupNorm(num_groups=4, num_channels=self.dim)
# self.bn2 = nn.BatchNorm1d(self.dim)
self.bn2 = nn.GroupNorm(num_groups=4, num_channels=self.dim)
def forward(self, x):
x_type = x.dtype
r = self.conv(x.to(self.conv.weight.dtype)).to(x_type)
b, c, h, w = x.size()
y = self.avg_pool(x).view(b, c * self.mul)
y = self.fc(y)
dy_phi = self.fc_phi(y).view(b, self.dim, self.dim)
dy_scale = self.hs(self.fc_scale(y)).view(b, -1, 1, 1)
r = dy_scale.expand_as(r) * r
x = self.conv_q(x.to(self.conv_q.weight.dtype)).to(self.bn1.weight.dtype)
x = self.bn1(x)
x = x.view(b, -1, h * w)
x = x + self.bn2(torch.matmul(dy_phi, x.to(dy_phi.dtype)).to(self.bn2.weight.dtype))
x = x.view(b, -1, h, w)
x = self.conv_p(x.to(self.conv_p.weight.dtype)).to(x_type)
return x + r
class SEModule_small(nn.Module):
def __init__(self, channel):
super(SEModule_small, self).__init__()
self.fc = nn.Sequential(
nn.Linear(channel, channel, bias=False),
Hsigmoid()
)
def forward(self, x):
y = self.fc(x)
return x * y
class SEModule(nn.Module):
def __init__(self, channel, reduction=4):
super(SEModule, self).__init__()
self.avg_pool = nn.AdaptiveAvgPool2d(1)
self.fc = nn.Sequential(
nn.Linear(channel, channel // reduction, bias=False),
nn.ReLU(inplace=True),
nn.Linear(channel // reduction, channel, bias=False),
Hsigmoid()
)
def forward(self, x):
b, c, _, _ = x.size()
y = self.avg_pool(x).view(b, c)
y = self.fc(y).view(b, c, 1, 1)
return x * y.expand_as(x)
+62
View File
@@ -0,0 +1,62 @@
import os
import torch
import numpy as np
from ..utils import log
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
import comfy.model_management as mm
from comfy.utils import load_torch_file, ProgressBar
import folder_paths
script_directory = os.path.dirname(os.path.abspath(__file__))
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoAddSteadyDancerEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"pose_latents_positive": ("LATENT",),
"pose_strength_spatial": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the pose embedding"}),
"pose_strength_temporal": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the pose embedding"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
},
"optional": {
"pose_latents_negative": ("LATENT",),
"clip_vision_embeds": ("WANVIDIMAGE_CLIPEMBEDS",),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, pose_latents_positive, pose_strength_spatial, pose_strength_temporal, start_percent=0.0, end_percent=1.0, pose_latents_negative=None, clip_vision_embeds=None):
sdancer_embeds = {
"cond_pos": pose_latents_positive["samples"][0],
"cond_neg": pose_latents_negative["samples"][0] if pose_latents_negative else None,
"pose_strength_spatial": pose_strength_spatial,
"pose_strength_temporal": pose_strength_temporal,
"start_percent": start_percent,
"end_percent": end_percent,
"clip_fea": clip_vision_embeds,
}
updated = dict(embeds)
updated["sdancer_embeds"] = sdancer_embeds
return (updated,)
NODE_CLASS_MAPPINGS = {
"WanVideoAddSteadyDancerEmbeds": WanVideoAddSteadyDancerEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddSteadyDancerEmbeds": "WanVideo Add SteadyDancer Embeds",
}
+133
View File
@@ -0,0 +1,133 @@
import torch
import torch.nn as nn
class FactorConv3d(nn.Module):
"""
(2+1)D decomposition of 3D convolution: 1xHxW spatial convolution → Swish → Tx1x1 temporal convolution
"""
def __init__(self,
in_channels: int,
out_channels: int,
kernel_size,
stride: int = 1,
dilation: int = 1):
super().__init__()
if isinstance(kernel_size, int):
k_t, k_h, k_w = kernel_size, kernel_size, kernel_size
else:
k_t, k_h, k_w = kernel_size
pad_t = (k_t - 1) * dilation // 2
pad_hw = (k_h - 1) * dilation // 2
self.spatial = nn.Conv3d(
in_channels, in_channels,
kernel_size=(1, k_h, k_w),
stride=(1, stride, stride),
padding=(0, pad_hw, pad_hw),
dilation=(1, dilation, dilation),
groups=in_channels,
bias=False
)
self.temporal = nn.Conv3d(
in_channels, out_channels,
kernel_size=(k_t, 1, 1),
stride=(stride, 1, 1),
padding=(pad_t, 0, 0),
dilation=(dilation, 1, 1),
bias=True
)
self.act = nn.SiLU()
def forward(self, x):
out_dtype = x.dtype
x = self.spatial(x.to(self.spatial.weight.dtype)).to(out_dtype)
x = self.act(x)
return self.temporal(x.to(self.temporal.weight.dtype)).to(out_dtype)
class LayerNorm2D(nn.Module):
"""
LayerNorm over C for a 4-D tensor (B, C, H, W)
"""
def __init__(self, num_channels, eps=1e-5, affine=True):
super().__init__()
self.num_channels = num_channels
self.eps = eps
self.affine = affine
if affine:
self.weight = nn.Parameter(torch.ones(1, num_channels, 1, 1))
self.bias = nn.Parameter(torch.zeros(1, num_channels, 1, 1))
def forward(self, x):
# x: (B, C, H, W)
mean = x.mean(dim=1, keepdim=True) # (B, 1, H, W)
var = x.var (dim=1, keepdim=True, unbiased=False)
x = (x - mean) / torch.sqrt(var + self.eps)
if self.affine:
x = x * self.weight + self.bias
return x
class PoseRefNetNoBNV3(nn.Module):
def __init__(self,
in_channels_c: int,
in_channels_x: int,
hidden_dim: int = 256,
num_heads: int = 8,
dropout: float = 0.1):
super().__init__()
self.d_model = hidden_dim
self.nhead = num_heads
self.proj_p = nn.Conv2d(in_channels_c, hidden_dim, kernel_size=1)
self.proj_r = nn.Conv2d(in_channels_x, hidden_dim, kernel_size=1)
self.proj_p_back = nn.Conv2d(hidden_dim, in_channels_c, kernel_size=1)
self.cross_attn = nn.MultiheadAttention(hidden_dim,
num_heads=num_heads,
dropout=dropout)
self.ffn_pose = nn.Sequential(
nn.Conv2d(hidden_dim, hidden_dim, kernel_size=1),
nn.SiLU(),
nn.Conv2d(hidden_dim, hidden_dim, kernel_size=1)
)
self.norm1 = LayerNorm2D(hidden_dim)
self.norm2 = LayerNorm2D(hidden_dim)
def forward(self, pose, ref, mask=None):
"""
pose : (B, C1, T, H, W)
ref : (B, C2, T, H, W)
mask : (B, T*H*W) optional key_padding_mask
return: (B, d_model, T, H, W)
"""
B, _, T, H, W = pose.shape
p_trans = pose.permute(0, 2, 1, 3, 4).contiguous().flatten(0, 1)
r_trans = ref.permute(0, 2, 1, 3, 4).contiguous().flatten(0, 1)
p_trans = self.proj_p(p_trans.to(self.proj_p.weight.dtype)).to(self.cross_attn.in_proj_weight.dtype).flatten(2).transpose(1, 2)
r_trans = self.proj_r(r_trans.to(self.proj_r.weight.dtype)).to(self.cross_attn.in_proj_weight.dtype).flatten(2).transpose(1, 2)
out = self.cross_attn(query=r_trans,
key=p_trans,
value=p_trans,
key_padding_mask=mask)[0]
out = self.norm1(out.transpose(1, 2).contiguous().view(B*T, -1, H, W))
out_type = out.dtype
out = out + self.ffn_pose(out.to(self.ffn_pose[0].weight.dtype)).to(out_type)
out = self.norm2(out)
out = self.proj_p_back(out.to(self.proj_p_back.weight.dtype)).to(out_type)
return out.view(B, T, -1, H, W).contiguous().transpose(1, 2)
+9 -2
View File
@@ -8,6 +8,7 @@ import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from tqdm.auto import tqdm from tqdm.auto import tqdm
from collections import namedtuple from collections import namedtuple
from ..wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38
DecoderResult = namedtuple("DecoderResult", ("frame", "memory")) DecoderResult = namedtuple("DecoderResult", ("frame", "memory"))
TWorkItem = namedtuple("TWorkItem", ("input_tensor", "block_index")) TWorkItem = namedtuple("TWorkItem", ("input_tensor", "block_index"))
@@ -146,7 +147,7 @@ def apply_model_with_memblocks(model, x, parallel, show_progress_bar):
return x return x
class TAEHV(nn.Module): class TAEHV(nn.Module):
def __init__(self, state_dict, parallel=False, decoder_time_upscale=(True, True), decoder_space_upscale=(True, True, True), dtype=torch.float16): def __init__(self, state_dict, parallel=False, decoder_time_upscale=(True, True), decoder_space_upscale=(True, True, True), dtype=torch.float16, model_name="taehv"):
"""Initialize pretrained TAEHV from the given checkpoint. """Initialize pretrained TAEHV from the given checkpoint.
Arg: Arg:
@@ -161,6 +162,7 @@ class TAEHV(nn.Module):
if self.latent_channels == 48: if self.latent_channels == 48:
self.patch_size = 2 self.patch_size = 2
self.dtype = dtype self.dtype = dtype
self.model_name = model_name
self.encoder = nn.Sequential( self.encoder = nn.Sequential(
conv(self.image_channels*self.patch_size**2, 64), nn.ReLU(inplace=True), conv(self.image_channels*self.patch_size**2, 64), nn.ReLU(inplace=True),
@@ -180,8 +182,11 @@ class TAEHV(nn.Module):
) )
if state_dict is not None: if state_dict is not None:
self.load_state_dict(self.patch_tgrow_layers(state_dict)) self.load_state_dict(self.patch_tgrow_layers(state_dict))
self.parallel = parallel self.parallel = parallel
orig_vae = WanVideoVAE38() if self.latent_channels == 48 else WanVideoVAE()
self.mean = orig_vae.mean.to(dtype).movedim(1, 2)
self.inv_std = orig_vae.inv_std.to(dtype).movedim(1, 2)
def patch_tgrow_layers(self, sd): def patch_tgrow_layers(self, sd):
"""Patch TGrow layers to use a smaller kernel if needed. """Patch TGrow layers to use a smaller kernel if needed.
@@ -221,6 +226,8 @@ class TAEHV(nn.Module):
if False, frames will be processed sequentially. if False, frames will be processed sequentially.
Returns NTCHW RGB tensor with ~[0, 1] values. Returns NTCHW RGB tensor with ~[0, 1] values.
""" """
if "light" in self.model_name.lower():
x = x / self.inv_std.to(x) + self.mean.to(x)
x = apply_model_with_memblocks(self.decoder, x, self.parallel, show_progress_bar) x = apply_model_with_memblocks(self.decoder, x, self.parallel, show_progress_bar)
if self.patch_size > 1: x = F.pixel_shuffle(x, self.patch_size) if self.patch_size > 1: x = F.pixel_shuffle(x, self.patch_size)
return x[:, self.frames_to_trim:] return x[:, self.frames_to_trim:]
@@ -0,0 +1,230 @@
# https://github.com/thu-ml/DiT-Extrapolation/blob/ultra-wan/sageattn/attn_qk_int8_per_block.py
import torch
import triton
import triton.language as tl
@triton.jit
def _attn_fwd_inner(acc, l_i, m_i, q, q_scale, kv_len, current_flag,
K_ptrs, K_scale_ptr, V_ptrs, stride_kn, stride_vn,
Block_bias_ptrs, stride_bbz, stride_bbh, stride_bm, stride_bn,
Decay_mask_ptrs, stride_dmz, stride_dmh, stride_dm, stride_dn,
start_m,
BLOCK_M: tl.constexpr, HEAD_DIM: tl.constexpr, BLOCK_N: tl.constexpr,
STAGE: tl.constexpr, offs_m: tl.constexpr, offs_n: tl.constexpr,
xpos_xi: tl.constexpr = 0.9999934149894527,
frame_tokens: tl.constexpr = 1560,
sigmoid_a: tl.constexpr = 1.0,
alpha_xpos_xi: tl.constexpr = 0.9999967941742395,
beta_xpos_xi: tl.constexpr = 0.9999860536252945,
sink_width: tl.constexpr = 4,
window_width: tl.constexpr = 16,
multi_factor: tl.constexpr = None,
entropy_factor: tl.constexpr = None,
):
lo, hi = 0, kv_len
for start_n in range(lo, hi, BLOCK_N):
start_n = tl.multiple_of(start_n, BLOCK_N)
k_mask = offs_n[None, :] < (kv_len - start_n)
k = tl.load(K_ptrs, mask = k_mask)
k_scale = tl.load(K_scale_ptr)
m = offs_m[:, None]
n = start_n + offs_n
qk = tl.dot(q, k).to(tl.float32) * q_scale * k_scale
window_th = frame_tokens * window_width / 2
dist2 = tl.abs(m - n).to(tl.int32)
dist_mask = dist2 <= window_th
negative_mask = (qk<0)
qk = tl.where(dist_mask | negative_mask, qk, qk*multi_factor)
window3 = (m <= frame_tokens) & (n > window_width*frame_tokens)
qk = tl.where(window3, -1e4, qk)
m_ij = tl.maximum(m_i, tl.max(qk, 1))
qk = qk - m_ij[:, None]
p = tl.math.exp2(qk)
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_ptrs, mask = offs_n[:, None] < (kv_len - start_n))
p = p.to(tl.float16)
acc += tl.dot(p, v, out_dtype=tl.float16)
m_i = m_ij
K_ptrs += BLOCK_N * stride_kn
K_scale_ptr += 1
V_ptrs += BLOCK_N * stride_vn
return acc, l_i
@triton.jit
def _attn_fwd(Q, K, V, Q_scale, K_scale, Out,
Block_bias, Decay_mask,
flags, stride_f_b, stride_f_h,
stride_qz, stride_qh, stride_qn,
stride_kz, stride_kh, stride_kn,
stride_vz, stride_vh, stride_vn,
stride_oz, stride_oh, stride_on,
stride_bbz, stride_bbh, stride_bm, stride_bn,
stride_dmz, stride_dmh, stride_dm, stride_dn,
qo_len, kv_len, H: tl.constexpr, num_kv_groups: tl.constexpr,
HEAD_DIM: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
STAGE: tl.constexpr,
xpos_xi: tl.constexpr = 0.9999934149894527,
frame_tokens: tl.constexpr = 1560,
sigmoid_a: tl.constexpr = 1.0,
alpha_xpos_xi: tl.constexpr = 0.9999967941742395,
beta_xpos_xi: tl.constexpr = 0.9999860536252945,
sink_width: tl.constexpr = 4,
window_width: tl.constexpr = 16,
multi_factor: tl.constexpr = None,
entropy_factor: tl.constexpr = None,
):
start_m = tl.program_id(0)
off_z = tl.program_id(2).to(tl.int64)
off_h = tl.program_id(1).to(tl.int64)
q_scale_offset = (off_z * H + off_h) * tl.cdiv(qo_len, BLOCK_M)
k_scale_offset = (off_z * (H // num_kv_groups) + off_h // num_kv_groups) * tl.cdiv(kv_len, BLOCK_N)
flag_ptr = flags + off_z * stride_f_b + off_h * stride_f_h
current_flag = tl.load(flag_ptr)
offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = tl.arange(0, BLOCK_N)
offs_k = tl.arange(0, HEAD_DIM)
Q_ptrs = Q + (off_z * stride_qz + off_h * stride_qh) + offs_m[:, None] * stride_qn + offs_k[None, :]
Q_scale_ptr = Q_scale + q_scale_offset + start_m
K_ptrs = K + (off_z * stride_kz + (off_h // num_kv_groups) * stride_kh) + offs_n[None, :] * stride_kn + offs_k[:, None]
K_scale_ptr = K_scale + k_scale_offset
V_ptrs = V + (off_z * stride_vz + (off_h // num_kv_groups) * stride_vh) + offs_n[:, None] * stride_vn + offs_k[None, :]
O_block_ptr = Out + (off_z * stride_oz + off_h * stride_oh) + offs_m[:, None] * stride_on + offs_k[None, :]
# # 计算block_bias指针
Block_bias_ptrs = Block_bias + off_z * stride_bbz + off_h * stride_bbh
# 计算decay_mask指针
Decay_mask_ptrs = Decay_mask + off_z * stride_dmz + off_h * stride_dmh
m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float("inf")
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
q = tl.load(Q_ptrs, mask = offs_m[:, None] < qo_len)
q_scale = tl.load(Q_scale_ptr)
acc, l_i = _attn_fwd_inner(acc, l_i, m_i, q, q_scale, kv_len, current_flag, K_ptrs, K_scale_ptr, V_ptrs,
stride_kn, stride_vn,
Block_bias_ptrs, stride_bbz, stride_bbh, stride_bm, stride_bn,
Decay_mask_ptrs, stride_dmz, stride_dmh, stride_dm, stride_dn,
start_m,
BLOCK_M, HEAD_DIM, BLOCK_N,
4 - STAGE, offs_m, offs_n,
xpos_xi=xpos_xi,
frame_tokens=frame_tokens,
sigmoid_a=sigmoid_a,
alpha_xpos_xi=alpha_xpos_xi,
beta_xpos_xi=beta_xpos_xi,
sink_width=sink_width,
window_width=window_width,
multi_factor=multi_factor,
entropy_factor=entropy_factor,
)
acc = acc / l_i[:, None]
tl.store(O_block_ptr, acc.to(Out.type.element_ty), mask = (offs_m[:, None] < qo_len))
def forward(q, k, v, flags, block_bias, decay_mask, q_scale, k_scale, tensor_layout="HND", output_dtype=torch.float16,
xpos_xi: tl.constexpr = 0.9999934149894527,
frame_tokens: tl.constexpr = 1560,
sigmoid_a: tl.constexpr = 1.0,
alpha_xpos_xi: tl.constexpr = 0.9999967941742395,
beta_xpos_xi: tl.constexpr = 0.9999860536252945,
BLOCK_M: tl.constexpr = 128,
BLOCK_N: tl.constexpr = 128,
sink_width: tl.constexpr = 4,
window_width: tl.constexpr = 16,
multi_factor: tl.constexpr = None,
entropy_factor: tl.constexpr = None,
):
stage = 1
o = torch.empty(q.shape, dtype=output_dtype, device=q.device)
b, h_qo, qo_len, head_dim = q.shape
if block_bias is None:
block_bias = torch.zeros((b, h_qo, (qo_len + BLOCK_M - 1) // BLOCK_M, (qo_len + BLOCK_N - 1) // BLOCK_N), dtype=torch.float16, device=q.device)
if decay_mask is None:
decay_mask = torch.zeros((b, h_qo, (qo_len + BLOCK_M - 1) // BLOCK_M, (qo_len + BLOCK_N - 1) // BLOCK_N), dtype=torch.bool, device=q.device)
if tensor_layout == "HND":
b, h_qo, qo_len, head_dim = q.shape
_, h_kv, kv_len, _ = k.shape
stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(1), q.stride(2)
stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(1), k.stride(2)
stride_bz_v, stride_h_v, stride_seq_v = v.stride(0), v.stride(1), v.stride(2)
stride_bz_o, stride_h_o, stride_seq_o = o.stride(0), o.stride(1), o.stride(2)
stride_bbz, stride_bbh, stride_bm, stride_bn = block_bias.stride()
stride_dmz, stride_dmh, stride_dm, stride_dn = decay_mask.stride()
# elif tensor_layout == "NHD":
# b, qo_len, h_qo, head_dim = q.shape
# _, kv_len, h_kv, _ = k.shape
# stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(2), q.stride(1)
# stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(2), k.stride(1)
# stride_bz_v, stride_h_v, stride_seq_v = v.stride(0), v.stride(2), v.stride(1)
# stride_bz_o, stride_h_o, stride_seq_o = o.stride(0), o.stride(2), o.stride(1)
# stride_bbz, stride_bbh, stride_bm, stride_bn = block_bias.stride(0), block_bias.stride(2), block_bias.stride(1), block_bias.stride(3)
else:
raise ValueError(f"tensor_layout {tensor_layout} not supported")
stride_f_b, stride_f_h = flags.stride()
HEAD_DIM_K = head_dim
num_kv_groups = h_qo // h_kv
grid = (triton.cdiv(qo_len, BLOCK_M), h_qo, b)
_attn_fwd[grid](
q, k, v, q_scale, k_scale, o,
block_bias, decay_mask,
flags,
stride_f_b, stride_f_h,
stride_bz_q, stride_h_q, stride_seq_q,
stride_bz_k, stride_h_k, stride_seq_k,
stride_bz_v, stride_h_v, stride_seq_v,
stride_bz_o, stride_h_o, stride_seq_o,
stride_bbz, stride_bbh, stride_bm, stride_bn,
stride_dmz, stride_dmh, stride_dm, stride_dn,
qo_len, kv_len,
h_qo, num_kv_groups,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, HEAD_DIM=HEAD_DIM_K,
STAGE=stage,
num_warps=4 if head_dim == 64 else 8,
num_stages=3 if head_dim == 64 else 4,
xpos_xi=xpos_xi,
frame_tokens=frame_tokens,
sigmoid_a=sigmoid_a,
alpha_xpos_xi=alpha_xpos_xi,
beta_xpos_xi=beta_xpos_xi,
sink_width=sink_width,
window_width=window_width,
multi_factor=multi_factor,
entropy_factor=entropy_factor,
)
return o
+64
View File
@@ -0,0 +1,64 @@
# source https://github.com/thu-ml/DiT-Extrapolation/blob/ultra-wan/sageattn/core.py
import torch
import triton.language as tl
from .quant_per_block import per_block_int8
from .attn_qk_int8_per_block import forward as attn_false
from typing import Optional
def sage_attention(
qkv: list[torch.Tensor],
tensor_layout: str ="HND",
is_causal=False,
sm_scale: Optional[float] = None,
smooth_k: bool =True,
xpos_xi: tl.constexpr = 0.9999934149894527,
flags = None,
block_bias = None,
sigmoid_a: float = 1.0,
alpha_xpos_xi: float = 0.97,
beta_xpos_xi: float = 0.8,
decay_mask = None,
sink_width: int = 4,
window_width: int = 21,
multi_factor: Optional[float] = None,
entropy_factor: Optional[float] = None,
block_size : int = 64,
**kwargs
) -> torch.Tensor:
dtype = qkv[0].dtype
q, k, v = qkv[0].transpose(1, 2), qkv[1].transpose(1, 2), qkv[2].transpose(1, 2) # to HND
if flags == None:
flags = torch.zeros([q.shape[0],q.shape[1]], dtype=torch.int32, device=q.device)
seq_dim = 2
if smooth_k:
km = k.mean(dim=seq_dim, keepdim=True)
k -= km
else:
km = None
if dtype == torch.bfloat16 or dtype == torch.float32:
v = v.to(torch.float16)
if q.dtype != k.dtype or q.dtype != v.dtype:
k, v = k.to(q.dtype), v.to(q.dtype)
q_int8, q_scale, k_int8, k_scale = per_block_int8(q, k, sm_scale=sm_scale, tensor_layout=tensor_layout, BLKQ=block_size, BLKK=block_size)
del q, k
o = attn_false(q_int8, k_int8, v, flags, block_bias, decay_mask, q_scale, k_scale,
tensor_layout=tensor_layout, output_dtype=dtype, xpos_xi=xpos_xi, sigmoid_a=sigmoid_a,
alpha_xpos_xi=alpha_xpos_xi, beta_xpos_xi=beta_xpos_xi,
BLOCK_M=block_size, BLOCK_N=block_size,
sink_width=sink_width,
window_width=window_width,
multi_factor=multi_factor,
entropy_factor=entropy_factor,
)
return o.transpose(1, 2).contiguous()
+84
View File
@@ -0,0 +1,84 @@
# https://github.com/thu-ml/DiT-Extrapolation/blob/ultra-wan/sageattn/quant_per_block.py
import torch
import triton
import triton.language as tl
@triton.jit
def quant_per_block_int8_kernel(Input, Output, Scale, L,
stride_iz, stride_ih, stride_in,
stride_oz, stride_oh, stride_on,
stride_sz, stride_sh,
sm_scale,
C: tl.constexpr, BLK: tl.constexpr):
off_blk = tl.program_id(0)
off_h = tl.program_id(1)
off_b = tl.program_id(2)
offs_n = off_blk * BLK + tl.arange(0, BLK)
offs_k = tl.arange(0, C)
input_ptrs = Input + off_b * stride_iz + off_h * stride_ih + offs_n[:, None] * stride_in + offs_k[None, :]
output_ptrs = Output + off_b * stride_oz + off_h * stride_oh + offs_n[:, None] * stride_on + offs_k[None, :]
scale_ptrs = Scale + off_b * stride_sz + off_h * stride_sh + off_blk
x = tl.load(input_ptrs, mask=offs_n[:, None] < L)
x = x.to(tl.float32)
x *= sm_scale
scale = tl.max(tl.abs(x)) / 127.
x_int8 = x / scale
x_int8 += 0.5 * tl.where(x_int8 >= 0, 1, -1)
x_int8 = x_int8.to(tl.int8)
tl.store(output_ptrs, x_int8, mask=offs_n[:, None] < L)
tl.store(scale_ptrs, scale)
def per_block_int8(q, k, BLKQ=128, BLKK=64, sm_scale=None, tensor_layout="HND"):
q_int8 = torch.empty(q.shape, dtype=torch.int8, device=q.device)
k_int8 = torch.empty(k.shape, dtype=torch.int8, device=k.device)
if tensor_layout == "HND":
b, h_qo, qo_len, head_dim = q.shape
_, h_kv, kv_len, _ = k.shape
stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(1), q.stride(2)
stride_bz_qo, stride_h_qo, stride_seq_qo = q_int8.stride(0), q_int8.stride(1), q_int8.stride(2)
stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(1), k.stride(2)
stride_bz_ko, stride_h_ko, stride_seq_ko = k_int8.stride(0), k_int8.stride(1), k_int8.stride(2)
# elif tensor_layout == "NHD":
# b, qo_len, h_qo, head_dim = q.shape
# _, kv_len, h_kv, _ = k.shape
# stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(2), q.stride(1)
# stride_bz_qo, stride_h_qo, stride_seq_qo = q_int8.stride(0), q_int8.stride(2), q_int8.stride(1)
# stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(2), k.stride(1)
# stride_bz_ko, stride_h_ko, stride_seq_ko = k_int8.stride(0), k_int8.stride(2), k_int8.stride(1)
else:
raise ValueError(f"Unknown tensor layout: {tensor_layout}")
q_scale = torch.empty((b, h_qo, (qo_len + BLKQ - 1) // BLKQ, 1), device=q.device, dtype=torch.float32)
k_scale = torch.empty((b, h_kv, (kv_len + BLKK - 1) // BLKK, 1), device=q.device, dtype=torch.float32)
if sm_scale is None:
sm_scale = head_dim**-0.5
grid = ((qo_len + BLKQ - 1) // BLKQ, h_qo, b)
quant_per_block_int8_kernel[grid](
q, q_int8, q_scale, qo_len,
stride_bz_q, stride_h_q, stride_seq_q,
stride_bz_qo, stride_h_qo, stride_seq_qo,
q_scale.stride(0), q_scale.stride(1),
sm_scale=(sm_scale * 1.44269504),
C=head_dim, BLK=BLKQ
)
grid = ((kv_len + BLKK - 1) // BLKK, h_kv, b)
quant_per_block_int8_kernel[grid](
k, k_int8, k_scale, kv_len,
stride_bz_k, stride_h_k, stride_seq_k,
stride_bz_ko, stride_h_ko, stride_seq_ko,
k_scale.stride(0), k_scale.stride(1),
sm_scale=1.0,
C=head_dim, BLK=BLKK
)
return q_int8, q_scale, k_int8, k_scale
+6 -13
View File
@@ -115,16 +115,9 @@ class WanRotaryPosEmbed(nn.Module):
freqs_w = freqs[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1) freqs_w = freqs[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1)
freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1).reshape(1, 1, ppf * pph * ppw, -1) freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1).reshape(1, 1, ppf * pph * ppw, -1)
return freqs return freqs
from ..wanvideo.modules.attention import sageattn_func from ..wanvideo.modules.attention import sageattn_func
def zero_module(module):
# Zero out the parameters of a module and return it.
for p in module.parameters():
p.detach().zero_()
return module
class SimpleAttnProcessor2_0: class SimpleAttnProcessor2_0:
def __init__(self, attention_mode): def __init__(self, attention_mode):
self.attention_mode = attention_mode self.attention_mode = attention_mode
@@ -278,7 +271,7 @@ class MaskCamEmbed(nn.Module):
mid_channels = controlnet_cfg.get("mid_channels", 64) mid_channels = controlnet_cfg.get("mid_channels", 64)
self.mask_proj = nn.Sequential(nn.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8)), self.mask_proj = nn.Sequential(nn.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8)),
nn.GroupNorm(mid_channels // 8, mid_channels), nn.SiLU()) nn.GroupNorm(mid_channels // 8, mid_channels), nn.SiLU())
self.mask_zero_proj = zero_module(nn.Conv3d(mid_channels, controlnet_cfg["conv_out_dim"], kernel_size=(1, 2, 2), stride=(1, 2, 2))) self.mask_zero_proj = nn.Conv3d(mid_channels, controlnet_cfg["conv_out_dim"], kernel_size=(1, 2, 2), stride=(1, 2, 2))
def forward(self, add_inputs: torch.Tensor): def forward(self, add_inputs: torch.Tensor):
# render_mask.shape [b,c,f,h,w] # render_mask.shape [b,c,f,h,w]
@@ -321,7 +314,7 @@ class WanControlNet(ModelMixin):
) )
self.proj_out = nn.ModuleList( self.proj_out = nn.ModuleList(
[ [
zero_module(nn.Linear(self.dim, 5120)) nn.Linear(self.dim, 5120)
for _ in range(controlnet_cfg["num_layers"]) for _ in range(controlnet_cfg["num_layers"])
] ]
) )
@@ -341,7 +334,7 @@ class WanControlNet(ModelMixin):
self.controlnet_mask_embedding = MaskCamEmbed(controlnet_cfg) self.controlnet_mask_embedding = MaskCamEmbed(controlnet_cfg)
def forward(self, render_latent, render_mask, camera_embedding, temb, device): def forward(self, render_latent, render_mask, camera_embedding, temb, out_device):
controlnet_rotary_emb = self.controlnet_rope(render_latent) controlnet_rotary_emb = self.controlnet_rope(render_latent)
controlnet_inputs = self.controlnet_patch_embedding(render_latent.to(torch.float32)) controlnet_inputs = self.controlnet_patch_embedding(render_latent.to(torch.float32))
if not self.quantized: if not self.quantized:
@@ -361,7 +354,7 @@ class WanControlNet(ModelMixin):
if add_inputs is not None: if add_inputs is not None:
add_inputs = self.controlnet_mask_embedding(add_inputs) add_inputs = self.controlnet_mask_embedding(add_inputs)
controlnet_inputs = controlnet_inputs + add_inputs controlnet_inputs = controlnet_inputs + add_inputs
hidden_states = self.proj_in(controlnet_inputs) hidden_states = self.proj_in(controlnet_inputs)
controlnet_states = [] controlnet_states = []
@@ -371,6 +364,6 @@ class WanControlNet(ModelMixin):
temb=temb, temb=temb,
rotary_emb=controlnet_rotary_emb rotary_emb=controlnet_rotary_emb
) )
controlnet_states.append(self.proj_out[i](hidden_states).to(device)) controlnet_states.append(self.proj_out[i](hidden_states).to(out_device))
return controlnet_states return controlnet_states
+29 -33
View File
@@ -10,8 +10,6 @@ from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device from accelerate.utils import set_module_tensor_to_device
import folder_paths import folder_paths
import json
import numpy as np
class WanVideoUni3C_ControlnetLoader: class WanVideoUni3C_ControlnetLoader:
@classmethod @classmethod
@@ -22,7 +20,7 @@ class WanVideoUni3C_ControlnetLoader:
"base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}), "base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}),
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e5m2'], {"default": 'disabled', "tooltip": "optional quantization method"}), "quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e5m2'], {"default": 'disabled', "tooltip": "optional quantization method"}),
"load_device": (["main_device", "offload_device"], {"default": "main_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}), "load_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
"attention_mode": ([ "attention_mode": ([
"sdpa", "sdpa",
"sageattn", "sageattn",
@@ -45,17 +43,17 @@ class WanVideoUni3C_ControlnetLoader:
offload_device = mm.unet_offload_device() offload_device = mm.unet_offload_device()
transformer_load_device = device if load_device == "main_device" else offload_device transformer_load_device = device if load_device == "main_device" else offload_device
base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[base_precision] base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[base_precision]
model_path = folder_paths.get_full_path_or_raise("controlnet", model) model_path = folder_paths.get_full_path_or_raise("controlnet", model)
sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True) sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True)
if not "controlnet_patch_embedding.weight" in sd: if not "controlnet_patch_embedding.weight" in sd:
raise ValueError("Invalid ControlNet model") raise ValueError("Invalid ControlNet model")
in_channels = sd["controlnet_patch_embedding.weight"].shape[1] in_channels = sd["controlnet_patch_embedding.weight"].shape[1]
ffn_dim = sd["controlnet_blocks.0.ffn.0.bias"].shape[0] ffn_dim = sd["controlnet_blocks.0.ffn.0.bias"].shape[0]
@@ -79,7 +77,7 @@ class WanVideoUni3C_ControlnetLoader:
with init_empty_weights(): with init_empty_weights():
controlnet = WanControlNet(controlnet_cfg) controlnet = WanControlNet(controlnet_cfg)
controlnet.eval() controlnet.eval()
if quantization == "disabled": if quantization == "disabled":
for k, v in sd.items(): for k, v in sd.items():
if isinstance(v, torch.Tensor): if isinstance(v, torch.Tensor):
@@ -97,18 +95,18 @@ class WanVideoUni3C_ControlnetLoader:
else: else:
dtype = base_dtype dtype = base_dtype
params_to_keep = {"norm", "head", "time_in", "vector_in", "controlnet_patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "proj_in"} params_to_keep = {"norm", "head", "time_in", "vector_in", "controlnet_patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "proj_in"}
log.info("Using accelerate to load and assign controlnet model weights to device...") log.info("Using accelerate to load and assign controlnet model weights to device...")
param_count = sum(1 for _ in controlnet.named_parameters()) param_count = sum(1 for _ in controlnet.named_parameters())
for name, param in tqdm(controlnet.named_parameters(), for name, param in tqdm(controlnet.named_parameters(),
desc=f"Loading transformer parameters to {transformer_load_device}", desc=f"Loading transformer parameters to {transformer_load_device}",
total=param_count, total=param_count,
leave=True): leave=True):
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
if "controlnet_patch_embedding" in name: if "controlnet_patch_embedding" in name:
dtype_to_use = torch.float32 dtype_to_use = torch.float32
set_module_tensor_to_device(controlnet, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) set_module_tensor_to_device(controlnet, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
del sd del sd
if compile_args is not None: if compile_args is not None:
@@ -123,8 +121,8 @@ class WanVideoUni3C_ControlnetLoader:
for i, block in enumerate(controlnet.controlnet_blocks): for i, block in enumerate(controlnet.controlnet_blocks):
controlnet.controlnet_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) controlnet.controlnet_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
else: else:
controlnet = torch.compile(controlnet, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) controlnet = torch.compile(controlnet, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
if load_device == "offload_device" and controlnet.device != offload_device: if load_device == "offload_device" and controlnet.device != offload_device:
log.info(f"Moving controlnet model from {controlnet.device} to {offload_device}") log.info(f"Moving controlnet model from {controlnet.device} to {offload_device}")
@@ -146,6 +144,7 @@ class WanVideoUni3C_embeds:
"optional": { "optional": {
"render_latent": ("LATENT",), "render_latent": ("LATENT",),
"render_mask": ("MASK", {"tooltip": "NOT IMPLEMENTED!"}), "render_mask": ("MASK", {"tooltip": "NOT IMPLEMENTED!"}),
"offload": ("BOOLEAN", {"default": True, "tooltip": "If enabled, the controlnet model will be offloaded before main model block processing to save VRAM."}),
}, },
} }
@@ -154,9 +153,7 @@ class WanVideoUni3C_embeds:
FUNCTION = "process" FUNCTION = "process"
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
def process(self, controlnet, strength, start_percent, end_percent, render_latent=None, render_mask=None): def process(self, controlnet, strength, start_percent, end_percent, render_latent=None, render_mask=None, offload=True):
device = mm.get_torch_device()
latent_mask = latents = None latent_mask = latents = None
if render_latent is not None: if render_latent is not None:
@@ -164,17 +161,17 @@ class WanVideoUni3C_embeds:
# nframe = latents.shape[2] * 4 # nframe = latents.shape[2] * 4
# height = latents.shape[3] * 8 # height = latents.shape[3] * 8
# width = latents.shape[4] * 8 # width = latents.shape[4] * 8
if render_mask is not None: if render_mask is not None:
raise NotImplementedError("render_mask is not implemented at this time") raise NotImplementedError("render_mask is not implemented at this time")
mask = torch.nn.functional.interpolate( # mask = torch.nn.functional.interpolate(
render_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W] # render_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
size=(nframe, height, width), # size=(nframe, height, width),
mode='trilinear', # mode='trilinear',
align_corners=False # align_corners=False
).squeeze(0) # ).squeeze(0)
latent_mask = mask.unsqueeze(0).to(device) # latent_mask = mask.unsqueeze(0).to(device)
log.info(f"latent mask shape {latent_mask.shape}") # log.info(f"latent mask shape {latent_mask.shape}")
# # load camera # # load camera
# cam_info = json.load(open(f"{render_path}/cam_info.json")) # cam_info = json.load(open(f"{render_path}/cam_info.json"))
@@ -199,7 +196,7 @@ class WanVideoUni3C_embeds:
# K_inv = K.inverse() # K_inv = K.inverse()
# intrinsic = K[None].repeat(nframe, 1, 1) # intrinsic = K[None].repeat(nframe, 1, 1)
# w2c_0, c2w_0 = set_initial_camera(start_elevation, depth_avg) # w2c_0, c2w_0 = set_initial_camera(start_elevation, depth_avg)
# w2cs, c2ws, intrinsic = build_cameras(cam_traj=cam_traj, # w2cs, c2ws, intrinsic = build_cameras(cam_traj=cam_traj,
# w2c_0=w2c_0, # w2c_0=w2c_0,
@@ -215,7 +212,7 @@ class WanVideoUni3C_embeds:
# y_offset=y_offset, # y_offset=y_offset,
# z_offset=z_offset) # z_offset=z_offset)
# from .camera import get_camera_embedding # from .camera import get_camera_embedding
# camera_embedding = get_camera_embedding(intrinsic, w2cs, nframe, height, width, normalize=True) # camera_embedding = get_camera_embedding(intrinsic, w2cs, nframe, height, width, normalize=True)
#print("camera embedding shape", camera_embedding.shape) #print("camera embedding shape", camera_embedding.shape)
@@ -227,11 +224,12 @@ class WanVideoUni3C_embeds:
"end": end_percent, "end": end_percent,
"render_latent": latents, "render_latent": latents,
"render_mask": latent_mask, "render_mask": latent_mask,
"camera_embedding": None "camera_embedding": None,
"offload": offload,
} }
return (uni3c_embeds,) return (uni3c_embeds,)
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"WanVideoUni3C_ControlnetLoader": WanVideoUni3C_ControlnetLoader, "WanVideoUni3C_ControlnetLoader": WanVideoUni3C_ControlnetLoader,
"WanVideoUni3C_embeds": WanVideoUni3C_embeds, "WanVideoUni3C_embeds": WanVideoUni3C_embeds,
@@ -240,5 +238,3 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoUni3C_ControlnetLoader": "WanVideo Uni3C Controlnet Loader", "WanVideoUni3C_ControlnetLoader": "WanVideo Uni3C Controlnet Loader",
"WanVideoUni3C_embeds": "WanVideo Uni3C Embeds", "WanVideoUni3C_embeds": "WanVideo Uni3C Embeds",
} }
+45 -56
View File
@@ -9,37 +9,28 @@ from ..utils import log
import comfy.model_management as mm import comfy.model_management as mm
from comfy.utils import ProgressBar from comfy.utils import ProgressBar
import comfy.ops
ops = comfy.ops.disable_weight_init
def update_transformer(transformer, state_dict): def update_transformer(transformer, state_dict):
concat_dim = 4 concat_dim = 4
transformer.dwpose_embedding = nn.Sequential( transformer.dwpose_embedding = nn.Sequential(
nn.Conv3d(3, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)), ops.Conv3d(3, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)), nn.SiLU(),
nn.SiLU(), ops.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)), nn.SiLU(),
nn.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)), ops.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)), nn.SiLU(),
nn.SiLU(), ops.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,2,2), padding=(1,1,1)), nn.SiLU(),
nn.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)), ops.Conv3d(concat_dim * 4, concat_dim * 4, 3, stride=(2,2,2), padding=1), nn.SiLU(),
nn.SiLU(), ops.Conv3d(concat_dim * 4, concat_dim * 4, 3, stride=(2,2,2), padding=1), nn.SiLU(),
nn.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,2,2), padding=(1,1,1)), ops.Conv3d(concat_dim * 4, 5120, (1,2,2), stride=(1,2,2), padding=0))
nn.SiLU(),
nn.Conv3d(concat_dim * 4, concat_dim * 4, 3, stride=(2,2,2), padding=1),
nn.SiLU(),
nn.Conv3d(concat_dim * 4, concat_dim * 4, 3, stride=(2,2,2), padding=1),
nn.SiLU(),
nn.Conv3d(concat_dim * 4, 5120, (1,2,2), stride=(1,2,2), padding=0))
randomref_dim = 20 randomref_dim = 20
transformer.randomref_embedding_pose = nn.Sequential( transformer.randomref_embedding_pose = nn.Sequential(
nn.Conv2d(3, concat_dim * 4, 3, stride=1, padding=1), nn.Conv2d(3, concat_dim * 4, 3, stride=1, padding=1), nn.SiLU(),
nn.SiLU(), nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=1, padding=1), nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=1, padding=1), nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=1, padding=1), nn.SiLU(),
nn.SiLU(), nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1), nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=1, padding=1), nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1), nn.SiLU(),
nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1),
nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1),
nn.SiLU(),
nn.Conv2d(concat_dim * 4, randomref_dim, 3, stride=2, padding=1), nn.Conv2d(concat_dim * 4, randomref_dim, 3, stride=2, padding=1),
) )
unianimate_sd = {} unianimate_sd = {}
@@ -123,7 +114,7 @@ class DWposeDetector:
body = candidate[:,:18].copy() body = candidate[:,:18].copy()
body = body.reshape(nums*18, locs) body = body.reshape(nums*18, locs)
score = subset[:,:18].copy() score = subset[:,:18].copy()
for i in range(len(score)): for i in range(len(score)):
for j in range(len(score[i])): for j in range(len(score[i])):
if score[i][j] > score_threshold: if score[i][j] > score_threshold:
@@ -142,17 +133,17 @@ class DWposeDetector:
else: else:
bodyfoot_score[i][j] = -1 bodyfoot_score[i][j] = -1
if -1 not in bodyfoot_score[:,18] and -1 not in bodyfoot_score[:,19]: if -1 not in bodyfoot_score[:,18] and -1 not in bodyfoot_score[:,19]:
bodyfoot_score[:,18] = np.array([18.]) bodyfoot_score[:,18] = np.array([18.])
else: else:
bodyfoot_score[:,18] = np.array([-1.]) bodyfoot_score[:,18] = np.array([-1.])
if -1 not in bodyfoot_score[:,21] and -1 not in bodyfoot_score[:,22]: if -1 not in bodyfoot_score[:,21] and -1 not in bodyfoot_score[:,22]:
bodyfoot_score[:,19] = np.array([19.]) bodyfoot_score[:,19] = np.array([19.])
else: else:
bodyfoot_score[:,19] = np.array([-1.]) bodyfoot_score[:,19] = np.array([-1.])
bodyfoot_score = bodyfoot_score[:, :20] bodyfoot_score = bodyfoot_score[:, :20]
bodyfoot = candidate[:,:24].copy() bodyfoot = candidate[:,:24].copy()
for i in range(nums): for i in range(nums):
if -1 not in bodyfoot[i][18] and -1 not in bodyfoot[i][19]: if -1 not in bodyfoot[i][18] and -1 not in bodyfoot[i][19]:
bodyfoot[i][18] = (bodyfoot[i][18]+bodyfoot[i][19])/2 bodyfoot[i][18] = (bodyfoot[i][18]+bodyfoot[i][19])/2
@@ -162,7 +153,7 @@ class DWposeDetector:
bodyfoot[i][19] = (bodyfoot[i][21]+bodyfoot[i][22])/2 bodyfoot[i][19] = (bodyfoot[i][21]+bodyfoot[i][22])/2
else: else:
bodyfoot[i][19] = np.array([-1., -1.]) bodyfoot[i][19] = np.array([-1., -1.])
bodyfoot = bodyfoot[:,:20,:] bodyfoot = bodyfoot[:,:20,:]
bodyfoot = bodyfoot.reshape(nums*20, locs) bodyfoot = bodyfoot.reshape(nums*20, locs)
@@ -172,7 +163,7 @@ class DWposeDetector:
hands = candidate[:,92:113] hands = candidate[:,92:113]
hands = np.vstack([hands, candidate[:,113:]]) hands = np.vstack([hands, candidate[:,113:]])
# bodies = dict(candidate=body, subset=score) # bodies = dict(candidate=body, subset=score)
bodies = dict(candidate=bodyfoot, subset=bodyfoot_score, score=bodyfoot_score) bodies = dict(candidate=bodyfoot, subset=bodyfoot_score, score=bodyfoot_score)
pose = dict(bodies=bodies, hands=hands, faces=faces) pose = dict(bodies=bodies, hands=hands, faces=faces)
@@ -180,7 +171,7 @@ class DWposeDetector:
# return draw_pose(pose, H, W) # return draw_pose(pose, H, W)
return pose return pose
def draw_pose(pose, H, W, stick_width=4,draw_body=True, draw_hands=True, draw_feet=True, def draw_pose(pose, H, W, stick_width=4,draw_body=True, draw_hands=True, draw_feet=True,
body_keypoint_size=4, hand_keypoint_size=4, draw_head=True): body_keypoint_size=4, hand_keypoint_size=4, draw_head=True):
from .dwpose.util import draw_body_and_foot, draw_handpose, draw_facepose from .dwpose.util import draw_body_and_foot, draw_handpose, draw_facepose
bodies = pose['bodies'] bodies = pose['bodies']
@@ -202,14 +193,14 @@ def draw_pose(pose, H, W, stick_width=4,draw_body=True, draw_hands=True, draw_fe
def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_threshold, stick_width, def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_threshold, stick_width,
draw_body=True, draw_hands=True, hand_keypoint_size=4, draw_feet=True, draw_body=True, draw_hands=True, hand_keypoint_size=4, draw_feet=True,
body_keypoint_size=4, handle_not_detected="repeat", draw_head=True): body_keypoint_size=4, handle_not_detected="repeat", draw_head=True):
results_vis = [] results_vis = []
comfy_pbar = ProgressBar(len(pose_images)) comfy_pbar = ProgressBar(len(pose_images))
if ref_image is not None: if ref_image is not None:
try: try:
pose_ref = dwpose_model(ref_image.squeeze(0), score_threshold=score_threshold) pose_ref = dwpose_model(ref_image.squeeze(0), score_threshold=score_threshold)
except: except Exception:
raise ValueError("No pose detected in reference image") raise ValueError("No pose detected in reference image")
prev_pose = None prev_pose = None
for img in tqdm(pose_images, desc="Pose Extraction", unit="image", total=len(pose_images)): for img in tqdm(pose_images, desc="Pose Extraction", unit="image", total=len(pose_images)):
@@ -217,14 +208,14 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
pose = dwpose_model(img, score_threshold=score_threshold) pose = dwpose_model(img, score_threshold=score_threshold)
if handle_not_detected == "repeat": if handle_not_detected == "repeat":
prev_pose = pose prev_pose = pose
except: except Exception:
if prev_pose is not None: if prev_pose is not None:
pose = prev_pose pose = prev_pose
else: else:
pose = np.zeros_like(img) pose = np.zeros_like(img)
results_vis.append(pose) results_vis.append(pose)
comfy_pbar.update(1) comfy_pbar.update(1)
bodies = results_vis[0]['bodies'] bodies = results_vis[0]['bodies']
faces = results_vis[0]['faces'] faces = results_vis[0]['faces']
hands = results_vis[0]['hands'] hands = results_vis[0]['hands']
@@ -268,7 +259,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
results_vis[0]['faces'][:,:,1] *= y_ratio results_vis[0]['faces'][:,:,1] *= y_ratio
results_vis[0]['hands'][:,:,0] *= x_ratio results_vis[0]['hands'][:,:,0] *= x_ratio
results_vis[0]['hands'][:,:,1] *= y_ratio results_vis[0]['hands'][:,:,1] *= y_ratio
########neck######## ########neck########
l_neck_ref = ((ref_candidate[0][0] - ref_candidate[1][0]) ** 2 + (ref_candidate[0][1] - ref_candidate[1][1]) ** 2) ** 0.5 l_neck_ref = ((ref_candidate[0][0] - ref_candidate[1][0]) ** 2 + (ref_candidate[0][1] - ref_candidate[1][1]) ** 2) ** 0.5
l_neck_0 = ((candidate[0][0] - candidate[1][0]) ** 2 + (candidate[0][1] - candidate[1][1]) ** 2) ** 0.5 l_neck_0 = ((candidate[0][0] - candidate[1][0]) ** 2 + (candidate[0][1] - candidate[1][1]) ** 2) ** 0.5
@@ -287,7 +278,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
results_vis[0]['bodies']['candidate'][16,1] += y_offset_neck results_vis[0]['bodies']['candidate'][16,1] += y_offset_neck
results_vis[0]['bodies']['candidate'][17,0] += x_offset_neck results_vis[0]['bodies']['candidate'][17,0] += x_offset_neck
results_vis[0]['bodies']['candidate'][17,1] += y_offset_neck results_vis[0]['bodies']['candidate'][17,1] += y_offset_neck
########shoulder2######## ########shoulder2########
l_shoulder2_ref = ((ref_candidate[2][0] - ref_candidate[1][0]) ** 2 + (ref_candidate[2][1] - ref_candidate[1][1]) ** 2) ** 0.5 l_shoulder2_ref = ((ref_candidate[2][0] - ref_candidate[1][0]) ** 2 + (ref_candidate[2][1] - ref_candidate[1][1]) ** 2) ** 0.5
l_shoulder2_0 = ((candidate[2][0] - candidate[1][0]) ** 2 + (candidate[2][1] - candidate[1][1]) ** 2) ** 0.5 l_shoulder2_0 = ((candidate[2][0] - candidate[1][0]) ** 2 + (candidate[2][1] - candidate[1][1]) ** 2) ** 0.5
@@ -435,9 +426,9 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
results_vis[0]['bodies']['candidate'][17,0] += x_offset_head17 results_vis[0]['bodies']['candidate'][17,0] += x_offset_head17
results_vis[0]['bodies']['candidate'][17,1] += y_offset_head17 results_vis[0]['bodies']['candidate'][17,1] += y_offset_head17
########MovingAverage######## ########MovingAverage########
########left leg######## ########left leg########
l_ll1_ref = ((ref_candidate[8][0] - ref_candidate[9][0]) ** 2 + (ref_candidate[8][1] - ref_candidate[9][1]) ** 2) ** 0.5 l_ll1_ref = ((ref_candidate[8][0] - ref_candidate[9][0]) ** 2 + (ref_candidate[8][1] - ref_candidate[9][1]) ** 2) ** 0.5
l_ll1_0 = ((candidate[8][0] - candidate[9][0]) ** 2 + (candidate[8][1] - candidate[9][1]) ** 2) ** 0.5 l_ll1_0 = ((candidate[8][0] - candidate[9][0]) ** 2 + (candidate[8][1] - candidate[9][1]) ** 2) ** 0.5
@@ -522,7 +513,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
results_vis[i]['bodies']['candidate'][17,1] += y_offset_neck results_vis[i]['bodies']['candidate'][17,1] += y_offset_neck
########shoulder2######## ########shoulder2########
x_offset_shoulder2 = (results_vis[i]['bodies']['candidate'][1][0]-results_vis[i]['bodies']['candidate'][2][0])*(1.-shoulder2_ratio) x_offset_shoulder2 = (results_vis[i]['bodies']['candidate'][1][0]-results_vis[i]['bodies']['candidate'][2][0])*(1.-shoulder2_ratio)
y_offset_shoulder2 = (results_vis[i]['bodies']['candidate'][1][1]-results_vis[i]['bodies']['candidate'][2][1])*(1.-shoulder2_ratio) y_offset_shoulder2 = (results_vis[i]['bodies']['candidate'][1][1]-results_vis[i]['bodies']['candidate'][2][1])*(1.-shoulder2_ratio)
@@ -676,7 +667,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
results_vis[i]['bodies']['candidate'] += offset[np.newaxis, :] results_vis[i]['bodies']['candidate'] += offset[np.newaxis, :]
results_vis[i]['faces'] += offset[np.newaxis, np.newaxis, :] results_vis[i]['faces'] += offset[np.newaxis, np.newaxis, :]
results_vis[i]['hands'] += offset[np.newaxis, np.newaxis, :] results_vis[i]['hands'] += offset[np.newaxis, np.newaxis, :]
dwpose_woface_list = [] dwpose_woface_list = []
for i in range(len(results_vis)): for i in range(len(results_vis)):
#try: #try:
@@ -684,7 +675,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
draw_body=draw_body, draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size, draw_body=draw_body, draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size,
draw_feet=draw_feet, body_keypoint_size=body_keypoint_size, draw_head=draw_head) draw_feet=draw_feet, body_keypoint_size=body_keypoint_size, draw_head=draw_head)
result = torch.from_numpy(dwpose_woface) result = torch.from_numpy(dwpose_woface)
#except: #except Exception:
# result = torch.zeros((height, width, 3), dtype=torch.uint8) # result = torch.zeros((height, width, 3), dtype=torch.uint8)
dwpose_woface_list.append(result) dwpose_woface_list.append(result)
dwpose_woface_tensor = torch.stack(dwpose_woface_list, dim=0) dwpose_woface_tensor = torch.stack(dwpose_woface_list, dim=0)
@@ -724,11 +715,11 @@ class WanVideoUniAnimateDWPoseDetector:
FUNCTION = "process" FUNCTION = "process"
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
def process(self, pose_images, score_threshold, stick_width, reference_pose_image=None, draw_body=True, body_keypoint_size=4, def process(self, pose_images, score_threshold, stick_width, reference_pose_image=None, draw_body=True, body_keypoint_size=4,
draw_feet=True, draw_hands=True, hand_keypoint_size=4, colorspace="RGB", handle_not_detected="empty", draw_head=True): draw_feet=True, draw_hands=True, hand_keypoint_size=4, colorspace="RGB", handle_not_detected="empty", draw_head=True):
device = mm.get_torch_device() device = mm.get_torch_device()
#model loading #model loading
dw_pose_model = "dw-ll_ucoco_384_bs5.torchscript.pt" dw_pose_model = "dw-ll_ucoco_384_bs5.torchscript.pt"
yolo_model = "yolox_l.torchscript.pt" yolo_model = "yolox_l.torchscript.pt"
@@ -742,27 +733,27 @@ class WanVideoUniAnimateDWPoseDetector:
if not os.path.exists(model_det): if not os.path.exists(model_det):
log.info(f"Downloading yolo model to: {model_base_path}") log.info(f"Downloading yolo model to: {model_base_path}")
from huggingface_hub import snapshot_download from huggingface_hub import snapshot_download
snapshot_download(repo_id="hr16/yolox-onnx", snapshot_download(repo_id="hr16/yolox-onnx",
allow_patterns=[f"*{yolo_model}*"], allow_patterns=[f"*{yolo_model}*"],
local_dir=model_base_path, local_dir=model_base_path,
local_dir_use_symlinks=False) local_dir_use_symlinks=False)
if not os.path.exists(model_pose): if not os.path.exists(model_pose):
log.info(f"Downloading dwpose model to: {model_base_path}") log.info(f"Downloading dwpose model to: {model_base_path}")
from huggingface_hub import snapshot_download from huggingface_hub import snapshot_download
snapshot_download(repo_id="hr16/DWPose-TorchScript-BatchSize5", snapshot_download(repo_id="hr16/DWPose-TorchScript-BatchSize5",
allow_patterns=[f"*{dw_pose_model}*"], allow_patterns=[f"*{dw_pose_model}*"],
local_dir=model_base_path, local_dir=model_base_path,
local_dir_use_symlinks=False) local_dir_use_symlinks=False)
if not hasattr(self, "det") or not hasattr(self, "pose"): if not hasattr(self, "det") or not hasattr(self, "pose"):
self.det = torch.jit.load(model_det, map_location=device) self.det = torch.jit.load(model_det, map_location=device)
self.pose = torch.jit.load(model_pose, map_location=device) self.pose = torch.jit.load(model_pose, map_location=device)
self.dwpose_detector = DWposeDetector(self.det, self.pose) self.dwpose_detector = DWposeDetector(self.det, self.pose)
#model inference #model inference
height, width = pose_images.shape[1:3] height, width = pose_images.shape[1:3]
pose_np = pose_images.cpu().numpy() * 255 pose_np = pose_images.cpu().numpy() * 255
ref_np = None ref_np = None
if reference_pose_image is not None: if reference_pose_image is not None:
@@ -772,11 +763,11 @@ class WanVideoUniAnimateDWPoseDetector:
prev_fuser_state = torch._C._jit_texpr_fuser_enabled() prev_fuser_state = torch._C._jit_texpr_fuser_enabled()
torch._C._jit_set_texpr_fuser_enabled(False) # removes warmup delay, may want to enable later torch._C._jit_set_texpr_fuser_enabled(False) # removes warmup delay, may want to enable later
poses, reference_pose = pose_extract(pose_np, ref_np, self.dwpose_detector, height, width, score_threshold, stick_width=stick_width, poses, reference_pose = pose_extract(pose_np, ref_np, self.dwpose_detector, height, width, score_threshold, stick_width=stick_width,
draw_body=draw_body, body_keypoint_size=body_keypoint_size, draw_feet=draw_feet, draw_body=draw_body, body_keypoint_size=body_keypoint_size, draw_feet=draw_feet,
draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size, handle_not_detected=handle_not_detected, draw_head=draw_head) draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size, handle_not_detected=handle_not_detected, draw_head=draw_head)
poses = poses / 255.0 poses = poses / 255.0
torch._C._jit_set_texpr_fuser_enabled(prev_fuser_state) torch._C._jit_set_texpr_fuser_enabled(prev_fuser_state)
if reference_pose_image is not None: if reference_pose_image is not None:
reference_pose = reference_pose.unsqueeze(0) / 255.0 reference_pose = reference_pose.unsqueeze(0) / 255.0
else: else:
@@ -828,11 +819,9 @@ class WanVideoUniAnimatePoseInput:
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"WanVideoUniAnimatePoseInput": WanVideoUniAnimatePoseInput, "WanVideoUniAnimatePoseInput": WanVideoUniAnimatePoseInput,
"WanVideoUniAnimateDWPoseDetector": WanVideoUniAnimateDWPoseDetector, "WanVideoUniAnimateDWPoseDetector": WanVideoUniAnimateDWPoseDetector,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoUniAnimatePoseInput": "WanVideo UniAnimate Pose Input", "WanVideoUniAnimatePoseInput": "WanVideo UniAnimate Pose Input",
"WanVideoUniAnimateDWPoseDetector": "WanVideo UniAnimate DWPose Detector", "WanVideoUniAnimateDWPoseDetector": "WanVideo UniAnimate DWPose Detector",
} }
+185 -18
View File
@@ -4,17 +4,121 @@ import logging
import math import math
from tqdm import tqdm from tqdm import tqdm
from pathlib import Path from pathlib import Path
import os import gc
import types, collections import types, collections
from comfy.utils import ProgressBar, copy_to_param, set_attr_param from comfy.utils import ProgressBar, copy_to_param, set_attr_param
from comfy.model_patcher import get_key_weight, string_to_seed from comfy.model_patcher import get_key_weight
from comfy.lora import calculate_weight from comfy.lora import calculate_weight
from comfy.model_management import cast_to_device
try:
from comfy.utils import string_to_seed
except Exception:
from comfy.model_patcher import string_to_seed
from comfy.float import stochastic_rounding from comfy.float import stochastic_rounding
from .custom_linear import remove_lora_from_module
import folder_paths import folder_paths
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s') logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
try:
from .gguf.gguf import GGUFParameter
except Exception:
pass
COLOR_CODES = {
"reset": "\033[0m",
"red": "\033[31m",
"green": "\033[32m",
"yellow": "\033[33m",
"blue": "\033[34m",
"magenta": "\033[35m",
"cyan": "\033[36m",
"white": "\033[37m",
}
def color_text(text, color):
try:
return f"{COLOR_CODES.get(color, COLOR_CODES['reset'])}{text}{COLOR_CODES['reset']}"
except Exception:
return text
class MetaParameter(torch.nn.Parameter):
def __new__(cls, dtype, quant_type=None):
data = torch.empty(0, dtype=dtype)
self = torch.nn.Parameter(data, requires_grad=False)
self.quant_type = quant_type
return self
def offload_transformer(transformer, remove_lora=True):
transformer.teacache_state.clear_all()
transformer.magcache_state.clear_all()
transformer.easycache_state.clear_all()
if transformer.patched_linear:
for name, param in transformer.named_parameters():
if "loras" in name or "controlnet" in name:
continue
module = transformer
subnames = name.split('.')
for subname in subnames[:-1]:
module = getattr(module, subname)
attr_name = subnames[-1]
if param.data.is_floating_point():
meta_param = torch.nn.Parameter(torch.empty_like(param.data, device='meta'), requires_grad=False)
setattr(module, attr_name, meta_param)
elif isinstance(param.data, GGUFParameter):
quant_type = getattr(param, 'quant_type', None)
setattr(module, attr_name, MetaParameter(param.data.dtype, quant_type))
else:
pass
if remove_lora:
remove_lora_from_module(transformer)
else:
transformer.to(offload_device)
for block in transformer.blocks:
block.kv_cache = None
if transformer.audio_model is not None and hasattr(block, 'audio_block'):
block.audio_block = None
mm.soft_empty_cache()
gc.collect()
def init_blockswap(transformer, block_swap_args, model):
if not transformer.patched_linear:
if block_swap_args is not None:
for name, param in transformer.named_parameters():
if "block" not in name or "control_adapter" in name or "face" in name:
param.data = param.data.to(device)
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
param.data = param.data.to(offload_device)
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
param.data = param.data.to(offload_device)
transformer.block_swap(
block_swap_args["blocks_to_swap"] - 1 ,
block_swap_args["offload_txt_emb"],
block_swap_args["offload_img_emb"],
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
)
elif model["auto_cpu_offload"]:
for module in transformer.modules():
if hasattr(module, "offload"):
module.offload()
if hasattr(module, "onload"):
module.onload()
for block in transformer.blocks:
block.modulation = torch.nn.Parameter(block.modulation.to(device))
transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device))
else:
transformer.to(device)
def check_device_same(first_device, second_device): def check_device_same(first_device, second_device):
if first_device.type != second_device.type: if first_device.type != second_device.type:
return False return False
@@ -91,9 +195,9 @@ def set_module_tensor_to_device(module, tensor_name, device, value=None, dtype=N
device = device_quantization device = device_quantization
if is_buffer: if is_buffer:
module._buffers[tensor_name] = new_value module._buffers[tensor_name] = new_value
elif value is not None or not check_device_same(torch.device(device), module._parameters[tensor_name].device): elif value is not None or not check_device_same(device, module._parameters[tensor_name].device):
param_cls = type(module._parameters[tensor_name]) param_cls = type(module._parameters[tensor_name])
new_value = param_cls(new_value, requires_grad=False).to(device) new_value = param_cls(new_value, requires_grad=False)
module._parameters[tensor_name] = new_value module._parameters[tensor_name] = new_value
#if device != "cpu": #if device != "cpu":
@@ -108,13 +212,11 @@ def check_diffusers_version():
except importlib.metadata.PackageNotFoundError: except importlib.metadata.PackageNotFoundError:
raise AssertionError("diffusers is not installed.") raise AssertionError("diffusers is not installed.")
def print_memory(device): def print_memory(device, process="Sampling"):
memory = torch.cuda.memory_allocated(device) / 1024**3
max_memory = torch.cuda.max_memory_allocated(device) / 1024**3 max_memory = torch.cuda.max_memory_allocated(device) / 1024**3
max_reserved = torch.cuda.max_memory_reserved(device) / 1024**3 max_reserved = torch.cuda.max_memory_reserved(device) / 1024**3
log.info(f"Allocated memory: {memory=:.3f} GB") log.info(f"[{process}] Max allocated memory: {max_memory=:.3f} GB")
log.info(f"Max allocated memory: {max_memory=:.3f} GB") log.info(f"[{process}] Max reserved memory: {max_reserved=:.3f} GB")
log.info(f"Max reserved memory: {max_reserved=:.3f} GB")
#memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False) #memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False)
#log.info(f"Memory Summary:\n{memory_summary}") #log.info(f"Memory Summary:\n{memory_summary}")
@@ -125,6 +227,18 @@ def get_module_memory_mb(module):
memory += param.nelement() * param.element_size() memory += param.nelement() * param.element_size()
return memory / (1024 * 1024) # Convert to MB return memory / (1024 * 1024) # Convert to MB
def get_module_memory_mb_per_device(module):
memory_per_device = {}
memory = 0
for param in module.parameters():
if param.data is not None:
device = str(param.device)
memory += param.nelement() * param.element_size()
memory_per_device[device] = memory_per_device.get(device, 0) + memory
memory_per_device = {dev: mem / (1024 * 1024) for dev, mem in memory_per_device.items()}
return memory_per_device
def get_tensor_memory(tensor): def get_tensor_memory(tensor):
memory_bytes = tensor.element_size() * tensor.nelement() memory_bytes = tensor.element_size() * tensor.nelement()
return f"{memory_bytes / (1024 * 1024):.2f} MB" return f"{memory_bytes / (1024 * 1024):.2f} MB"
@@ -140,7 +254,7 @@ def patch_weight_to_device(self, key, device_to=None, inplace_update=False, back
self.backup[key] = collections.namedtuple('Dimension', ['weight', 'inplace_update'])(weight.to(device=self.offload_device, copy=inplace_update), inplace_update) self.backup[key] = collections.namedtuple('Dimension', ['weight', 'inplace_update'])(weight.to(device=self.offload_device, copy=inplace_update), inplace_update)
if device_to is not None: if device_to is not None:
temp_weight = cast_to_device(weight, device_to, torch.float32, copy=True) temp_weight = mm.cast_to_device(weight, device_to, torch.float32, copy=True)
else: else:
temp_weight = weight.to(torch.float32, copy=True) temp_weight = weight.to(torch.float32, copy=True)
if convert_func is not None: if convert_func is not None:
@@ -195,7 +309,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
key = f"{name.replace('diffusion_model.', '')}.{param}" key = f"{name.replace('diffusion_model.', '')}.{param}"
try: try:
set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[key]) set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[key])
except: except Exception:
continue continue
key = f"{name}.{param}" key = f"{name}.{param}"
if scale_weights is not None: if scale_weights is not None:
@@ -209,7 +323,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
if low_mem_load: if low_mem_load:
try: try:
set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=model.model.diffusion_model.state_dict()[key]) set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=model.model.diffusion_model.state_dict()[key])
except: except Exception:
continue continue
m.comfy_patched_weights = True m.comfy_patched_weights = True
cnt += 1 cnt += 1
@@ -238,7 +352,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
dtype_to_use = torch.float32 dtype_to_use = torch.float32
try: try:
set_module_tensor_to_device(model.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[name]) set_module_tensor_to_device(model.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[name])
except: except Exception:
continue continue
return model return model
@@ -584,17 +698,18 @@ def check_duplicate_nodes():
"""Check ComfyUI custom_nodes directory for duplicate installations""" """Check ComfyUI custom_nodes directory for duplicate installations"""
custom_nodes_dir = Path(folder_paths.folder_names_and_paths["custom_nodes"][0][0]) custom_nodes_dir = Path(folder_paths.folder_names_and_paths["custom_nodes"][0][0])
current_path = Path(__file__).parent current_path = Path(__file__).parent
wanvideo_dirs = [] wanvideo_dirs = []
# Check all directories in custom_nodes # Check all directories in custom_nodes
for path in custom_nodes_dir.iterdir(): for path in custom_nodes_dir.iterdir():
if (path.is_dir() and if (path.is_dir() and
path != current_path and path != current_path and
not path.name.endswith('.disabled') and
'wanvideo' in path.name.lower() and 'wanvideo' in path.name.lower() and
'wrapper' in path.name.lower()): 'wrapper' in path.name.lower()):
wanvideo_dirs.append(str(path)) wanvideo_dirs.append(str(path))
return wanvideo_dirs return wanvideo_dirs
#https://github.com/temporalscorerescaling/TSR/ #https://github.com/temporalscorerescaling/TSR/
@@ -609,3 +724,55 @@ def temporal_score_rescaling(model_output, sample, timestep, k=1.0, tsr_sigma=0.
if not t == 1.0: if not t == 1.0:
model_output = (ratio * ((1-t) * model_output + sample) - sample) / (1 - t) model_output = (ratio * ((1-t) * model_output + sample) - sample) / (1 - t)
return model_output return model_output
def match_and_blend_colors(
source_chunk: torch.Tensor, # (C, T, H, W), range [-1, 1]
reference_image: torch.Tensor, # (C, 1, H, W), range [-1, 1]
strength: float,
) -> torch.Tensor:
import kornia
if strength == 0.0:
return source_chunk
source_chunk = source_chunk.unsqueeze(0) # (1, C, T, H, W)
# shapes
B, C, T, H, W = source_chunk.shape
input_dtype = source_chunk.dtype
# [-1,1] -> [0,1]
src_01 = (source_chunk + 1.0) * 0.5
ref_01 = (reference_image + 1.0) * 0.5
src32 = src_01.to(torch.float32)
ref32 = ref_01.to(torch.float32)
# (B, C, T, H, W) -> (B*T, C, H, W)
src_bt = src32.permute(0, 2, 1, 3, 4).contiguous().view(B * T, C, H, W)
ref_bchw = ref32[:, :, 0, :, :].contiguous()
# RGB->Lab
src_lab = kornia.color.rgb_to_lab(src_bt) # (B*T, C, H, W)
ref_lab = kornia.color.rgb_to_lab(ref_bchw) # (B, C, H, W)
src_lab_flat = src_lab.view(B * T, C, -1) # (B*T, C, HW)
ref_lab_flat = ref_lab.view(B, C, -1) # (B, C, HW)
src_std, src_mean = torch.std_mean(src_lab_flat, dim=-1, keepdim=True, unbiased=False)
ref_std, ref_mean = torch.std_mean(ref_lab_flat, dim=-1, keepdim=True, unbiased=False)
src_std = src_std.clamp_min_(1e-6)
ref_mean_bt = ref_mean.repeat_interleave(T, dim=0) # (B*T, C, 1)
ref_std_bt = ref_std.repeat_interleave(T, dim=0) # (B*T, C, 1)
corrected_lab_flat = (src_lab_flat - src_mean) * (ref_std_bt / src_std) + ref_mean_bt
corrected_lab = corrected_lab_flat.view(B * T, C, H, W)
# Lab->RGB
corrected_rgb_01 = kornia.color.lab_to_rgb(corrected_lab) # (B*T, C, H, W)
blended_rgb_01 = (1.0 - strength) * src_bt + strength * corrected_rgb_01
# (B, C, T, H, W)
blended_rgb_01 = blended_rgb_01.view(B, T, C, H, W).permute(0, 2, 1, 3, 4).contiguous()
# [0,1] -> [-1,1]
return (blended_rgb_01 * 2.0 - 1.0)[0].to(dtype=input_dtype)
+71 -195
View File
@@ -1,25 +1,21 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import torch import torch
from ...utils import log from ...utils import log
# Flash Attention imports from comfy.ldm.modules.attention import optimized_attention
try:
import flash_attn_interface def attention_func_error(*args, **kwargs):
FLASH_ATTN_3_AVAILABLE = True raise ImportError("Selected attention mode not available. Please ensure required packages are installed correctly.")
except Exception as e:
FLASH_ATTN_3_AVAILABLE = False from .attention_flash import flash_attention
try:
import flash_attn
FLASH_ATTN_2_AVAILABLE = True
except Exception as e:
FLASH_ATTN_2_AVAILABLE = False
# Sage Attention imports # Sage Attention imports
# using custom ops to avoid graph breaks with torch.compile
try: try:
from sageattention import sageattn from sageattention import sageattn
@torch.compiler.disable()
def sageattn_func(q, k, v, attn_mask=None, dropout_p=0, is_causal=False, tensor_layout="HND"): @torch.library.custom_op("wanvideo::sageattn", mutates_args=())
def sageattn_func(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, attn_mask: torch.Tensor | None = None, dropout_p: float = 0.0, is_causal: bool = False, tensor_layout: str = "HND"
) -> torch.Tensor:
if not (q.dtype == k.dtype == v.dtype): if not (q.dtype == k.dtype == v.dtype):
return sageattn(q, k.to(q.dtype), v.to(q.dtype), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout) return sageattn(q, k.to(q.dtype), v.to(q.dtype), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout)
elif q.dtype == torch.float32: elif q.dtype == torch.float32:
@@ -27,6 +23,13 @@ try:
else: else:
return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout) return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout)
@sageattn_func.register_fake
def _(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=False, tensor_layout="HND"):
# Return tensor with same shape as q
return q.clone()
sageattn_func = torch.ops.wanvideo.sageattn
def sageattn_func_compiled(q, k, v, attn_mask=None, dropout_p=0, is_causal=False, tensor_layout="HND"): def sageattn_func_compiled(q, k, v, attn_mask=None, dropout_p=0, is_causal=False, tensor_layout="HND"):
if not (q.dtype == k.dtype == v.dtype): if not (q.dtype == k.dtype == v.dtype):
return sageattn(q, k.to(q.dtype), v.to(q.dtype), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout) return sageattn(q, k.to(q.dtype), v.to(q.dtype), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout)
@@ -40,20 +43,14 @@ except Exception as e:
log.warning("sageattention package is not installed, sageattention will not be available") log.warning("sageattention package is not installed, sageattention will not be available")
elif isinstance(e, ImportError) and "DLL" in str(e): elif isinstance(e, ImportError) and "DLL" in str(e):
log.warning("sageattention DLL loading error, sageattention will not be available") log.warning("sageattention DLL loading error, sageattention will not be available")
sageattn_func = None sageattn_func = attention_func_error
try: try:
from sageattn3 import sageattn3_blackwell as sageattn_blackwell
except:
try:
from sageattn import sageattn_blackwell
except:
SAGE3_AVAILABLE = False
try:
from sageattention import sageattn_varlen from sageattention import sageattn_varlen
@torch.compiler.disable() from typing import List
def sageattn_varlen_func(q, k, v, q_lens, k_lens, max_seqlen_q, max_seqlen_k, dropout_p=0, is_causal=False):
@torch.library.custom_op("wanvideo::sageattn_varlen", mutates_args=())
def sageattn_varlen_func(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, q_lens: List[int], k_lens: List[int], max_seqlen_q: int, max_seqlen_k: int, dropout_p: float = 0.0, is_causal: bool = False) -> torch.Tensor:
cu_seqlens_q = torch.tensor([0] + list(torch.cumsum(torch.tensor(q_lens), dim=0)), device=q.device, dtype=torch.int32) cu_seqlens_q = torch.tensor([0] + list(torch.cumsum(torch.tensor(q_lens), dim=0)), device=q.device, dtype=torch.int32)
cu_seqlens_k = torch.tensor([0] + list(torch.cumsum(torch.tensor(k_lens), dim=0)), device=q.device, dtype=torch.int32) cu_seqlens_k = torch.tensor([0] + list(torch.cumsum(torch.tensor(k_lens), dim=0)), device=q.device, dtype=torch.int32)
if not (q.dtype == k.dtype == v.dtype): if not (q.dtype == k.dtype == v.dtype):
@@ -62,180 +59,59 @@ try:
return sageattn_varlen(q.to(torch.float16), k.to(torch.float16), v.to(torch.float16), cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, dropout_p=dropout_p, is_causal=is_causal).to(torch.float32) return sageattn_varlen(q.to(torch.float16), k.to(torch.float16), v.to(torch.float16), cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, dropout_p=dropout_p, is_causal=is_causal).to(torch.float32)
else: else:
return sageattn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, dropout_p=dropout_p, is_causal=is_causal) return sageattn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, dropout_p=dropout_p, is_causal=is_causal)
except:
sageattn_varlen_func = None
__all__ = [ @sageattn_varlen_func.register_fake
'flash_attention', def _(q, k, v, q_lens, k_lens, max_seqlen_q, max_seqlen_k, dropout_p=0.0, is_causal=False):
'attention', # Return tensor with same shape as q
] return q.clone()
sageattn_varlen_func = torch.ops.wanvideo.sageattn_varlen
except Exception:
sageattn_varlen_func = attention_func_error
# sage3
try:
from sageattn3 import sageattn3_blackwell as sageattn_blackwell
except Exception:
try:
from sageattn import sageattn_blackwell
except Exception:
sageattn_blackwell = attention_func_error
try:
from ...ultravico.sageattn.core import sage_attention as sageattn_ultravico
@torch.library.custom_op("wanvideo::sageattn_ultravico", mutates_args=())
def sageattn_func_ultravico(qkv: List[torch.Tensor], attn_mask: torch.Tensor | None = None, dropout_p: float = 0.0, is_causal: bool = False, multi_factor: float = 0.9, frame_tokens: int = 1536
) -> torch.Tensor:
return sageattn_ultravico(qkv, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, multi_factor=multi_factor, frame_tokens=frame_tokens)
@sageattn_func_ultravico.register_fake
def _(qkv, attn_mask=None, dropout_p=0.0, is_causal=False, multi_factor=0.9):
return torch.empty_like(qkv[0]).contiguous()
sageattn_func_ultravico = torch.ops.wanvideo.sageattn_ultravico
except Exception:
sageattn_func_ultravico = attention_func_error
def flash_attention( def attention(q, k, v, q_lens=None, k_lens=None, max_seqlen_q=None, max_seqlen_k=None, dropout_p=0.,
q, softmax_scale=None, q_scale=None, causal=False, window_size=(-1, -1), deterministic=False, dtype=torch.bfloat16,
k, attention_mode='sdpa', attn_mask=None, transformer_options={}, frame_tokens=1536, heads=128):
v,
q_lens=None,
k_lens=None,
dropout_p=0.,
softmax_scale=None,
q_scale=None,
causal=False,
window_size=(-1, -1),
deterministic=False,
dtype=torch.bfloat16,
version=None,
):
"""
q: [B, Lq, Nq, C1].
k: [B, Lk, Nk, C1].
v: [B, Lk, Nk, C2]. Nq must be divisible by Nk.
q_lens: [B].
k_lens: [B].
dropout_p: float. Dropout probability.
softmax_scale: float. The scaling of QK^T before applying softmax.
causal: bool. Whether to apply causal attention mask.
window_size: (left right). If not (-1, -1), apply sliding window local attention.
deterministic: bool. If True, slightly slower and uses more memory.
dtype: torch.dtype. Apply when dtype of q/k/v is not float16/bfloat16.
"""
half_dtypes = (torch.float16, torch.bfloat16)
#assert dtype in half_dtypes
#assert q.device.type == 'cuda' and q.size(-1) <= 256
# params
b, lq, lk, out_dtype = q.size(0), q.size(1), k.size(1), q.dtype
def half(x):
return x if x.dtype in half_dtypes else x.to(dtype)
# preprocess query
if q_lens is None:
q = half(q.flatten(0, 1))
q_lens = torch.tensor(
[lq] * b, dtype=torch.int32).to(
device=q.device, non_blocking=True)
else:
q = half(torch.cat([u[:v] for u, v in zip(q, q_lens)]))
# preprocess key, value
if k_lens is None:
k = half(k.flatten(0, 1))
v = half(v.flatten(0, 1))
k_lens = torch.tensor(
[lk] * b, dtype=torch.int32).to(
device=k.device, non_blocking=True)
else:
k = half(torch.cat([u[:v] for u, v in zip(k, k_lens)]))
v = half(torch.cat([u[:v] for u, v in zip(v, k_lens)]))
q = q.to(v.dtype)
k = k.to(v.dtype)
if q_scale is not None:
q = q * q_scale
if version is not None and version == 3 and not FLASH_ATTN_3_AVAILABLE:
log.warning('Flash attention 3 is not available, use flash attention 2 instead.')
# apply attention
if (version is None or version == 3) and FLASH_ATTN_3_AVAILABLE:
# Note: dropout_p, window_size are not supported in FA3 now.
x = flash_attn_interface.flash_attn_varlen_func(
q=q,
k=k,
v=v,
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
seqused_q=None,
seqused_k=None,
max_seqlen_q=lq,
max_seqlen_k=lk,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic)[0].unflatten(0, (b, lq))
else:
assert FLASH_ATTN_2_AVAILABLE
x = flash_attn.flash_attn_varlen_func(
q=q,
k=k,
v=v,
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
max_seqlen_q=lq,
max_seqlen_k=lk,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
causal=causal,
window_size=window_size,
deterministic=deterministic).unflatten(0, (b, lq))
# output
return x.type(out_dtype)
def attention(
q,
k,
v,
q_lens=None,
k_lens=None,
max_seqlen_q=None,
max_seqlen_k=None,
dropout_p=0.,
softmax_scale=None,
q_scale=None,
causal=False,
window_size=(-1, -1),
deterministic=False,
dtype=torch.bfloat16,
attention_mode='sdpa',
attn_mask=None,
):
if "flash" in attention_mode: if "flash" in attention_mode:
if attention_mode == 'flash_attn_2': return flash_attention(q, k, v, q_lens=q_lens, k_lens=k_lens, dropout_p=dropout_p, softmax_scale=softmax_scale,
fa_version = 2 q_scale=q_scale, causal=causal, window_size=window_size, deterministic=deterministic, dtype=dtype, version=2 if attention_mode == 'flash_attn_2' else 3,
elif attention_mode == 'flash_attn_3':
fa_version = 3
return flash_attention(
q=q,
k=k,
v=v,
q_lens=q_lens,
k_lens=k_lens,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
q_scale=q_scale,
causal=causal,
window_size=window_size,
deterministic=deterministic,
dtype=dtype,
version=fa_version,
) )
elif attention_mode == 'sdpa': elif attention_mode == 'sageattn_3':
return sageattn_blackwell(q.transpose(1,2), k.transpose(1,2), v.transpose(1,2), per_block_mean=False).transpose(1,2).contiguous()
elif attention_mode == 'sageattn_varlen':
return sageattn_varlen_func(q,k,v, q_lens=q_lens, k_lens=k_lens, max_seqlen_k=max_seqlen_k, max_seqlen_q=max_seqlen_q)
elif attention_mode == 'sageattn_compiled': # for sage versions that allow torch.compile, may be redundant now as other sageattn ops are wrapper in custom ops
return sageattn_func_compiled(q, k, v, tensor_layout="NHD").contiguous()
elif attention_mode == 'sageattn':
return sageattn_func(q, k, v, tensor_layout="NHD").contiguous()
elif attention_mode == 'sageattn_ultravico':
return sageattn_func_ultravico([q, k, v], multi_factor=transformer_options.get("ultravico_alpha", 0.9), frame_tokens=frame_tokens).contiguous()
elif attention_mode == 'comfy':
return optimized_attention(q.transpose(1,2), k.transpose(1,2), v.transpose(1,2), heads=heads, skip_reshape=True)
else: # sdpa
if not (q.dtype == k.dtype == v.dtype): if not (q.dtype == k.dtype == v.dtype):
return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2).to(q.dtype), v.transpose(1, 2).to(q.dtype), attn_mask=attn_mask).transpose(1, 2).contiguous() return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2).to(q.dtype), v.transpose(1, 2).to(q.dtype), attn_mask=attn_mask).transpose(1, 2).contiguous()
return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), attn_mask=attn_mask).transpose(1, 2).contiguous() return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), attn_mask=attn_mask).transpose(1, 2).contiguous()
elif attention_mode == 'sageattn_3':
return sageattn_blackwell(
q.transpose(1,2),
k.transpose(1,2),
v.transpose(1,2),
per_block_mean=False #seems necessary for reasonable VRAM usage, not sure of other implications
).transpose(1,2).contiguous()
elif attention_mode == 'sageattn_varlen':
return sageattn_varlen_func(
q,k,v,
q_lens=q_lens,
k_lens=k_lens,
max_seqlen_k=max_seqlen_k,
max_seqlen_q=max_seqlen_q
)
elif attention_mode == 'sageattn_compiled':
return sageattn_func_compiled(q, k, v, tensor_layout="NHD").contiguous()
else:
return sageattn_func(q, k, v, tensor_layout="NHD").contiguous()
+80
View File
@@ -0,0 +1,80 @@
import torch
from ...utils import log
def attention_func_error(*args, **kwargs):
raise ImportError("Selected attention mode not available. Please ensure required packages are installed correctly.")
try:
import flash_attn_interface
FLASH_ATTN_3_AVAILABLE = True
except Exception as e:
FLASH_ATTN_3_AVAILABLE = False
try:
import flash_attn
FLASH_ATTN_2_AVAILABLE = True
except Exception as e:
FLASH_ATTN_2_AVAILABLE = False
if not FLASH_ATTN_2_AVAILABLE and not FLASH_ATTN_3_AVAILABLE:
flash_attention = attention_func_error
else:
def flash_attention(q, k, v, q_lens=None, k_lens=None, dropout_p=0., softmax_scale=None, q_scale=None, causal=False, window_size=(-1, -1), deterministic=False, dtype=torch.bfloat16, version=None):
half_dtypes = (torch.float16, torch.bfloat16)
# params
b, lq, lk, out_dtype = q.size(0), q.size(1), k.size(1), q.dtype
def half(x):
return x if x.dtype in half_dtypes else x.to(dtype)
# preprocess query
if q_lens is None:
q = half(q.flatten(0, 1))
q_lens = torch.tensor(
[lq] * b, dtype=torch.int32).to(
device=q.device, non_blocking=True)
else:
q = half(torch.cat([u[:v] for u, v in zip(q, q_lens)]))
# preprocess key, value
if k_lens is None:
k = half(k.flatten(0, 1))
v = half(v.flatten(0, 1))
k_lens = torch.tensor(
[lk] * b, dtype=torch.int32).to(
device=k.device, non_blocking=True)
else:
k = half(torch.cat([u[:v] for u, v in zip(k, k_lens)]))
v = half(torch.cat([u[:v] for u, v in zip(v, k_lens)]))
q = q.to(v.dtype)
k = k.to(v.dtype)
if q_scale is not None:
q = q * q_scale
if version is not None and version == 3 and not FLASH_ATTN_3_AVAILABLE:
log.warning('Flash attention 3 is not available, use flash attention 2 instead.')
if (version is None or version == 3) and FLASH_ATTN_3_AVAILABLE:
# Note: dropout_p, window_size are not supported in FA3 now.
x = flash_attn_interface.flash_attn_varlen_func(q=q, k=k, v=v,
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
seqused_q=None, seqused_k=None, max_seqlen_q=lq, max_seqlen_k=lk,
softmax_scale=softmax_scale, causal=causal,
deterministic=deterministic).unflatten(0, (b, lq))
else:
assert FLASH_ATTN_2_AVAILABLE
x = flash_attn.flash_attn_varlen_func(q=q, k=k, v=v,
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
max_seqlen_q=lq, max_seqlen_k=lk, dropout_p=dropout_p,
softmax_scale=softmax_scale, causal=causal, window_size=window_size,
deterministic=deterministic).unflatten(0, (b, lq))
return x.type(out_dtype)
+869 -487
View File
File diff suppressed because it is too large Load Diff
+3 -3
View File
@@ -4,15 +4,15 @@ import torch
try: try:
from spas_sage_attn import block_sparse_sage2_attn_cuda from spas_sage_attn import block_sparse_sage2_attn_cuda
sparse_attn_func = block_sparse_sage2_attn_cuda sparse_attn_func = block_sparse_sage2_attn_cuda
except: except Exception:
try: try:
from sparse_sageattn import sparse_sageattn from sparse_sageattn import sparse_sageattn
sparse_attn_func = sparse_sageattn sparse_attn_func = sparse_sageattn
except: except Exception:
try: try:
from .sparse_sage.core import sparse_sageattn from .sparse_sage.core import sparse_sageattn
sparse_attn_func = sparse_sageattn sparse_attn_func = sparse_sageattn
except: except Exception:
sparse_sageattn = None sparse_sageattn = None
raise ImportError("sparse_sageattn is not available. Please install the sparse_sageattn package or check your import path.") raise ImportError("sparse_sageattn is not available. Please install the sparse_sageattn package or check your import path.")
+73 -25
View File
@@ -1,12 +1,15 @@
import torch import torch
from .fm_solvers import (FlowDPMSolverMultistepScheduler, get_sampling_sigmas, retrieve_timesteps) import numpy as np
from .fm_solvers import (FlowDPMSolverMultistepScheduler)
from .fm_solvers_unipc import FlowUniPCMultistepScheduler from .fm_solvers_unipc import FlowUniPCMultistepScheduler
from .basic_flowmatch import FlowMatchScheduler from .basic_flowmatch import FlowMatchScheduler
from .flowmatch_pusa import FlowMatchSchedulerPusa from .flowmatch_pusa import FlowMatchSchedulerPusa
from .flowmatch_res_multistep import FlowMatchSchedulerResMultistep from .flowmatch_res_multistep import FlowMatchSchedulerResMultistep
from .ersde_scheduler import ERSDEScheduler
from .scheduling_flow_match_lcm import FlowMatchLCMScheduler from .scheduling_flow_match_lcm import FlowMatchLCMScheduler
from .fm_sa_ode import FlowMatchSAODEStableScheduler from .fm_sa_ode import FlowMatchSAODEStableScheduler
from .fm_rcm import rCMFlowMatchScheduler from .fm_rcm import rCMFlowMatchScheduler
from .vitb_unipc import ViBTScheduler
from ...utils import log from ...utils import log
try: try:
@@ -24,25 +27,35 @@ scheduler_list = [
"deis", "deis",
"lcm", "lcm/beta", "lcm", "lcm/beta",
"res_multistep", "res_multistep",
"er_sde",
"flowmatch_causvid", "flowmatch_causvid",
"flowmatch_distill", "flowmatch_distill",
"flowmatch_pusa", "flowmatch_pusa",
"multitalk", "multitalk",
"sa_ode_stable", "sa_ode_stable",
"rcm" "rcm",
"vibt_unipc",
] ]
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, flowedit_args=None, denoise_strength=1.0, sigmas=None, log_timesteps=False, **kwargs): def _apply_custom_sigmas(sample_scheduler, sigmas, device):
sample_scheduler.sigmas = sigmas.to(device)
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, denoise_strength=1.0, sigmas=None, log_timesteps=False, enhance_hf=False, **kwargs):
timesteps = None timesteps = None
if 'unipc' in scheduler: if sigmas is not None:
steps = len(sigmas) - 1
if scheduler == 'vibt_unipc':
sample_scheduler = ViBTScheduler()
sample_scheduler.set_parameters(shift=shift)
sample_scheduler.set_timesteps(steps, device=device)
elif 'unipc' in scheduler:
sample_scheduler = FlowUniPCMultistepScheduler(shift=shift) sample_scheduler = FlowUniPCMultistepScheduler(shift=shift)
if sigmas is None: if sigmas is None:
sample_scheduler.set_timesteps(steps, device=device, shift=shift, use_beta_sigmas=('beta' in scheduler)) sample_scheduler.set_timesteps(steps, device=device, shift=shift, use_beta_sigmas=('beta' in scheduler))
else: else:
sample_scheduler.sigmas = sigmas.to(device) _apply_custom_sigmas(sample_scheduler, sigmas, device)
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
elif scheduler in ['euler/beta', 'euler', 'longcat_distill_euler']: elif scheduler in ['euler/beta', 'euler', 'longcat_distill_euler']:
if 'longcat' in scheduler: if 'longcat' in scheduler:
num_distill_sample_steps = 50 num_distill_sample_steps = 50
@@ -57,10 +70,10 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas) sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas)
else: else:
sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta')) sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta'))
if flowedit_args: #seems to work better if sigmas is None:
timesteps, _ = retrieve_timesteps(sample_scheduler, device=device, sigmas=get_sampling_sigmas(steps, shift)) sample_scheduler.set_timesteps(steps, device=device)
else: else:
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].tolist() if sigmas is not None else None) _apply_custom_sigmas(sample_scheduler, sigmas, device)
elif 'dpm' in scheduler: elif 'dpm' in scheduler:
if 'sde' in scheduler: if 'sde' in scheduler:
algorithm_type = "sde-dpmsolver++" algorithm_type = "sde-dpmsolver++"
@@ -70,16 +83,20 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
if sigmas is None: if sigmas is None:
sample_scheduler.set_timesteps(steps, device=device, use_beta_sigmas=('beta' in scheduler)) sample_scheduler.set_timesteps(steps, device=device, use_beta_sigmas=('beta' in scheduler))
else: else:
sample_scheduler.sigmas = sigmas.to(device) _apply_custom_sigmas(sample_scheduler, sigmas, device)
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
elif scheduler == 'deis': elif scheduler == 'deis':
sample_scheduler = DEISMultistepScheduler(use_flow_sigmas=True, prediction_type="flow_prediction", flow_shift=shift) sample_scheduler = DEISMultistepScheduler(use_flow_sigmas=True, prediction_type="flow_prediction", flow_shift=shift)
sample_scheduler.set_timesteps(steps, device=device) if sigmas is None:
sample_scheduler.sigmas[-1] = 1e-6 sample_scheduler.set_timesteps(steps, device=device)
sample_scheduler.sigmas[-1] = 1e-6
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif 'lcm' in scheduler: elif 'lcm' in scheduler:
sample_scheduler = FlowMatchLCMScheduler(shift=shift, use_beta_sigmas=(scheduler == 'lcm/beta')) sample_scheduler = FlowMatchLCMScheduler(shift=shift, use_beta_sigmas=(scheduler == 'lcm/beta'))
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].tolist() if sigmas is not None else None) if sigmas is None:
sample_scheduler.set_timesteps(steps, device=device)
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif 'flowmatch_causvid' in scheduler: elif 'flowmatch_causvid' in scheduler:
if sigmas is not None: if sigmas is not None:
raise NotImplementedError("This scheduler does not support custom sigmas") raise NotImplementedError("This scheduler does not support custom sigmas")
@@ -112,21 +129,50 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)]) sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)])
elif 'flowmatch_pusa' in scheduler: elif 'flowmatch_pusa' in scheduler:
sample_scheduler = FlowMatchSchedulerPusa(shift=shift, sigma_min=0.0, extra_one_step=True) sample_scheduler = FlowMatchSchedulerPusa(shift=shift, sigma_min=0.0, extra_one_step=True)
sample_scheduler.set_timesteps(steps+1, denoising_strength=denoise_strength, shift=shift, if sigmas is None:
sigmas=sigmas[:-1].tolist() if sigmas is not None else None) sample_scheduler.set_timesteps(steps+1, denoising_strength=denoise_strength, shift=shift)
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif scheduler == 'res_multistep': elif scheduler == 'res_multistep':
sample_scheduler = FlowMatchSchedulerResMultistep(shift=shift) sample_scheduler = FlowMatchSchedulerResMultistep(shift=shift)
sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, sigmas=sigmas[:-1].tolist() if sigmas is not None else None) if sigmas is None:
sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength)
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif scheduler == 'er_sde':
sample_scheduler = ERSDEScheduler(shift=shift)
if sigmas is None:
sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength)
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif "sa_ode_stable" in scheduler: elif "sa_ode_stable" in scheduler:
sample_scheduler = FlowMatchSAODEStableScheduler(shift=shift, **kwargs) sample_scheduler = FlowMatchSAODEStableScheduler(shift=shift, **kwargs)
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].tolist() if sigmas is not None else None) if sigmas is None:
sample_scheduler.set_timesteps(steps, device=device)
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif 'rcm' in scheduler: elif 'rcm' in scheduler:
sample_scheduler = rCMFlowMatchScheduler() sample_scheduler = rCMFlowMatchScheduler()
sample_scheduler.set_timesteps(steps, sigma_max=120) if sigmas is None:
sample_scheduler.set_timesteps(steps, sigma_max=120)
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
if timesteps is None: if timesteps is None:
timesteps = sample_scheduler.timesteps timesteps = sample_scheduler.timesteps
if enhance_hf:
num_tail_uniform_steps = max(3, min(15, int(len(timesteps) * 0.2))) # Use 20% of steps for uniform tail (minimum 3, maximum 15)
tail_uniform_start = float(timesteps.max()) * 0.5 # Split at 50% of the timestep range
tail_uniform_end = 0
timesteps_uniform_tail = list(np.linspace(tail_uniform_start, tail_uniform_end, num_tail_uniform_steps, dtype=np.float32, endpoint=(tail_uniform_end != 0)))
timesteps_uniform_tail = [torch.tensor(t, device=device).unsqueeze(0) for t in timesteps_uniform_tail]
filtered_timesteps = [timestep.unsqueeze(0).to(device) for timestep in timesteps if timestep > tail_uniform_start]
timesteps = torch.cat(filtered_timesteps + timesteps_uniform_tail)
sample_scheduler.timesteps = timesteps
sample_scheduler.sigmas = torch.cat([timesteps / 1000, torch.zeros(1, device=timesteps.device)])
steps = len(timesteps) steps = len(timesteps)
if (isinstance(start_step, int) and end_step != -1 and start_step >= end_step) or (not isinstance(start_step, int) and start_step != -1 and end_step >= start_step): if (isinstance(start_step, int) and end_step != -1 and start_step >= end_step) or (not isinstance(start_step, int) and start_step != -1 and end_step >= start_step):
raise ValueError("start_step must be less than end_step") raise ValueError("start_step must be less than end_step")
@@ -136,7 +182,7 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
end_idx = len(timesteps) - 1 end_idx = len(timesteps) - 1
if log_timesteps: if log_timesteps:
log.info(f"------- Scheduler info -------") log.info("------- Scheduler info -------")
log.info(f"Total timesteps: {timesteps}") log.info(f"Total timesteps: {timesteps}")
if isinstance(start_step, float): if isinstance(start_step, float):
@@ -156,6 +202,7 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
end_idx = end_step - 1 end_idx = end_step - 1
# Slice timesteps and sigmas once, based on indices # Slice timesteps and sigmas once, based on indices
all_timesteps = timesteps
timesteps = timesteps[start_idx:end_idx+1] timesteps = timesteps[start_idx:end_idx+1]
sample_scheduler.full_sigmas = sample_scheduler.sigmas.clone() sample_scheduler.full_sigmas = sample_scheduler.sigmas.clone()
sample_scheduler.sigmas = sample_scheduler.sigmas[start_idx:start_idx+len(timesteps)+1] # always one longer sample_scheduler.sigmas = sample_scheduler.sigmas[start_idx:start_idx+len(timesteps)+1] # always one longer
@@ -163,9 +210,10 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
if log_timesteps: if log_timesteps:
log.info(f"Using timesteps: {timesteps}") log.info(f"Using timesteps: {timesteps}")
log.info(f"Using sigmas: {sample_scheduler.sigmas}") log.info(f"Using sigmas: {sample_scheduler.sigmas}")
log.info(f"------------------------------") log.info("------------------------------")
if hasattr(sample_scheduler, 'timesteps'): if hasattr(sample_scheduler, 'timesteps'):
sample_scheduler.timesteps = timesteps sample_scheduler.timesteps = timesteps
setattr(sample_scheduler, 'all_timesteps', all_timesteps)
return sample_scheduler, timesteps, start_idx, end_idx return sample_scheduler, timesteps, start_idx, end_idx
+154
View File
@@ -0,0 +1,154 @@
import torch
class ERSDEScheduler():
"""Extended Reverse-Time SDE solver (VP ER-SDE-Solver-3).
Based on: arXiv: https://arxiv.org/abs/2309.06169
Code reference: https://github.com/QinpengCui/ER-SDE-Solver/blob/main/er_sde_solver.py
"""
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0,
sigma_max=1.0, sigma_min=0.003 / 1.002, max_stage=3, s_noise=1.0,
num_integration_points=200):
self.num_train_timesteps = num_train_timesteps
self.shift = shift
self.sigma_max = sigma_max
self.sigma_min = sigma_min
self.max_stage = max_stage
self.s_noise = s_noise
self.num_integration_points = num_integration_points
self.set_timesteps(num_inference_steps)
self.old_denoised = None
self.old_denoised_d = None
self.step_index = 0
def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, sigmas=None):
"""Generate the full sigma schedule (from max to min)."""
full_sigmas = torch.linspace(self.sigma_max, self.sigma_min, self.num_train_timesteps)
ss = len(full_sigmas) / num_inference_steps
if sigmas is None:
sigmas = []
for x in range(num_inference_steps):
idx = int(round(x * ss))
sigmas.append(float(full_sigmas[idx]))
sigmas.append(0.0)
self.sigmas = torch.FloatTensor(sigmas)
self.sigmas = self.shift * self.sigmas / (1 + (self.shift - 1) * self.sigmas)
self.timesteps = self.sigmas[:-1] * self.num_train_timesteps
self.step_index = 0
self.old_denoised = None
self.old_denoised_d = None
def default_er_sde_noise_scaler(self, x):
return x * ((x ** 0.3).exp() + 10.0)
def step(self, model_output, timestep, sample, generator):
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
self.sigmas = self.sigmas.to(model_output.device)
self.timesteps = self.timesteps.to(model_output.device)
if timestep.ndim == 0:
timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0)
else:
timestep_id = torch.argmin((self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
noise_scaler = self.default_er_sde_noise_scaler
# Get current and next sigma
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
if (timestep_id + 1 >= len(self.sigmas)).any():
sigma_next = torch.zeros_like(sigma)
else:
sigma_next = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
er_lambda_s = sigma
er_lambda_t = sigma_next
# Calculate alpha values
alpha_s = sigma / (er_lambda_s + 1e-10)
alpha_t = sigma_next / (er_lambda_t + 1e-10)
r_alpha = alpha_t / (alpha_s + 1e-10)
# Denoised prediction (x_0 estimate)
denoised = sample - sigma * model_output
# Determine which stage to use
stage_used = min(self.max_stage, self.step_index + 1)
if sigma_next == 0 or (sigma_next == 0.0).all():
# Final step - return denoised
x = denoised
else:
r = noise_scaler(er_lambda_t) / (noise_scaler(er_lambda_s) + 1e-10)
# Stage 1: Euler step
x = r_alpha * r * sample + alpha_t * (1 - r) * denoised
if stage_used >= 2 and self.old_denoised is not None:
dt = er_lambda_t - er_lambda_s
lambda_step_size = -dt / self.num_integration_points
# Create integration points
point_indice = torch.arange(0, self.num_integration_points,
dtype=torch.float32, device=sample.device)
lambda_pos = er_lambda_t + point_indice * lambda_step_size
scaled_pos = noise_scaler(lambda_pos)
# Stage 2: Second-order correction
s = torch.sum(1 / (scaled_pos + 1e-10)) * lambda_step_size
# Get previous sigma for derivative calculation
if timestep_id > 0:
sigma_prev = self.sigmas[timestep_id - 1].reshape(-1, 1, 1, 1)
er_lambda_prev = sigma_prev
else:
er_lambda_prev = er_lambda_s
denoised_d = (denoised - self.old_denoised) / ((er_lambda_s - er_lambda_prev) + 1e-10)
x = x + alpha_t * (dt + s * noise_scaler(er_lambda_t)) * denoised_d
if stage_used >= 3 and self.old_denoised_d is not None:
# Stage 3: Third-order correction
s_u = torch.sum((lambda_pos - er_lambda_s) / (scaled_pos + 1e-10)) * lambda_step_size
# Get sigma from two steps ago
if timestep_id > 1:
sigma_prev_prev = self.sigmas[timestep_id - 2].reshape(-1, 1, 1, 1)
er_lambda_prev_prev = sigma_prev_prev
else:
er_lambda_prev_prev = er_lambda_prev
denoised_u = (denoised_d - self.old_denoised_d) / (((er_lambda_s - er_lambda_prev_prev) / 2) + 1e-10)
x = x + alpha_t * ((dt ** 2) / 2 + s_u * noise_scaler(er_lambda_t)) * denoised_u
self.old_denoised_d = denoised_d
# Add stochastic noise
if self.s_noise > 0:
noise_term = (er_lambda_t ** 2 - er_lambda_s ** 2 * r ** 2).sqrt()
noise_term = torch.nan_to_num(noise_term, nan=0.0)
noise = torch.randn(*x.shape, dtype=torch.float32, device=torch.device("cpu"), generator=generator).to(x)
x = x + alpha_t * noise * self.s_noise * noise_term
# Store current denoised for next iteration
self.old_denoised = denoised
self.step_index += 1
return x
def add_noise(self, original_samples, noise, timestep):
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
self.sigmas = self.sigmas.to(noise.device)
self.timesteps = self.timesteps.to(noise.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
sample = (1 - sigma) * original_samples + sigma * noise
return sample.type_as(noise)
@@ -35,9 +35,7 @@ class FlowMatchSchedulerResMultistep():
self.sigmas = torch.FloatTensor(sigmas) self.sigmas = torch.FloatTensor(sigmas)
self.sigmas = self.shift * self.sigmas / \ self.sigmas = self.shift * self.sigmas / \
(1 + (self.shift - 1) * self.sigmas) (1 + (self.shift - 1) * self.sigmas)
self.timesteps = self.sigmas * self.num_train_timesteps self.timesteps = self.sigmas[:-1] * self.num_train_timesteps
#print(f"Timesteps: {self.timesteps}, Sigmas: {self.sigmas}")
def step(self, model_output, timestep, sample): def step(self, model_output, timestep, sample):
if timestep.ndim == 2: if timestep.ndim == 2:
@@ -48,14 +46,14 @@ class FlowMatchSchedulerResMultistep():
timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0) timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0)
else: else:
timestep_id = torch.argmin((self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1) timestep_id = torch.argmin((self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1) sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
sigma_prev = self.sigmas[timestep_id - 1].reshape(-1, 1, 1, 1) if timestep_id > 0 else sigma sigma_prev = self.sigmas[timestep_id - 1].reshape(-1, 1, 1, 1) if timestep_id > 0 else sigma
if (timestep_id + 1 >= len(self.sigmas)).any(): if (timestep_id + 1 >= len(self.sigmas)).any():
sigma_next = torch.tensor(0) sigma_next = torch.tensor(0)
else: else:
sigma_next = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1) sigma_next = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
x0_pred = (sample - sigma * model_output) x0_pred = (sample - sigma * model_output)
if sigma_next == 0 or self.prev_model_output is None: if sigma_next == 0 or self.prev_model_output is None:
@@ -73,7 +71,7 @@ class FlowMatchSchedulerResMultistep():
self.old_sigma_next = sigma_next self.old_sigma_next = sigma_next
self.prev_model_output = x0_pred self.prev_model_output = x0_pred
return x return x
def add_noise(self, original_samples, noise, timestep): def add_noise(self, original_samples, noise, timestep):
""" """
+41
View File
@@ -0,0 +1,41 @@
from diffusers.schedulers import UniPCMultistepScheduler
import torch
class ViBTScheduler(UniPCMultistepScheduler):
def __init__(self, **kwargs):
super().__init__(**{**kwargs, "use_flow_sigmas": True})
self.set_parameters()
def set_parameters(self, noise_scale=1.0, shift=5.0, seed=None):
self.noise_scale = noise_scale
self.config.flow_shift = shift
def step(self, model_output, timestep, sample, generator, **kwargs):
delta_t = (
max(self.timesteps[self.timesteps < timestep]) - timestep
if any(self.timesteps < timestep)
else -timestep - 1
) / 1000
current_t = (timestep + 1) / 1000.0
eta = (-delta_t * (current_t + delta_t) / current_t) ** 0.5
noise = torch.randn(
sample.shape,
generator=generator,
device=torch.device("cpu"),
dtype=sample.dtype,
).to(sample.device)
latents = sample + delta_t * model_output + eta * self.noise_scale * noise
return (latents,)
@classmethod
def from_scheduler(
cls, scheduler: UniPCMultistepScheduler, noise_scale=1.0, shift_gamma=5.0
):
obj = cls.__new__(cls)
obj.__dict__ = scheduler.__dict__.copy()
obj.set_parameters(noise_scale, shift_gamma)
return obj
+145 -45
View File
@@ -5,6 +5,9 @@ import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from tqdm import tqdm from tqdm import tqdm
from comfy.utils import ProgressBar from comfy.utils import ProgressBar
from ..utils import print_memory, log
from comfy import model_management as mm
device = mm.get_torch_device()
import comfy.ops import comfy.ops
ops = comfy.ops.disable_weight_init ops = comfy.ops.disable_weight_init
@@ -254,10 +257,11 @@ class Resample38(Resample):
class ResidualBlock(nn.Module): class ResidualBlock(nn.Module):
def __init__(self, in_dim, out_dim, dropout=0.0): def __init__(self, in_dim, out_dim, dropout=0.0, cpu_cache=False):
super().__init__() super().__init__()
self.in_dim = in_dim self.in_dim = in_dim
self.out_dim = out_dim self.out_dim = out_dim
self.cpu_cache = cpu_cache
# layers # layers
self.residual = nn.Sequential( self.residual = nn.Sequential(
@@ -269,6 +273,12 @@ class ResidualBlock(nn.Module):
if in_dim != out_dim else nn.Identity() if in_dim != out_dim else nn.Identity()
def forward(self, x, feat_cache=None, feat_idx=[0]): def forward(self, x, feat_cache=None, feat_idx=[0]):
if self.cpu_cache:
return self._forward_cpu_cache(x, feat_cache, feat_idx)
else:
return self._forward(x, feat_cache, feat_idx)
def _forward(self, x, feat_cache=None, feat_idx=[0]):
h = self.shortcut(x) h = self.shortcut(x)
for layer in self.residual: for layer in self.residual:
if check_is_instance(layer, CausalConv3d) and feat_cache is not None: if check_is_instance(layer, CausalConv3d) and feat_cache is not None:
@@ -288,6 +298,26 @@ class ResidualBlock(nn.Module):
x = layer(x) x = layer(x)
return x + h return x + h
def _forward_cpu_cache(self, x, feat_cache=None, feat_idx=[0]):
h = self.shortcut(x)
for layer in self.residual:
if check_is_instance(layer, CausalConv3d) and feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
cached_frame = feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device)
cache_x = torch.cat([cached_frame, cache_x], dim=2)
prev_cache = feat_cache[idx].to(x.device) if feat_cache[idx] is not None else None
x = layer(x, prev_cache)
feat_cache[idx] = cache_x.to("cpu", non_blocking=True)
feat_idx[0] += 1
else:
x = layer(x)
return x + h
class AttentionBlock(nn.Module): class AttentionBlock(nn.Module):
""" """
@@ -507,7 +537,8 @@ class Encoder3d(nn.Module):
attn_scales=[], attn_scales=[],
temperal_downsample=[True, True, False], temperal_downsample=[True, True, False],
dropout=0.0, dropout=0.0,
pruning_rate=0.0): pruning_rate=0.0,
cpu_cache=False):
super().__init__() super().__init__()
self.dim = dim self.dim = dim
self.z_dim = z_dim self.z_dim = z_dim
@@ -515,6 +546,7 @@ class Encoder3d(nn.Module):
self.num_res_blocks = num_res_blocks self.num_res_blocks = num_res_blocks
self.attn_scales = attn_scales self.attn_scales = attn_scales
self.temperal_downsample = temperal_downsample self.temperal_downsample = temperal_downsample
self.cpu_cache = cpu_cache
# dimensions # dimensions
dims = [dim * u for u in [1] + dim_mult] dims = [dim * u for u in [1] + dim_mult]
@@ -529,7 +561,7 @@ class Encoder3d(nn.Module):
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])): for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
# residual (+attention) blocks # residual (+attention) blocks
for _ in range(num_res_blocks): for _ in range(num_res_blocks):
downsamples.append(ResidualBlock(in_dim, out_dim, dropout)) downsamples.append(ResidualBlock(in_dim, out_dim, dropout, cpu_cache=cpu_cache))
if scale in attn_scales: if scale in attn_scales:
downsamples.append(AttentionBlock(out_dim)) downsamples.append(AttentionBlock(out_dim))
in_dim = out_dim in_dim = out_dim
@@ -543,9 +575,9 @@ class Encoder3d(nn.Module):
self.downsamples = nn.Sequential(*downsamples) self.downsamples = nn.Sequential(*downsamples)
# middle blocks # middle blocks
self.middle = nn.Sequential(ResidualBlock(out_dim, out_dim, dropout), self.middle = nn.Sequential(ResidualBlock(out_dim, out_dim, dropout, cpu_cache=cpu_cache),
AttentionBlock(out_dim), AttentionBlock(out_dim),
ResidualBlock(out_dim, out_dim, dropout)) ResidualBlock(out_dim, out_dim, dropout, cpu_cache=cpu_cache))
# output blocks # output blocks
self.head = nn.Sequential(RMS_norm(out_dim, images=False), nn.SiLU(), self.head = nn.Sequential(RMS_norm(out_dim, images=False), nn.SiLU(),
@@ -612,7 +644,8 @@ class Encoder3d_38(nn.Module):
attn_scales=[], attn_scales=[],
temperal_downsample=[False, True, True], temperal_downsample=[False, True, True],
dropout=0.0, dropout=0.0,
pruning_rate=0.0): pruning_rate=0.0,
cpu_cache=False):
super().__init__() super().__init__()
self.dim = dim self.dim = dim
self.z_dim = z_dim self.z_dim = z_dim
@@ -620,6 +653,7 @@ class Encoder3d_38(nn.Module):
self.num_res_blocks = num_res_blocks self.num_res_blocks = num_res_blocks
self.attn_scales = attn_scales self.attn_scales = attn_scales
self.temperal_downsample = temperal_downsample self.temperal_downsample = temperal_downsample
self.cpu_cache = cpu_cache
# dimensions # dimensions
dims = [dim * u for u in [1] + dim_mult] dims = [dim * u for u in [1] + dim_mult]
@@ -650,9 +684,9 @@ class Encoder3d_38(nn.Module):
# middle blocks # middle blocks
self.middle = nn.Sequential( self.middle = nn.Sequential(
ResidualBlock(out_dim, out_dim, dropout), ResidualBlock(out_dim, out_dim, dropout, cpu_cache=cpu_cache),
AttentionBlock(out_dim), AttentionBlock(out_dim),
ResidualBlock(out_dim, out_dim, dropout), ResidualBlock(out_dim, out_dim, dropout, cpu_cache=cpu_cache),
) )
# # output blocks # # output blocks
@@ -730,7 +764,8 @@ class Decoder3d(nn.Module):
attn_scales=[], attn_scales=[],
temperal_upsample=[False, True, True], temperal_upsample=[False, True, True],
dropout=0.0, dropout=0.0,
pruning_rate=0.0): pruning_rate=0.0,
cpu_cache=False):
super().__init__() super().__init__()
self.dim = dim self.dim = dim
self.z_dim = z_dim self.z_dim = z_dim
@@ -738,6 +773,7 @@ class Decoder3d(nn.Module):
self.num_res_blocks = num_res_blocks self.num_res_blocks = num_res_blocks
self.attn_scales = attn_scales self.attn_scales = attn_scales
self.temperal_upsample = temperal_upsample self.temperal_upsample = temperal_upsample
self.cpu_cache = cpu_cache
# dimensions # dimensions
dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]] dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
@@ -748,9 +784,9 @@ class Decoder3d(nn.Module):
self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1) self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
# middle blocks # middle blocks
self.middle = nn.Sequential(ResidualBlock(dims[0], dims[0], dropout), self.middle = nn.Sequential(ResidualBlock(dims[0], dims[0], dropout, cpu_cache=cpu_cache),
AttentionBlock(dims[0]), AttentionBlock(dims[0]),
ResidualBlock(dims[0], dims[0], dropout)) ResidualBlock(dims[0], dims[0], dropout, cpu_cache=cpu_cache))
# upsample blocks # upsample blocks
upsamples = [] upsamples = []
@@ -759,7 +795,7 @@ class Decoder3d(nn.Module):
if i == 1 or i == 2 or i == 3: if i == 1 or i == 2 or i == 3:
in_dim = in_dim // 2 in_dim = in_dim // 2
for _ in range(num_res_blocks + 1): for _ in range(num_res_blocks + 1):
upsamples.append(ResidualBlock(in_dim, out_dim, dropout)) upsamples.append(ResidualBlock(in_dim, out_dim, dropout, cpu_cache=cpu_cache))
if scale in attn_scales: if scale in attn_scales:
upsamples.append(AttentionBlock(out_dim)) upsamples.append(AttentionBlock(out_dim))
in_dim = out_dim in_dim = out_dim
@@ -838,7 +874,8 @@ class Decoder3d_38(nn.Module):
attn_scales=[], attn_scales=[],
temperal_upsample=[False, True, True], temperal_upsample=[False, True, True],
dropout=0.0, dropout=0.0,
pruning_rate=0.0): pruning_rate=0.0,
cpu_cache=False):
super().__init__() super().__init__()
self.dim = dim self.dim = dim
self.z_dim = z_dim self.z_dim = z_dim
@@ -855,9 +892,9 @@ class Decoder3d_38(nn.Module):
self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1) self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
# middle blocks # middle blocks
self.middle = nn.Sequential(ResidualBlock(dims[0], dims[0], dropout), self.middle = nn.Sequential(ResidualBlock(dims[0], dims[0], dropout, cpu_cache=cpu_cache),
AttentionBlock(dims[0]), AttentionBlock(dims[0]),
ResidualBlock(dims[0], dims[0], dropout)) ResidualBlock(dims[0], dims[0], dropout, cpu_cache=cpu_cache))
# upsample blocks # upsample blocks
upsamples = [] upsamples = []
@@ -951,7 +988,9 @@ class VideoVAE_(nn.Module):
dropout=0.0, dropout=0.0,
mean=None, mean=None,
inv_std=None, inv_std=None,
pruning_rate=0.0): pruning_rate=0.0,
cpu_cache=False,
verbose=False):
super().__init__() super().__init__()
self.dim = dim self.dim = dim
self.z_dim = z_dim self.z_dim = z_dim
@@ -962,14 +1001,15 @@ class VideoVAE_(nn.Module):
self.temperal_upsample = temperal_downsample[::-1] self.temperal_upsample = temperal_downsample[::-1]
self.mean = mean self.mean = mean
self.inv_std = inv_std self.inv_std = inv_std
self.verbose = verbose
# modules # modules
self.encoder = Encoder3d(dim, z_dim * 2, dim_mult, num_res_blocks, self.encoder = Encoder3d(dim, z_dim * 2, dim_mult, num_res_blocks,
attn_scales, self.temperal_downsample, dropout, pruning_rate) attn_scales, self.temperal_downsample, dropout, pruning_rate, cpu_cache=cpu_cache)
self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1) self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1)
self.conv2 = CausalConv3d(z_dim, z_dim, 1) self.conv2 = CausalConv3d(z_dim, z_dim, 1)
self.decoder = Decoder3d(dim, z_dim, dim_mult, num_res_blocks, self.decoder = Decoder3d(dim, z_dim, dim_mult, num_res_blocks,
attn_scales, self.temperal_upsample, dropout, pruning_rate) attn_scales, self.temperal_upsample, dropout, pruning_rate, cpu_cache=cpu_cache)
def forward(self, x): def forward(self, x):
mu, log_var = self.encode(x) mu, log_var = self.encode(x)
@@ -1016,10 +1056,15 @@ class VideoVAE_(nn.Module):
def encode(self, x, pbar=True, sample=False): def encode(self, x, pbar=True, sample=False):
t = x.shape[2] t = x.shape[2]
iter_ = 1 + (t - 1) // 4 iter_ = 1 + (t - 1) // 4
input_shape = x.shape
if pbar: if pbar:
pbar = ProgressBar(iter_) pbar = ProgressBar(iter_)
try:
torch.cuda.reset_peak_memory_stats(device)
except Exception:
pass
for i in range(iter_): for i in tqdm(range(iter_), desc="WanVAE encoding frames", disable=not pbar):
self._enc_conv_idx = [0] self._enc_conv_idx = [0]
if i == 0: if i == 0:
out = self.encoder(x[:, :, :1, :, :], out = self.encoder(x[:, :, :1, :, :],
@@ -1042,15 +1087,22 @@ class VideoVAE_(nn.Module):
std = torch.exp(0.5 * log_var.clamp(-30.0, 20.0)) std = torch.exp(0.5 * log_var.clamp(-30.0, 20.0))
eps = torch.randn_like(std) eps = torch.randn_like(std)
return mu + std * eps return mu + std * eps
if self.verbose:
try:
log.info(f"WanVAE encoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE encode")
torch.cuda.reset_peak_memory_stats(device)
except Exception:
pass
return mu return mu
#modification originally by @raindrop313 https://github.com/raindrop313/ComfyUI-WanVideoStartEndFrames #modification originally by @raindrop313 https://github.com/raindrop313/ComfyUI-WanVideoStartEndFrames
def decode_2(self, z): def decode_2(self, z):
# z: [b,c,t,h,w] # z: [b,c,t,h,w]
z = z / self.inv_std.to(z) + self.mean.to(z) z = z / self.inv_std.to(z) + self.mean.to(z)
iter_ = z.shape[2] iter_ = z.shape[2]
z_head=z[:,:,:-1,:,:] z_head=z[:,:,:-1,:,:]
z_tail=z[:,:,-1,:,:].unsqueeze(2) z_tail=z[:,:,-1,:,:].unsqueeze(2)
@@ -1065,12 +1117,12 @@ class VideoVAE_(nn.Module):
out_ = self.decoder(x[:, :, -1, :, :].unsqueeze(2), out_ = self.decoder(x[:, :, -1, :, :].unsqueeze(2),
feat_cache=None, feat_cache=None,
feat_idx=self._conv_idx) feat_idx=self._conv_idx)
out = torch.cat([out, out_], 2) # may add tensor offload out = torch.cat([out, out_], 2)
else: else:
out_ = self.decoder(x[:, :, i:i + 1, :, :], out_ = self.decoder(x[:, :, i:i + 1, :, :],
feat_cache=self._feat_map, feat_cache=self._feat_map,
feat_idx=self._conv_idx) feat_idx=self._conv_idx)
out = torch.cat([out, out_], 2) # may add tensor offload out = torch.cat([out, out_], 2)
self.clear_cache() self.clear_cache()
return out return out
@@ -1079,11 +1131,16 @@ class VideoVAE_(nn.Module):
def decode(self, z, pbar=True): def decode(self, z, pbar=True):
# z: [b,c,t,h,w] # z: [b,c,t,h,w]
z = z / self.inv_std.to(z) + self.mean.to(z) z = z / self.inv_std.to(z) + self.mean.to(z)
input_shape = z.shape
iter_ = z.shape[2] iter_ = z.shape[2]
if pbar: if pbar:
pbar = ProgressBar(iter_) pbar = ProgressBar(iter_)
try:
torch.cuda.reset_peak_memory_stats(device)
except Exception:
pass
x = self.conv2(z) x = self.conv2(z)
for i in range(iter_): for i in tqdm(range(iter_), desc="WanVAE decoding frames", disable=not pbar):
self._conv_idx = [0] self._conv_idx = [0]
if i == 0: if i == 0:
out = self.decoder(x[:, :, i:i + 1, :, :], out = self.decoder(x[:, :, i:i + 1, :, :],
@@ -1093,13 +1150,20 @@ class VideoVAE_(nn.Module):
out_ = self.decoder(x[:, :, i:i + 1, :, :], out_ = self.decoder(x[:, :, i:i + 1, :, :],
feat_cache=self._feat_map, feat_cache=self._feat_map,
feat_idx=self._conv_idx) feat_idx=self._conv_idx)
out = torch.cat([out, out_], 2) # may add tensor offload out = torch.cat([out, out_], 2)
if pbar: if pbar:
pbar.update(1) pbar.update(1)
if pbar: if pbar:
pbar.update_absolute(0) pbar.update_absolute(0)
self.clear_cache() self.clear_cache()
if self.verbose:
try:
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device)
except Exception:
pass
return out return out
def reparameterize(self, mu, log_var): def reparameterize(self, mu, log_var):
@@ -1126,11 +1190,12 @@ class VideoVAE_(nn.Module):
class WanVideoVAE(nn.Module): class WanVideoVAE(nn.Module):
def __init__(self, z_dim=16, dtype=torch.float32, pruning_rate=0.0): def __init__(self, z_dim=16, dtype=torch.float32, pruning_rate=0.0, cpu_cache=False, verbose=False):
super().__init__() super().__init__()
self.dtype = dtype self.dtype = dtype
self.cpu_cache = cpu_cache
self.verbose = verbose
mean = [ mean = [
-0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508, -0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508,
0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921 0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921
@@ -1144,7 +1209,7 @@ class WanVideoVAE(nn.Module):
self.z_dim = z_dim self.z_dim = z_dim
# init model # init model
self.model = VideoVAE_(z_dim=z_dim, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate).eval().requires_grad_(False) self.model = VideoVAE_(z_dim=z_dim, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=self.cpu_cache, verbose=self.verbose).eval().requires_grad_(False)
self.upsampling_factor = 8 self.upsampling_factor = 8
@@ -1170,7 +1235,7 @@ class WanVideoVAE(nn.Module):
return mask return mask
def tiled_decode(self, hidden_states, device, tile_size, tile_stride, pbar=True): def tiled_decode(self, hidden_states, device, tile_size, tile_stride, end_=False, pbar=True):
_, _, T, H, W = hidden_states.shape _, _, T, H, W = hidden_states.shape
size_h, size_w = tile_size size_h, size_w = tile_size
stride_h, stride_w = tile_stride stride_h, stride_w = tile_stride
@@ -1187,14 +1252,20 @@ class WanVideoVAE(nn.Module):
data_device = "cpu" data_device = "cpu"
computation_device = device computation_device = device
out_T = T * 4 - 3 weight, values = None, None
weight = torch.zeros((1, 1, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
values = torch.zeros((1, 3, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
if pbar: if pbar:
pbar = ProgressBar(len(tasks)) pbar = ProgressBar(len(tasks))
for h, h_, w, w_ in tqdm(tasks, desc="VAE decoding"): for h, h_, w, w_ in tqdm(tasks, desc="VAE decoding"):
hidden_states_batch = hidden_states[:, :, :, h:h_, w:w_].to(computation_device) hidden_states_batch = hidden_states[:, :, :, h:h_, w:w_].to(computation_device)
hidden_states_batch = self.model.decode(hidden_states_batch).to(data_device) if end_:
hidden_states_batch = self.model.decode_2(hidden_states_batch).to(data_device)
else:
hidden_states_batch = self.model.decode(hidden_states_batch).to(data_device)
if weight is None:
weight = torch.zeros((1, 1, hidden_states_batch.shape[2], H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
if values is None:
values = torch.zeros((1, 3, hidden_states_batch.shape[2], H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
mask = self.build_mask( mask = self.build_mask(
hidden_states_batch, hidden_states_batch,
@@ -1227,7 +1298,7 @@ class WanVideoVAE(nn.Module):
def tiled_encode(self, video, device, tile_size, tile_stride, end_=False, pbar=True): def tiled_encode(self, video, device, tile_size, tile_stride, end_=False, pbar=True):
_, _, T, H, W = video.shape _, _, T, H, W = video.shape
if tile_size is None and tile_stride is None: if tile_size is None and tile_stride is None:
size_h, size_w = H //2, W // 2 size_h, size_w = H //2, W // 2
stride_h, stride_w = size_h // 2, size_w // 2 stride_h, stride_w = size_h // 2, size_w // 2
@@ -1338,7 +1409,7 @@ class WanVideoVAE(nn.Module):
for hidden_state in hidden_states: for hidden_state in hidden_states:
hidden_state = hidden_state.unsqueeze(0) hidden_state = hidden_state.unsqueeze(0)
if tiled: if tiled:
video = self.tiled_decode(hidden_state, device, tile_size, tile_stride, pbar=pbar) video = self.tiled_decode(hidden_state, device, tile_size, tile_stride, end_=end_, pbar=pbar)
else: else:
if end_: if end_:
video = self.double_decode(hidden_state, device) video = self.double_decode(hidden_state, device)
@@ -1363,7 +1434,9 @@ class VideoVAE38_(VideoVAE_):
dtype=torch.bfloat16, dtype=torch.bfloat16,
mean=None, mean=None,
inv_std=None, inv_std=None,
pruning_rate=0.0): pruning_rate=0.0,
cpu_cache=False,
verbose=False):
super(VideoVAE_, self).__init__() super(VideoVAE_, self).__init__()
self.dim = dim self.dim = dim
self.z_dim = z_dim self.z_dim = z_dim
@@ -1375,24 +1448,30 @@ class VideoVAE38_(VideoVAE_):
self.dtype = dtype self.dtype = dtype
self.mean = mean self.mean = mean
self.inv_std = inv_std self.inv_std = inv_std
self.cpu_cache = cpu_cache
self.verbose = verbose
# modules # modules
self.encoder = Encoder3d_38(dim, z_dim * 2, dim_mult, num_res_blocks, self.encoder = Encoder3d_38(dim, z_dim * 2, dim_mult, num_res_blocks,
attn_scales, self.temperal_downsample, dropout, pruning_rate) attn_scales, self.temperal_downsample, dropout, pruning_rate, cpu_cache=cpu_cache)
self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1) self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1)
self.conv2 = CausalConv3d(z_dim, z_dim, 1) self.conv2 = CausalConv3d(z_dim, z_dim, 1)
self.decoder = Decoder3d_38(dec_dim, z_dim, dim_mult, num_res_blocks, self.decoder = Decoder3d_38(dec_dim, z_dim, dim_mult, num_res_blocks,
attn_scales, self.temperal_upsample, dropout, pruning_rate) attn_scales, self.temperal_upsample, dropout, pruning_rate, cpu_cache=cpu_cache)
def encode(self, x, pbar=True, sample=False): def encode(self, x, pbar=True, sample=False):
input_shape = x.shape
self.clear_cache() self.clear_cache()
try:
torch.cuda.reset_peak_memory_stats(device)
except Exception:
pass
x = patchify(x, patch_size=2) x = patchify(x, patch_size=2)
t = x.shape[2] t = x.shape[2]
iter_ = 1 + (t - 1) // 4 iter_ = 1 + (t - 1) // 4
if pbar: if pbar:
pbar = ProgressBar(iter_) pbar = ProgressBar(iter_)
for i in range(iter_): for i in tqdm(range(iter_), desc="WanVAE encoding frames", disable=not pbar):
self._enc_conv_idx = [0] self._enc_conv_idx = [0]
if i == 0: if i == 0:
out = self.encoder(x[:, :, :1, :, :], out = self.encoder(x[:, :, :1, :, :],
@@ -1408,18 +1487,30 @@ class VideoVAE38_(VideoVAE_):
mu = self.conv1(out).chunk(2, dim=1)[0] mu = self.conv1(out).chunk(2, dim=1)[0]
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu) mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
self.clear_cache() self.clear_cache()
if self.verbose:
try:
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device)
except Exception:
pass
return mu return mu
def decode(self, z, pbar=True): def decode(self, z, pbar=True):
self.clear_cache() self.clear_cache()
input_shape = z.shape
try:
torch.cuda.reset_peak_memory_stats(device)
except Exception:
pass
z = z / self.inv_std.to(z) + self.mean.to(z) z = z / self.inv_std.to(z) + self.mean.to(z)
iter_ = z.shape[2] iter_ = z.shape[2]
if pbar: if pbar:
pbar = ProgressBar(iter_) pbar = ProgressBar(iter_)
x = self.conv2(z) x = self.conv2(z)
for i in range(iter_): for i in tqdm(range(iter_), desc="WanVAE decoding frames", disable=not pbar):
self._conv_idx = [0] self._conv_idx = [0]
if i == 0: if i == 0:
out = self.decoder(x[:, :, i:i + 1, :, :], out = self.decoder(x[:, :, i:i + 1, :, :],
@@ -1435,12 +1526,19 @@ class VideoVAE38_(VideoVAE_):
pbar.update(1) pbar.update(1)
out = unpatchify(out, patch_size=2) out = unpatchify(out, patch_size=2)
self.clear_cache() self.clear_cache()
if self.verbose:
try:
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device)
except Exception:
pass
return out return out
class WanVideoVAE38(WanVideoVAE): class WanVideoVAE38(WanVideoVAE):
def __init__(self, z_dim=48, dim=160, dtype=torch.bfloat16, pruning_rate=0.0): def __init__(self, z_dim=48, dim=160, dtype=torch.bfloat16, pruning_rate=0.0, cpu_cache=False, verbose=False):
super(WanVideoVAE, self).__init__() super(WanVideoVAE, self).__init__()
mean = [ mean = [
@@ -1463,7 +1561,9 @@ class WanVideoVAE38(WanVideoVAE):
self.inv_std = (1.0 / torch.tensor(std)).view(1, z_dim, 1, 1, 1) self.inv_std = (1.0 / torch.tensor(std)).view(1, z_dim, 1, 1, 1)
self.dtype = dtype self.dtype = dtype
self.z_dim = z_dim self.z_dim = z_dim
self.cpu_cache = cpu_cache
self.verbose = verbose
# init model # init model
self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate).eval().requires_grad_(False) self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=cpu_cache, verbose=verbose).eval().requires_grad_(False)
self.upsampling_factor = 16 self.upsampling_factor = 16