983 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
kijai 4e31081262 Better errors when trying to load models that don't exist 2025-11-13 16:19:05 +02:00
kijai ff26836cab Create wanvideo_2_2_5B_Ovi_image_to_video_audio_10_seconds_example_01.json 2025-11-11 18:53:11 +02:00
kijai 22037243ab Fix Ovi audio negative prompt
Had rather bad bug here which made Ovi audio always use the video negative prompt...
2025-11-11 17:47:33 +02:00
kijai e926f7a069 version bump 1.3.9 2025-11-11 10:57:46 +02:00
kijai e01e34da1f Update nodes_model_loading.py 2025-11-11 10:01:06 +02:00
kijai 47514f678d Allow loading original Ovi -models 2025-11-11 09:46:13 +02:00
kijai de3c9c895a Create wanvideo_1_3B_UniLumos_relight_example_01.json 2025-11-10 19:41:22 +02:00
kijai 4576ddb35e Add node to create input for UniLumos 2025-11-10 19:41:19 +02:00
kijai 68392684b5 Add node to use UniLumos
Simply allows fore and background latent inputs for UniLumos relight model, example inputs seem to work: https://github.com/alibaba-damo-academy/Lumos-Custom/tree/main/UniLumos/UniLumos/examples
2025-11-10 19:06:45 +02:00
Jukka Seppänen d3f33a9f09 Update readme.md 2025-11-06 16:38:46 +02:00
kijai d0ef3b5601 Update readme.md 2025-11-06 16:37:50 +02:00
kijai 475f96aede Fix accidental positional arg 2025-11-04 20:53:46 +02:00
kijai 1d0516a2a9 Avoid graph break for LongCat 2025-11-04 10:26:37 +02:00
kijai 8002d8a2f9 This still needed for some reason too 2025-11-04 09:58:40 +02:00
kijai 9a588a42ec Fix some precision issues with unmerged lora 2025-11-04 09:47:38 +02:00
kijai 509d6922f5 Update custom_linear.py 2025-11-04 09:44:07 +02:00
kijai 9fa4140159 Make lora torch.compile optional for unmerged lora application
This change has caused issues especially with LoRAs that have dynamic rank. Will now be disabled by default, to allow full graph with unmerged LoRAs the option to allow compile is available in the Torch Compile Settings -node
2025-11-04 01:34:51 +02:00
kijai 8ce6916d72 Fix for some cases of using comfy_chunked rope 2025-11-03 10:36:55 +02:00
kijai 0d0d28569a Fix cases where text encoder isn't used (eg. Minimax remover) 2025-11-03 10:29:10 +02:00
kijai 75109fdb79 Fix custom sigmas with euler 2025-11-03 10:12:05 +02:00
kijai 5eae7087fa Fix for S2V 2025-11-02 01:24:50 +02:00
kijai 393fe78ec2 Update model.py 2025-10-31 23:39:55 +02:00
kijai 5f4020b12d Fix a possible issue with Ovi audio model loading 2025-10-31 17:24:24 +02:00
kijai 5da8a6b169 Fix MultiTalk on some models 2025-10-31 16:50:00 +02:00
kijai ce6e7b501d Fix unmerged LoRA application for certain LoRAs 2025-10-30 23:28:38 +02:00
kijai 366f740d28 Update readme.md 2025-10-30 23:10:06 +02:00
Jukka Seppänen d45fe1ee22 Add note about blocking new accounts from posting issues
Added a note regarding issue posting restrictions due to bot activity.
2025-10-30 18:19:45 +02:00
kijai da24890d53 Update pyproject.toml 2025-10-30 17:56:20 +02:00
kijai 95391f403d Update nodes_sampler.py 2025-10-30 17:55:55 +02:00
kijai 9e0b3afe4e version checkpoint 2025-10-30 17:53:29 +02:00
kijai ba1beba982 Create LongCat_TI2V_example_01.json 2025-10-30 17:51:52 +02:00
kijai 64c195167b Update nodes.py 2025-10-30 17:33:42 +02:00
kijai cc9bf1e4f5 Store lora diffs in buffers for GGUF as well 2025-10-30 16:44:03 +02:00
kijai a64f115d35 Fix to previous 2025-10-29 11:06:38 +02:00
kijai e45f6f2fc4 Allow WanVideoScheduler -node to work with the looping samplers
Was broken for Multitalk/WanAnimate/S2V
2025-10-29 10:41:33 +02:00
kijai 1cd8df5c00 Update custom_linear.py 2025-10-29 02:50:11 +02:00
kijai d2614a9a49 Merge branch 'main' into longcat 2025-10-29 02:33:37 +02:00
kijai 083a8458c4 Register lora diffs as buffers to allow them to work with block swap
unmerged loras (non GGUF for now) will now be moved with block swap instead of always loaded from cpu to reduce device transfers and allow torch compile full graph
2025-10-29 02:33:26 +02:00
kijai 1c2f17e8d7 Add utility node to split sampler from settings
For cleaner previews
2025-10-29 02:24:24 +02:00
kijai 9d45b9f0de Use comfy core Conv3D workaround for VAE rather than the fp32 cast 2025-10-29 02:23:49 +02:00
Jukka Seppänen 833c6f50c7 Merge pull request #1581 from chengzeyi/fix-ref-conv-dtype-mismatch
Fix dtype mismatch in ref_conv forward pass
2025-10-28 14:27:43 +02:00
chengzeyiandClaude d15cf3001f Fix dtype mismatch in ref_conv forward pass
This commit fixes a RuntimeError that occurs when using Fun-Control
reference images: "Input type (float) and bias type (c10::Half)
should be the same"

Root cause:
- Commit 1ba1a16 changed the dtype handling strategy to convert
  the main latent `x` to `base_dtype` instead of converting
  embeddings to match `x.dtype`
- This caused `fun_ref` input to be in a different dtype than
  the `ref_conv` layer's weights and bias
- Line 2324 already handles this correctly for `attn_cond` by
  converting to `self.attn_conv_in.weight.dtype`

Solution:
- Convert `fun_ref` to match `self.ref_conv.weight.dtype` before
  passing through the convolution layer
- This follows the same pattern used for `attn_cond` on line 2324

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude <noreply@anthropic.com>
2025-10-28 12:08:23 +00:00
kijai eebbcd5ee0 Update model.py 2025-10-28 02:04:50 +02:00
kijai 2633119505 Update custom_linear.py 2025-10-28 01:55:25 +02:00
kijai 90908df260 Update model.py 2025-10-28 01:54:42 +02:00
kijai c80a488f70 Use fp32 norms for other models too and other fixes 2025-10-28 01:52:48 +02:00
kijai e69e068b57 Update gguf.py 2025-10-27 21:02:59 +02:00
kijai 54c45500b0 Apply lora diffs with unmerged loras too 2025-10-27 21:01:45 +02:00
kijai e560366600 Update model.py 2025-10-27 18:55:30 +02:00
kijai f880b321c6 Allow compile in lora application 2025-10-27 01:34:32 +02:00
kijai 51fcbd6b3d Revert "Allow compile here"
This reverts commit f583b56878.
2025-10-27 01:07:03 +02:00
kijai f583b56878 Allow compile here 2025-10-27 01:06:39 +02:00
kijai c59e52ca44 Precision adjustments 2025-10-27 00:23:32 +02:00
kijai a0bdf20817 Some cleanup and allow full block swap 2025-10-26 23:05:17 +02:00
kijai 8ad7e50f33 Fix cross attention split point 2025-10-26 22:07:19 +02:00
kijai d504c96174 Separate attention for input images like in original 2025-10-26 19:14:47 +02:00
kijai 43acf83adb Update model.py 2025-10-26 16:57:25 +02:00
kijai fb00932cad Init
https://huggingface.co/Kijai/LongCat-Video_comfy/tree/main
2025-10-26 16:36:20 +02:00
wzxysf 6a37c0b2d6 Enable vae tiling with end frame 2025-10-24 22:23:31 +08:00
kijai d74cfc54e8 Don't zero the extra frames
Possible fix for SVI "shot" method
2025-10-24 15:28:35 +03:00
kijai cfa883767f Update utils.py 2025-10-24 12:50:49 +03:00
kijai 88a60d71ab Fix 2025-10-23 23:25:20 +03:00
kijai b3ad381a65 not all torch versions have this 2025-10-23 21:12:45 +03:00
kijai 41168b1e82 Support light VAE
https://huggingface.co/lightx2v/Autoencoders/tree/main
2025-10-23 12:31:55 +03:00
kijai 7251c996d2 Load these LoRA keys too
Even though doesn't seem to do anything? At least silences the errors
2025-10-22 20:43:45 +03:00
kijai 67fcf0ba52 Reduce needless torch.compile recompiles 2025-10-22 13:29:13 +03:00
kijai aa610cde2b VACE: Support per context window reference images 2025-10-22 11:24:30 +03:00
kijai 6b286552b7 MocHa: remove possibly alpha channel from inputs 2025-10-21 23:26:17 +03:00
kijai 2ea824b9ab Mocha: Fix context window slicing 2025-10-21 20:33:03 +03:00
kijai 001040abc7 Update nodes.py 2025-10-21 20:27:37 +03:00
kijai 746d46e137 WanVideoScheduler: Always show split step sigma 2025-10-21 20:14:54 +03:00
kijai 089329ef2e MoCha: revert some RoPE changes 2025-10-21 20:12:40 +03: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
kijai 5c1f64197e Check for actual scale_weights even if the model for some reason doesn't have the scaled_fp8 key 2025-10-21 19:40:50 +03:00
kijai 54e938cd70 MoCha: Fix context window mask 2025-10-21 19:40:01 +03:00
kijai 7495db7669 Update nodes_sampler.py 2025-10-21 19:15:57 +03:00
kijai 4e36aee658 MoCha: experimental context windows support 2025-10-21 18:32:40 +03:00
kijai 56120d633e Update nodes.py 2025-10-21 17:59:37 +03:00
kijai 1f0861b649 MoCha: modify RoPE function to be more torch.compile friendly 2025-10-21 17:54:25 +03:00
kijai 7916f89c33 Add updated MoCha example for lower VRAM 2025-10-21 17:27:41 +03:00
Jukka Seppänen e294f417c0 Merge pull request #1501 from unrealMJ/mocha
[Feature] Request to add MoCha: End-to-End Video Character Replacement without Structural Guidance
2025-10-21 17:17:21 +03:00
unrealMJ ee011e66e9 update workflow 2025-10-21 17:28:50 +08:00
unrealMJ d7fc563581 add mocha workflow 2025-10-21 10:00:51 +08:00
unrealMJ 88defbfdd1 add MoCha 2025-10-21 09:40:54 +08:00
Jukka Seppänen 74f33df658 Merge pull request #1493 from HM-RunningHub/fix-uni3c-context-options
Fix Uni3C + Context Options compatibility (Issue #1491)
2025-10-20 22:33:38 +03:00
wenjian 487c400e8e Fix Uni3C + Context Options compatibility issue (Issue #1491)
Add uni3c_data parameter to predict_with_cfg() call in context windowing loop.

Without this parameter, Uni3C camera effects were not working when using
Context Options for long video generation. This makes the context windowing
behavior consistent with multitalk and wananimate sampling modes.
2025-10-21 02:31:47 +08:00
kijai 200f6943e3 Add sageattn mode that allows torch.compile
Latest wheel from woct0rdho includes the torch.compile fix:

https://github.com/woct0rdho/SageAttention/releases

Based on my quick testing this reduces peak VRAM usage a bit when running sageattn + torch.compile
2025-10-20 15:16:43 +03:00
kijai 8081e1337c Fix downcasting for fp8_fast and add error to indicate you can't do scaled downcast currently 2025-10-20 13:46:05 +03:00
kijai 9ecb80a58b Support loading VACE LoRAs (for ditto) 2025-10-20 11:49:49 +03:00
kijai 967d15321d fix rcm scheduler check 2025-10-19 19:49:24 +03:00
kijai 9cd79d3d4a Add experimental rCM scheduler
Based on the original code, works but doesn't feel better than dpm++_sde so far
2025-10-19 19:34:17 +03:00
kijai 7ecd55e92a bump version 2025-10-19 18:13:09 +03:00
kijai 99c0efdb45 Expose force_parameter_static_shapes in torch.compile options 2025-10-19 18:01:03 +03:00
kijai a721fe3d0f Better (hopefully) FlashVSR frame selection
First frame was completely discarded before
2025-10-19 13:49:40 +03:00
kijai e8ecda7240 Update wanvideo_1_3B_FlashVSR_upscale_example.json 2025-10-18 16:16:25 +03:00
kijai 7d88316bf8 Fix for the case of applying unmerged lora to already compiled model 2025-10-18 14:00:18 +03:00
kijai 6bcbfc54cf Fix diffusion forcing 2025-10-17 18:20:45 +03:00
kijai 55463a6290 Fix audio_cfg_scale 2025-10-17 18:12:54 +03:00
kijai 6aa0446706 Ovi: fix decoded output audio shape
I don't know how it ever worked with any nodes... but now it should work with all audio nodes.
2025-10-16 18:03:42 +03:00
kijai 848ceb2beb Don't require whole facexlib for lynx 2025-10-16 17:51:21 +03:00
kijai cc06d71ca0 Workaround for bug in pytorch 2.9.0 that makes the VAE use crazy amounts of VRAM
This is probably caused by:

https://github.com/pytorch/pytorch/pull/164027/files

That disables cudnn when using half precision VAE, this workaround simply uses fp32 for the Conv3D operations.
2025-10-16 17:10:05 +03:00
kijai e221fdd7ac Fix normal tiny vae 2025-10-16 14:54:00 +03:00
kijai c876108059 Update fp8_e4m3fn compile warning on older arch to indicate it should now work with latest Triton
https://github.com/woct0rdho/triton-windows/releases/tag/v3.5.0-windows.post21
2025-10-16 13:08:42 +03:00
kijai 9cae986648 Fix for HuMo/Phantom when using cfg 2025-10-16 13:01:21 +03:00
kijai d2ce150490 Update nodes_sampler.py 2025-10-16 10:49:54 +03:00
Anastasiy Safari 904a019491 Another fix 2025-10-15 17:52:26 -07:00
Anastasiy Safari be02ec93cc Fix WanAnimate to use ref_latent without bg_images 2025-10-15 17:28:19 -07:00
kijai 50483a3b54 FlashVSR tweaks 2025-10-16 00:42:33 +03:00
kijai 367c612ddb FlashVSR: Better input frame handling 2025-10-15 23:58:32 +03:00
kijai 9a88b9e40a FlashVSR: Add strength setting 2025-10-15 19:11:41 +03:00
kijai 2fd5bb6ffa Basic context window support for FlashVSR 2025-10-15 18:45:42 +03:00
kijai bb75cddd60 Add minimal FlashVSR upscale support
https://zhuang2002.github.io/FlashVSR/

This only implements the projection model and the VAE, which seems to be enough for upscaling. This does NOT implement any of the streaming and sparse attention code.
2025-10-15 18:24:28 +03:00
kijai 3d42bf62ce Relocate Ovi example 2025-10-14 00:17:47 +03:00
kijai da43a51683 Unmerged lora application tweak
Should resolve some issues when using same model with multiple sampler and different loras
2025-10-13 20:18:43 +03:00
kijai 139bdf827f Squashed commit of the following:
commit 73dd1a06d33953912f5dd684f168028b14e42a36
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Oct 13 19:47:38 2025 +0300

    cleanup

commit 39bc2cecf493e2eb176b55e8841d933f0da1ec39
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Oct 13 19:24:20 2025 +0300

    Allow scheduling ovi cfg

commit 2c153c5f324dbd59670ad9c51a7995459504a3cd
Merge: dba7667 32eb6b4
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Oct 13 17:48:20 2025 +0300

    Merge branch 'main' into ovi

commit dba76674c71af7bf94c82834a0b0e40d94043c99
Merge: 0f11a43 5a0456e
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Oct 12 22:45:43 2025 +0300

    Merge branch 'main' into ovi

commit 0f11a439622799ad8070f8a2b8cc8e6a041b761d
Merge: 0999f50 e2d8c9b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Oct 11 07:48:06 2025 +0300

    Merge branch 'main' into ovi

commit 0999f50cfe025290cd7ce88a8dd1acff0b38d9bd
Merge: d45df1f f1d1c83
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Oct 10 22:16:09 2025 +0300

    Merge branch 'main' into ovi

commit d45df1fb5b7c629b15eabc197357d62bdc232aaf
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Oct 9 20:21:37 2025 +0300

    Remove dependency for librosa

commit d8e7533fdf7eab1d2489c3e025a908c02d997444
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Oct 9 19:57:28 2025 +0300

    Remove omegaconf dependency

commit f4e27ff018e98cb5b09655dceda399baea36b240
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Oct 9 19:31:06 2025 +0300

    Fix VACE

commit 35d3df39294831e5e7568b6f7e16d2ecf2d790a0
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Oct 9 00:26:40 2025 +0300

    small update

commit 96f8ea1d26869ab7e49e12a07f19d5d5a2023253
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Oct 8 22:32:57 2025 +0300

    Create wanvideo_2_2_5B_ovi_testing.json

commit a2511be73b9da7019fd21aeb0b521af941c09150
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Oct 8 22:32:54 2025 +0300

    Update nodes_sampler.py

commit d3688b8db71452ea1f7c9a2bc0216441d524e56c
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Oct 8 21:43:02 2025 +0300

    Allow EasyCache to work with ovi

commit 586d9148a0306ef5d30e9a971a9c3be4cd3ecc97
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Oct 8 19:09:06 2025 +0300

    Update model.py

commit 61eedd2839decdb7d4c2ddd5f1310fdaf49d36ad
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Oct 8 19:09:02 2025 +0300

    I2V fix

commit a97fcb1b9ae9fb7bbfdf668c24816e014a1b58d1
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Oct 8 17:57:28 2025 +0300

    Add nodes to set audio latent size

commit d41e42a697f3d561dabbc22566f633b5f1bbd952
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Oct 8 16:42:04 2025 +0300

    Support loading mmaudio vae from .safetensors

commit 1b0e28ec41e3c97fe1f2f057fef9b9bbcb87bca7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Oct 8 16:19:53 2025 +0300

    Update nodes_sampler.py

commit fbd18f45fe85ede8edcb5aebaea7ceb5b6eab5a2
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Oct 8 10:16:44 2025 +0300

    Fixes for other workflows

commit b06993b637198f7fad92208f3b3dc9a7d7f57c7f
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Wed Oct 8 09:46:27 2025 +0300

    initial commit

    T2V works
2025-10-13 20:16:53 +03:00
kijai 32eb6b480d Allow canceling mid step 2025-10-13 17:47:28 +03:00
kijai 6c38f8cb24 WanAnimate: Force dimensions to be divisible by 16 2025-10-13 17:31:54 +03:00
kijai 64ed8c4183 WanAnimate: Actually use clip embeds when provided
Kinda forgot this model even has those layers...
2025-10-13 16:44:49 +03:00
kijai 5a0456ed9e Merge branch 'pr/1321' 2025-10-12 22:45:33 +03:00
kijai e2d8c9bef5 Allow using tiny vae in other precisions
The model is available in fp16 only though
2025-10-11 07:47:52 +03:00
kijai f1d1c83713 Support 2.2 tiny VAE 2025-10-10 22:13:12 +03:00
kijai 6d2ff33466 Fix double encode 2025-10-10 09:36:37 +03:00
kijai 4c4e7defc2 Fix for 2.2 Fun-camera 2025-10-08 10:14:50 +03:00
kijai 7c09f1e171 Fix some variable names 2025-10-08 00:05:51 +03:00
kijai d6864ec46e Fix for S2V 2025-10-07 23:48:26 +03:00
kijai 260c181e0e Allow lynx reference adapter to run with I2V models 2025-10-07 15:07:31 +03:00
kijai 6075807e03 Fix for some fp8_scaled models 2025-10-07 13:30:41 +03:00
kijai 76ea2aa6a6 Update wan_video_vae.py 2025-10-07 11:41:54 +03:00
kijai 1e56fd1114 Fix weight loading when not using lora or patched linear 2025-10-07 10:56:30 +03:00
kijai 1d94082563 Add lynx requirements.txt 2025-10-06 22:11:00 +03:00
kijai dd888f7a8d GGUF dtype fix 2025-10-06 22:09:58 +03:00
kijai b27c1e5fcd Fix VACE 2025-10-06 21:29:16 +03:00
kijai 269a14f413 Merge branch 'lynx' 2025-10-06 19:39:06 +03:00
kijai 5c74e5c99a bumb version 2025-10-06 19:38:56 +03:00
kijai a088425268 Create wanvideo_T2V_14B_lynx_example_01.json 2025-10-06 19:38:20 +03:00
kijai bb23263f78 Separate full model's ip and ref layer loading
Allows using full ref with lite ip adapter, or full ref alone without loading the ip weights
2025-10-06 18:42:14 +03:00
kijai 174cba5759 Fix some precision issues with some models 2025-10-06 13:20:20 +03:00
kijai c517cfbc17 Merge branch 'main' into lynx 2025-10-06 13:02:23 +03:00
kijai d3010e37ab Fix some torch compile + fp8_fast issues 2025-10-05 12:31:45 +03:00
kijai 3167e088ed Update fp8_optimization.py 2025-10-05 12:26:44 +03:00
kijai d20d271a24 Update nodes_model_loading.py 2025-10-05 12:19:43 +03:00
kijai 05232fb966 Update nodes.py 2025-10-05 01:37:12 +03:00
kijai 00a5cb8e5c Update nodes.py 2025-10-05 01:27:32 +03:00
kijai d380c7d219 Add TSR (Temporal Score Rescaling) to experimental args
https://github.com/temporalscorerescaling/TSR
2025-10-05 01:14:25 +03:00
kijai 71c8a4961b Merge branch 'pr/1240' 2025-10-03 19:24:40 +03:00
kijai dd7c55f5d2 Merge branch 'main' of https://github.com/kijai/ComfyUI-WanVideoWrapper 2025-10-03 19:23:41 +03:00
kijai 300e8a4786 Update readme.md 2025-10-03 19:23:39 +03:00
Jukka Seppänen f66c5df07d Update readme with example links
Added links to WanAnimate and ReCamMaster examples.
2025-10-02 14:39:18 +03:00
kijai 700aae60af Add example for WanAnimate with the new preprocessor 2025-10-02 13:31:31 +03:00
kijai 1ad36cc5f9 Merge branch 'main' into lynx 2025-10-02 00:55:48 +03:00
kijai 697fff7442 Add FlowMatchSAODEStableScheduler
source: https://github.com/eddyhhlure1Eddy/ode-ComfyUI-WanVideoWrapper
2025-10-02 00:53:49 +03:00
kijai 8360666470 Update model.py 2025-10-01 19:34:17 +03:00
kijai 4b079bdf5a Update nodes_sampler.py 2025-10-01 19:30:44 +03:00
kijai 1316f57390 Update model.py 2025-10-01 19:29:05 +03:00
kijai 0730929558 Allow selecting which blocks to use for lynx ref 2025-10-01 19:04:56 +03:00
kijai 1e88104558 compile fixes and cleanup 2025-10-01 18:44:56 +03:00
kijai 28ae01e11b fixes 2025-10-01 10:05:29 +03:00
kijai 5d3c4223a8 Update model.py 2025-09-30 18:59:14 +03:00
kijai ed167a0135 dtype fixes 2025-09-30 18:33:05 +03:00
kijai 1ba1a1662b cleanup and use fp32 text/time embed if available 2025-09-30 18:10:18 +03:00
kijai 4f42dbfacf cfg 2025-09-29 22:46:28 +03:00
kijai 9b1ab5b5c7 Merge branch 'main' into lynx 2025-09-29 22:09:25 +03:00
kijai bea3fcad8f Fix diff diff masking with InfiniteTalk 2025-09-29 19:53:35 +03:00
kabachuha d79686bc20 sageattn3 import rename
Renaming because the official release moved the import names
2025-09-28 14:53:42 +03:00
kijai aefff3796e uncond 2025-09-28 11:21:05 +03:00
kijai 0ed0f54d11 initial lynx full support 2025-09-28 02:01:01 +03:00
kijai 4ba6c88ad8 Update nodes_model_loading.py 2025-09-27 16:11:07 +03:00
kijai 358baa2c2a Lynx + VACE 2025-09-27 02:05:17 +03:00
kijai cf235c0728 update 2025-09-27 00:49:38 +03:00
kijai ab31158673 init: support lite model 2025-09-27 00:49:02 +03:00
kijai 37365817e8 Revert unicode character in warning strings
Seems it may cause issues on some systems
2025-09-26 12:39:39 +03:00
kijai 86193e3cc9 Add a warning for possible custom node version conflict from having multiple WanVideoWrappers installed 2025-09-26 01:01:25 +03:00
kijai 05474da487 Fix denoise_strength usage 2025-09-26 00:32:13 +03:00
kijai cfdae3b49f WanAnimate: Fix case of no mask and no face when looping 2025-09-23 21:08:30 +03:00
kijai f05865de63 Fix VACE 2025-09-23 19:06:15 +03:00
kijai 840553b261 Update model.py 2025-09-23 19:04:37 +03:00
kijai b8c2bb0a93 Add torch native RMSNorm as a choice 2025-09-23 17:41:24 +03:00
kijai 0ba45a1645 Fix flash attention 2025-09-23 00:10:41 +03:00
kijai 17e2452d30 WanAnimate: Fix bug in previous 2025-09-22 19:15:05 +03:00
kijai 1c7e32d8af Support lucy edit 2025-09-22 15:34:48 +03:00
kijai 95cd0ef690 WanAnimate: Use empty frames for face when none are provided and masking is used
Necessary for the mask to work properly in this scenario
2025-09-22 15:18:17 +03:00
kijai 2c4dfaf4c8 WanAnimate: Fix regression in image quality when using fp8
Face blocks not doing well in fp8
2025-09-22 03:12:48 +03:00
kijai 1ab9803bda Update nodes_sampler.py 2025-09-22 02:57:19 +03:00
kijai 02583db554 WanAnimate: Encode pose latents in the loop too for better sync and seams between windows 2025-09-22 02:36:17 +03:00
kijai 7ebdcd4abc Fix context windows with infinite talk. 2025-09-22 01:54:04 +03:00
kijai 0ad92fad01 Update wanvideo_WanAnimate_example_01.json 2025-09-22 01:50:05 +03:00
kijai 2c42ac65e9 Update wanvideo_WanAnimate_example_01.json 2025-09-22 00:58:32 +03:00
kijai 718946ef13 Revert using torch RMSNorm as there are quality concerns 2025-09-21 18:19:27 +03:00
kijai fc3a684b0b WanAnimate: Move face adapter blocks to corresponding main blocks to include them in block swap 2025-09-21 16:34:17 +03:00
kijai 0482667c78 Use torch native (fused) RMSNorm if available 2025-09-21 15:29:33 +03:00
kijai 6cb1241d48 Refactor: Move sampler code to it's own file 2025-09-21 13:43:18 +03:00
Jukka Seppänen 28387fe37f Merge pull request #1216 from MinaWanasTA/fix-display-name-typo
Fix: Correct display name typo for TeaCache node
2025-09-21 13:42:29 +03:00
kijai b3ea982f6f version 2025-09-21 13:28:40 +03:00
kijai a6dc8aeb64 offload before vae encode too in WanAnimate loop 2025-09-20 21:52:37 +03:00
kijai b9fd93c19e WanAnimate: Fix first latent issue when not using bg_images and looping 2025-09-20 20:14:39 +03:00
kijai a90ee9773f Fix VACE first latent auto trimming 2025-09-20 14:51:15 +03:00
kijai 1cc8082452 Fix InfiniteTalk offloading 2025-09-20 14:46:35 +03:00
kijai 0dc8a84564 WanAnimate: zero face pixels for uncond 2025-09-20 01:03:13 +03:00
kijai 795a74d756 Update nodes.py 2025-09-19 22:28:27 +03:00
kijai 14a1f8bd71 Update nodes.py 2025-09-19 21:56:33 +03:00
kijai de4fcd6fb3 Update nodes.py 2025-09-19 18:40:49 +03:00
kijai ca48988742 Update wanvideo_WanAnimate_example_01.json 2025-09-19 18:11:47 +03:00
kijai bca7236c73 Update nodes.py 2025-09-19 18:03:24 +03:00
kijai e4627466f5 WanAnimate block swap fix 2025-09-19 16:59:57 +03:00
kijai 0f1ba64b80 Update wanvideo_WanAnimate_example_01.json 2025-09-19 16:26:36 +03:00
kijai 7aa9086b48 WanAnimate bugfixes 2025-09-19 16:23:16 +03:00
kijai 72f5423d6a Allow uni3c with WanAnimate 2025-09-19 15:28:01 +03:00
kijai 3acfb513ea WanAnimate input frame count tweaks 2025-09-19 14:51:19 +03:00
kijai 543df461ff Allow not using face 2025-09-19 13:57:39 +03:00
kijai b2623b3a04 Update wanvideo_WanAnimate_example_01.json 2025-09-19 13:31:46 +03:00
kijai de304b10d1 face adapter offloading 2025-09-19 12:45:08 +03:00
kijai 623ce5a9fa Update wanvideo_WanAnimate_example_01.json 2025-09-19 12:30:56 +03:00
kijai b2c8cf969f Fix pose and face strength setting 2025-09-19 12:04:07 +03:00
kijai 9d002ddfbc handle undetected faces better 2025-09-19 11:59:07 +03:00
kijai b06562823b Fix GGUF with WanAnimate 2025-09-19 11:35:14 +03:00
kijai 8bc1d6651a Create wanvideo_WanAnimate_example_01.json 2025-09-19 10:19:20 +03:00
kijai 79583d4be3 Support WanAnimate 2025-09-19 10:14:08 +03:00
kabachuha 7eafe77ba0 update magcache ratios 2025-09-16 21:07:39 +03:00
kijai 95f3307d5c Fix InfiniteTalk regression 2025-09-16 18:39:04 +03:00
kijai e9672e8d56 Support HuMo 1.7B 2025-09-16 17:41:07 +03:00
kijai 6d05cc5cf9 Update wanvideo_HuMo_example_01.json 2025-09-16 15:43:34 +03:00
kijai 8f2c1b7824 Allow passing through control_images only on WanVideoVACEStartToEndFrame 2025-09-15 20:45:37 +03:00
kijai 076c4a5c7e Experimental: Allow HuMo to work with InfiniteTalk
Doesn't work that great, but pushing this anyway for possible future usecases
2025-09-15 19:07:34 +03:00
Jukka Seppänen 61c68365d2 Fix dimension ordering in HuMo reference images processing 2025-09-14 17:35:55 +03:00
kijai 70f8ac77dd Correct tooltip 2025-09-13 21:06:53 +03:00
kijai 7fa6be05b8 Fix HuMo context windows when using cfg 2025-09-13 20:41:16 +03:00
kijai f6fc9cee6e Update nodes.py 2025-09-13 20:29:05 +03:00
kijai 699bad3887 InfiniteTalk: Don't stack or return frames when using output path 2025-09-13 20:10:13 +03:00
kijai d47ac9cbe0 Update nodes.py 2025-09-13 19:43:12 +03:00
kijai 1ff6bab5f5 Update nodes.py 2025-09-13 19:40:52 +03:00
kijai 538534c749 Update wanvideo_HuMo_example_01.json 2025-09-13 19:24:03 +03:00
kijai 9063b99bd1 Allow T2V with HuMo 2025-09-13 19:20:19 +03:00
kijai 949647c4d0 Allow setting HuMo frame count to audio length
num_frames -1 = audio length
2025-09-13 17:43:40 +03:00
kijai d50ec8135e HuMo + context windows 2025-09-13 17:21:15 +03:00
kijai 38fd791a77 Squashed commit of the following:
commit fda0fe6e0c21eb10276ae302cd88b6cbcf5b36b5
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 16:55:00 2025 +0300

    Create wanvideo_HuMo_example_01.json

commit cffe3039c3d2fbacd4803329bf31b5fdc45215ba
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 16:30:49 2025 +0300

    Update model.py

commit ddce018a5a6ffeb926860342889b12efb7343ec0
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 16:29:27 2025 +0300

    cleanup

commit 8c021b8b3f66144804e500e74aa6f2be52f9f9fc
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 16:23:27 2025 +0300

    avoid compile graph break

commit ef9c7732042261581b4bba6d980def78633f56dc
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 16:16:13 2025 +0300

    Allow using whisper model without decoder layers

commit 8d0ba29ee84d14be6084ecbcee1cbc9414128fb3
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 15:55:26 2025 +0300

    start/end percent for HuMo audio

commit bfe0d358a8820240f262351e61cfb979cb9a47ff
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 15:37:11 2025 +0300

    cleanup

commit e563ae317f24a7f5751b43cdef4bffbaeaea5114
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Sep 13 14:02:21 2025 +0300

    Make audio work

commit 95855196c51b1124a19079b746c5d12ec70d9026
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Sep 12 18:10:04 2025 +0300

    cfg

commit d5a18b090fe719b7f0b00a0f68e6824685598313
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Sep 12 03:10:15 2025 +0300

    wrong way around

commit 34c8c4842c14002fe4694dfa23c24a65b7ea39d0
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Sep 12 03:01:45 2025 +0300

    Update nodes.py

commit 47d1e2ab5f3e1483782d0739f5b51cdd33707c36
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Sep 12 02:47:03 2025 +0300

    update

    image inputs are working but audio still doesn't do anything

commit 67890d816a64459944091cb01478c1e0ec4c4a82
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Sep 11 21:13:12 2025 +0300

    update

commit dbcef53405bb78feae4c5d2c6b310b76e4ef9949
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Sep 11 17:09:37 2025 +0300

    Update model.py

commit 92c9aac51f4d37988757510a4e57179834cc5de2
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Sep 11 16:15:39 2025 +0300

    init

    untested as no weights released as of yet
2025-09-13 16:55:28 +03:00
kijai 8ce4432fef bump version 2025-09-13 10:34:42 +03:00
kijai 57586bed1d Error when trying to use GGUF extra model with non-GGUF main model 2025-09-13 10:34:34 +03:00
kijai 617f89e0f5 Allow WanVideoExtraModelSelect to see GGUF files too 2025-09-12 11:40:41 +03:00
Mina Wanas 53d26032dd Fix: Correct display name typo for TeaCache node 2025-09-11 11:10:24 +03:00
kijai 9cefe309e3 Update nodes.py 2025-09-10 20:58:45 +03:00
kijai b9e49af4d1 Just resize any dimension mismatch on uni3c input 2025-09-10 17:18:02 +03:00
kijai c254ae4235 Rather interpolate any Uni3C input frame count differences 2025-09-10 16:47:13 +03:00
kijai dd36ab46c6 Pad Uni3c input when using VACE ref 2025-09-10 16:26:20 +03:00
kijai 7cb46ca8c3 Allow Uni3C to work with T2V 2025-09-10 16:09:44 +03:00
kijai e5ef9752a7 Fix fp8_fast when not using loras 2025-09-08 18:54:29 +03:00
kijai fe379fa77f InfiniteTalk: Add option to save results during the process, make it clearer no latents are returned 2025-09-08 18:50:53 +03:00
kijai e5955d8395 InfiniteTalk: Don't create 2 windows if total frames requested is under frame_window_size 2025-09-06 23:10:26 +03:00
kijai 011c0ce38d InfiniteTalk: Pad with silent embeds instead of repeat, add MultiTalkSilentEmbeds -node 2025-09-06 21:58:03 +03:00
kijai ae768f53a4 Make merge_lora switches with multiple loras behave like it used to and how the tooltip indicates 2025-09-06 19:28:16 +03:00
kijai bafd5503ac Fix unianim error when using ref pose 2025-09-05 23:38:08 +03:00
kijai edea93286c Small potential tweaks to Pusa vram use 2025-09-05 16:00:18 +03:00
kijai fbf8f169b6 Update nodes.py 2025-09-05 15:59:01 +03:00
kijai 0ad2b00cc0 Fix force_offload on Multi/InfiniteTalk long I2V node 2025-09-05 15:54:11 +03:00
kijai 9b864b987b Don't error even if redundant offloading is enabled for single text encode 2025-09-04 16:20:44 +03:00
kijai cbbf44f3bc Add Pusa 2.2 example 2025-09-04 02:22:54 +03:00
kijai c38da2242d version 2025-09-03 22:49:02 +03:00
kijai d05796ffb3 Update requirements.txt 2025-09-03 22:48:22 +03:00
kijai 488b6628c7 typo 2025-09-03 22:47:42 +03:00
kijai 8131eaef23 update tooltip 2025-09-03 22:26:07 +03:00
kijai 242637db5c Accept per latent noise multiplier list for Pusa as well 2025-09-03 22:18:37 +03:00
kijai 24ec3f6960 Add WanVideoLoraSelectByName
select lora by name string
2025-09-03 20:22:41 +03:00
kijai 02da8fed47 sigma plot tweaks 2025-09-03 20:01:55 +03:00
kijai 7a29bcb18a Fix sigma graph 2025-09-03 19:49:00 +03:00
kijai 617ce55938 Fix up pusa noise multipliers indexing 2025-09-03 19:24:47 +03:00
kijai e76bf4d899 Add Pusa additional noise handling 2025-09-03 18:32:34 +03:00
kijai 493106b555 Revert this as it didn't end up working
Pusa extra latents seem to need to be encoded individually
2025-09-03 15:48:14 +03:00
kijai 54391300b1 Never use main_device with unmerged loras 2025-09-03 10:00:51 +03:00
kijai c57ffd8150 update example 2025-09-03 00:41:23 +03:00
kijai 838e095dbd Update model.py 2025-09-02 23:46:04 +03:00
kijai 863ea45d23 Update model.py 2025-09-02 23:34:34 +03:00
kijai 44f9584f16 Fix S2V 2025-09-02 23:31:47 +03:00
kijai cc3b79ab18 Update nodes.py 2025-09-02 21:08:02 +03:00
kijai cdfa3a8661 Update nodes.py 2025-09-02 21:05:04 +03:00
kijai 1492310544 Update nodes.py 2025-09-02 20:49:24 +03:00
kijai 0ad24cd7e7 Allow for multiple I2V images for context windows
Similarly as multiple prompts already work, use the new WanVideoEncodeLatentBatch -node to encode them
2025-09-02 19:28:33 +03:00
kijai c538cd4c8a Fix unmerged lora
oops
2025-09-02 18:48:25 +03:00
kijai 3183589eff Merge branch 'main' of https://github.com/kijai/ComfyUI-WanVideoWrapper 2025-09-02 17:09:12 +03:00
kijai 4cf23a5a40 Can't do this inplace with all models 2025-09-02 17:09:05 +03:00
Jukka Seppänen 0405dde5a2 Merge pull request #1165 from zinigor/fix/unbound-local-error
Moved dimension update to where the count is initialized.
2025-09-02 15:06:44 +03:00
Igor Zinovyev 42054724b7 Moved dimension update to where the count is initialized. 2025-09-02 14:32:01 +03:00
kijai 42d7180759 keep scale_weight on gpu to avoid needless casts 2025-09-02 01:04:16 +03:00
kijai 82608009c8 use a copy of the scheduler if provided by the scheduler node to make sure it's always reset 2025-09-01 23:04:05 +03:00
kijai f9755820de handle negative latent indices with extra latents 2025-09-01 14:19:51 +03:00
kijai 56e3366c26 Fix pusa steps 2025-09-01 13:50:41 +03:00
kijai 43d8a2c8a5 flowmatch_pusa split sampling fixes 2025-09-01 13:34:49 +03:00
kijai d4d942c5ee Update nodes.py 2025-09-01 11:17:22 +03:00
kijai 30c2fca799 WanVideoScheduler fixes
support flowmatch pusa when using the node
2025-09-01 11:07:41 +03:00
kijai d9038d575f Allow setting both start and end step by sigma 2025-08-31 19:04:01 +03:00
kijai b8acd3ee48 Allow using fp16 multitalk model with GGUF again 2025-08-31 18:57:59 +03:00
kijai 9f84cd2254 Fix offloading when using stand-in lora 2025-08-31 18:45:13 +03:00
kijai b9b46638ac Fix FantasyPortrait + GGUF main model 2025-08-31 18:41:08 +03:00
kijai b1c8b8a280 remove flex attention as redundant 2025-08-31 18:14:38 +03:00
kijai 79395b872c Update pyproject.toml 2025-08-31 18:10:07 +03:00
kijai cedac626e2 Update __init__.py 2025-08-31 18:09:24 +03:00
kijai 4870b8f081 Update __init__.py 2025-08-31 18:06:20 +03:00
kijai 81dac31735 Merge branch 's2v' 2025-08-31 18:05:30 +03:00
kijai 26a11c3044 Update nodes.py 2025-08-31 18:01:53 +03:00
kijai 10d3a0fbef Update model.py 2025-08-30 23:34:14 +03:00
kijai 40308f1c70 Allow basic FantasyPortrait + S2V 2025-08-30 22:33:32 +03:00
kijai df0b8c419a Fix 2025-08-30 18:01:05 +03:00
kijai 45f129b287 Revert this for now to avoid issues on some systems
Prefetching still works, but using cuda stream caused issue on some systems, investigate reasons later
2025-08-30 14:09:16 +03:00
kijai a7cf6d4ed1 Update nodes.py 2025-08-29 21:33:12 +03:00
kijai 7dd0e1ac61 Update nodes_model_loading.py 2025-08-29 21:21:24 +03:00
kijai f36fa45c29 Loading fixes when using merged loras 2025-08-29 20:57:28 +03:00
kijai 3a290cfd15 Update nodes_model_loading.py 2025-08-29 20:41:40 +03:00
kijai 8f6507fa64 Update nodes.py 2025-08-29 18:10:18 +03:00
kijai 1e4cc84a30 Fix previous 2025-08-29 17:35:27 +03:00
kijai 88cf36936b skip this if not on cuda 2025-08-29 17:29:25 +03:00
kijai 876b5bda5c Update nodes.py 2025-08-29 17:17:37 +03:00
kijai 5f7a5d533b Reduce memory use in the Framepack loop 2025-08-29 16:59:23 +03:00
kijai 6a2053c9d1 Update nodes_model_loading.py 2025-08-29 15:26:47 +03:00
kijai 0c8a883a36 Update nodes.py 2025-08-29 15:25:31 +03:00
kijai 044ff9f25e don't use model management for cuda stream... 2025-08-29 15:21:49 +03:00
kijai 6978686272 Update model.py 2025-08-29 15:16:02 +03:00
kijai 740157bd61 Fix MultiTalk when using the Multi/InfiniteTalk node
InfiniteTalk was also mistakenly using only first frame mask... but honestly after trying the reference method, the result was worse, so I'll default to single frame mask in all cases for now.
2025-08-29 14:48:12 +03:00
kijai 704cca215b Fix control lora when not last lora to be loaded 2025-08-29 14:45:03 +03:00
kijai 999a314652 Fix GGUF offloading 2025-08-29 12:58:36 +03:00
kijai f9be754980 cleanup 2025-08-29 00:17:05 +03:00
kijai a21e4b3210 Update model.py 2025-08-28 20:54:36 +03:00
kijai 469d9168da Update wanvideo2_2_S2V_framepack_pose_testing.json 2025-08-28 19:31:47 +03:00
kijai a053336c1a Update nodes.py 2025-08-28 18:52:10 +03:00
kijai 6faa24b7f2 Implement Framepack long geneneration method
RoPE handling is comfyanon's code
2025-08-28 18:48:39 +03:00
kijai a5621b8739 Add pose input 2025-08-27 22:43:28 +03:00
kijai 5266959a93 Create wanvideo2_2_S2V_context_window_testing.json 2025-08-27 21:41:34 +03:00
kijai a37b12235f Add loudness norm node 2025-08-27 20:51:08 +03:00
kijai 90c3bbb6c2 better context window indices, cleanup 2025-08-27 18:47:31 +03:00
kijai 4bc1ffee24 Update wanvideo_Fun_2_2_control_example_03.json 2025-08-27 16:12:31 +03:00
kijai 56bd2ec5f6 Update WanVideoScheduler -node 2025-08-27 13:19:07 +03:00
kijai a507a5763d update WanVideoScheduler -node 2025-08-27 13:03:24 +03:00
kijai c89b82e479 Merge branch 'main' into s2v 2025-08-27 12:24:08 +03:00
kijai e241998f53 Fix Uni3C 2025-08-27 12:21:25 +03:00
kijai 0538f5d600 Update nodes.py 2025-08-27 02:41:24 +03:00
kijai a84d142cb2 Update nodes.py 2025-08-27 02:40:27 +03:00
kijai 5d17484cc5 continue 2025-08-27 02:25:39 +03:00
kijai 636b252c7e Merge branch 'main' into s2v 2025-08-26 22:41:37 +03:00
kijai 2a3ca69c83 Fix offloading when using merged loras 2025-08-26 22:34:12 +03:00
kijai e85e1ff9e6 Update model.py 2025-08-26 22:05:22 +03:00
kijai 8e65fae3c5 Update model.py 2025-08-26 22:02:36 +03:00
kijai 96f7f6accd fix normal model loading 2025-08-26 22:02:02 +03:00
kijai 3c79851230 ref latent 2025-08-26 21:30:10 +03:00
kijai 63d4b6aada Update nodes.py 2025-08-26 19:07:38 +03:00
kijai 748ec89aa8 init 2025-08-26 19:02:24 +03:00
kijai a1ca0985ec Merge branch 'dev' 2025-08-26 16:39:26 +03:00
kijai 30f185a19b Update pyproject.toml 2025-08-26 16:39:17 +03:00
kijai 4f13aa46d4 Update nodes.py 2025-08-26 01:12:22 +03:00
kijai dfb2a59211 Update nodes.py 2025-08-26 01:11:20 +03:00
kijai 5f5f8ae30a Merge branch 'main' into dev 2025-08-26 00:33:49 +03:00
kijai 82d2ab44e5 update workflows 2025-08-26 00:27:08 +03:00
kijai efcef40b00 fix name 2025-08-26 00:18:05 +03:00
kijai 6fce0e2d3b Add Wav2VecModelLoader to load wav2vec2 from single .safetensors
https://huggingface.co/Kijai/wav2vec2_safetensors/
2025-08-26 00:09:46 +03:00
kijai 665ddd6811 Merge branch 'main' into dev 2025-08-25 19:35:13 +03:00
kijai e836134b90 UniAnimate fixes and optimization 2025-08-25 19:35:05 +03:00
kijai 17f02a91bf Merge branch 'main' into dev 2025-08-25 17:29:04 +03:00
kijai 4d52c10148 Faster Multi/InfiniteTalk model load 2025-08-25 15:06:31 +03:00
kijai 9c63d8d04f Fix temporal_mask dtype 2025-08-25 14:59:43 +03:00
kijai 288a7e96a9 Add experimental scheduler node with sigma plot 2025-08-25 14:53:36 +03:00
kijai d9def84332 Fix UniAnimate and MultiTalk loading 2025-08-25 14:09:08 +03:00
kijai 42df277488 Merge branch 'main' into dev 2025-08-25 12:51:46 +03:00
kijai 66d4c5b90d Slice possible alpha channel away before encoding 2025-08-25 12:50:41 +03:00
kijai bc22008b66 Fix InfiniteTalk v2v 2025-08-24 22:58:10 +03:00
kijai c65f81d089 Handle missing faces better in FantasyPortrait detection 2025-08-24 20:50:11 +03:00
kijai 71c93f87fd Update nodes.py 2025-08-24 20:28:26 +03:00
kijai 0c3030a7d8 Add a node to easier split sigmas, EasyCache fixes 2025-08-24 20:07:34 +03:00
kijai 3d9b52cb10 Fix: can't del the end image at this point... 2025-08-24 14:18:54 +03:00
kijai 61aa0e2ffd Drastically reduce VRAM usage on I2V encode node. 2025-08-24 11:39:56 +03:00
kijai 8a151b5402 Multi/InfiniteTalk sampling loop cleanup and optimizations, support FantasyPortrait within the loop 2025-08-24 11:39:18 +03:00
Jukka Seppänen 54f745bb63 Update nodes.py 2025-08-24 02:03:32 +03:00
kijai d39155213e This can be float for compatibility 2025-08-24 01:06:53 +03:00
kijai 9b98413e1e Update nodes.py 2025-08-23 21:00:48 +03:00
kijai 11f3e41e9a Merge branch 'main' into dev 2025-08-23 21:00:45 +03:00
kijai 9076a3aae8 bump version 2025-08-23 20:57:41 +03:00
kijai b7d7f9afe5 cleanup unianimate stuff 2025-08-23 18:59:21 +03:00
kijai 77d79f6878 This is redundant 2025-08-23 17:55:12 +03:00
kijai 9a2a13498a Cleanup MultiTalk model code 2025-08-23 17:49:04 +03:00
kijai 73cff6ebab Cleanup Multi/InfiniteTalk sampling loop some 2025-08-23 16:27:10 +03:00
kijai 5f521dc169 Add generic CreateScheduleFloatList to assist with LoRA scheduling 2025-08-23 13:04:19 +03:00
kijai 4eeaf1ea19 Update nodes.py 2025-08-23 01:06:57 +03:00
kijai d0b9f2c907 Allow UniAnimate to work with unmerged LoRAs 2025-08-23 01:06:22 +03:00
kijai 948fdd369d cleanup 2025-08-23 01:05:41 +03:00
kijai b1ceaaeb2e Support Wan22 VAE in WanVideoLatentReScale 2025-08-22 15:04:29 +03:00
kijai bb5503c9bb remove prints 2025-08-22 12:57:57 +03:00
kijai ff1fe75919 RoPE ntk scaling 2025-08-22 12:48:12 +03:00
kijai 49951a85b6 Update nodes.py 2025-08-22 01:51:16 +03:00
kijai bfee623d77 Pass original image to 2nd sampler for 2.2 diff diff 2025-08-22 01:48:35 +03:00
kijai c2923068e6 Differential diffusion for Multi/InfiniteTalk long I2V as well 2025-08-22 01:24:16 +03:00
kijai 644f9aeaa3 Fix in multitalk sampling 2025-08-21 22:14:37 +03:00
kijai 65973b0251 gguf loading fixes 2025-08-21 22:13:07 +03:00
kijai 2dfaeb0093 Update nodes_model_loading.py 2025-08-21 19:07:11 +03:00
kijai b4d2290859 Merge branch 'main' into dev 2025-08-21 18:42:47 +03:00
kijai bd0a634ec0 Update nodes.py 2025-08-21 18:26:06 +03:00
kijai bd5f76a3c8 Fix differential diffusion, VACE tweaks 2025-08-21 12:26:29 +03:00
kijai bcdd6dc664 Support loading VACE GGUF modules 2025-08-21 02:36:52 +03:00
kijai fd6404bb2f Update nodes.py 2025-08-21 01:45:31 +03:00
kijai 8bd7a6318c Allow multiple prompts for Multi/InfiniteTalk loop 2025-08-21 01:39:34 +03:00
kijai 6d51934ae8 Update wanvideo_I2V_InfiniteTalk_example_01.json 2025-08-20 22:05:25 +03:00
kijai ad008d369b Create wanvideo_InfiniteTalk_V2V_example_01.json 2025-08-20 22:03:04 +03:00
kijai 1edc6ddb32 Create wanvideo_I2V_InfiniteTalk_example_01.json 2025-08-20 21:31:25 +03:00
kijai 655e6557eb Update nodes_model_loading.py 2025-08-20 21:13:49 +03:00
kijai 70bced06f4 Return exact audio frame count if num_frames above it 2025-08-20 21:08:25 +03:00
kijai 68ef3ac468 Fix 2.2VAE progressbar issue 2025-08-20 18:58:55 +03:00
kijai 6db2d113ba Update nodes.py 2025-08-20 17:12:04 +03:00
kijai b39d6abc1b Fix bug 2025-08-20 17:11:46 +03:00
kijai a0cf985bc3 Update nodes.py 2025-08-20 16:51:23 +03:00
kijai 053f26ce82 Multi/InfiniteTalk v2v 2025-08-20 16:40:13 +03:00
kijai edf2a24519 Fix encode progress bar 2025-08-20 13:58:14 +03:00
kijai 8b4546896e Update nodes.py 2025-08-20 13:39:46 +03:00
kijai a1220f7f36 Handle InfiniteTalk first frame when not using the looping sampling 2025-08-20 13:33:44 +03:00
kijai ff779c9171 Update nodes.py 2025-08-20 12:18:53 +03:00
kijai b9f32539fd This should never have been here... 2025-08-20 12:11:45 +03:00
kijai 35adf24d97 Support Fun 5B control 2025-08-20 12:11:34 +03:00
kijai 5375a01fb6 Add error 2025-08-19 20:30:59 +03:00
kijai 04a99fd4e2 Add tooltips 2025-08-19 20:19:01 +03:00
kijai cae24fc7e0 uni3c adjustment 2025-08-19 20:02:24 +03:00
kijai 4aac86a828 offloading fixes 2025-08-19 19:11:33 +03:00
kijai 39908c9aea Fix fp8_fast with multitalk 2025-08-19 18:21:21 +03:00
kijai 2a0bd1c42d Merge branch 'main' into dev 2025-08-19 18:06:49 +03:00
kijai b6d4c172f0 Allow using GGUF Multi/InfiniteTalk models 2025-08-19 17:48:58 +03:00
kijai d5399e8667 Merge branch 'main' into dev 2025-08-19 16:56:00 +03:00
kijai 345c286c86 Update utils.py 2025-08-19 16:55:51 +03:00
kijai eec70ff3fb Fix multitalk progress bar 2025-08-19 14:50:31 +03:00
kijai c3d75d5e60 Fix multitalk progressbar 2025-08-19 14:49:51 +03:00
kijai 10e1dd68ea Merge branch 'main' into dev 2025-08-19 13:51:02 +03:00
kijai 4f9ef5b7a9 Update nodes.py 2025-08-19 13:34:56 +03:00
kijai 5a8f8730bd Fix for using additional scaled model and merging loras 2025-08-19 13:31:56 +03:00
kijai e2a5fe8b4b Merge branch 'main' into dev 2025-08-19 11:51:08 +03:00
kijai 3cd6a930c3 Update nodes.py 2025-08-19 11:10:16 +03:00
kijai e0eee40bde Add InfiniteTalk loop method 2025-08-19 10:45:42 +03:00
kijai daa638befb fix UniAnimate 2025-08-19 01:04:26 +03:00
kijai 61ec1ff641 Create wanvideo_MTV_Crafter_example_WIP.json 2025-08-19 00:51:15 +03:00
kijai f182b976f8 update 2025-08-19 00:36:13 +03:00
kijai 0b1fa14ade Update nodes.py 2025-08-18 23:11:17 +03:00
kijai 09710f9ca0 Basic MTV Crafter support
https://github.com/DINGYANB/MTVCrafter
2025-08-18 22:42:11 +03:00
kijai c16a7b5a7d Update nodes_model_loading.py 2025-08-18 17:06:25 +03:00
kijai 282947d93a GGUF fix 2025-08-18 16:57:07 +03:00
kijai e4f422cf99 Update multitalk.py 2025-08-17 19:33:20 +03:00
kijai aa62d59d9e Update nodes.py 2025-08-17 19:23:51 +03:00
kijai 497fafe2ea Merge branch 'main' into dev 2025-08-17 19:23:34 +03:00
kijai 427e7be6c2 fix multitalk 2025-08-17 18:24:48 +03:00
kijai 1d7e0848e0 Allow Phantom to work with MultiTalk 2025-08-17 16:46:50 +03:00
kijai 9a7dba06c1 Update multitalk.py 2025-08-17 16:43:46 +03:00
kijai be285f5258 multitalk fix 2025-08-17 15:14:27 +03:00
kijai 689797bdb7 fix 2025-08-17 15:02:45 +03:00
kijai 6181f30b37 Update nodes_model_loading.py 2025-08-17 14:11:15 +03:00
kijai 6f83ca96bc Update nodes.py 2025-08-17 13:57:01 +03:00
kijai 1eb4db5ff9 Merge branch 'main' into dev 2025-08-17 13:52:30 +03:00
kijai b748ed8bdd Squashed commit of the following:
commit f06ce25f6c50174b0e7b12c3517827a6650f33fa
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Aug 17 13:10:02 2025 +0300

    redundant

commit 9957f765e48f68b21713c33db334231207107b01
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Aug 17 12:57:12 2025 +0300

    Update nodes.py

commit 9797f15ed1e5a9b9b93f777e92cd82ed584c2b57
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Aug 17 12:35:46 2025 +0300

    Remove the "FM" scheduler

    This is just dpm++

commit 244fc2b125253a4c7b6e91ca8dd0b4dc8564c2f0
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Aug 17 02:05:26 2025 +0300

    just use default rope I guess...

commit 5c3a60d3b189864f8f2f0977a552499b202326ba
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Aug 17 01:35:49 2025 +0300

    Update nodes.py

commit ee47ed17c086f8468fd5a256318274c0077f6dfc
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Aug 17 01:35:36 2025 +0300

    Update model.py

commit 8a17ee7d41ae5789b8d7ed01eb8b48c0c3b52ce7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Aug 16 23:20:10 2025 +0300

    init
2025-08-17 13:51:23 +03:00
kijai 4b4a3b4f8f Update nodes.py 2025-08-17 10:55:06 +03:00
kijai d6ac1842fa remove print 2025-08-17 01:41:49 +03:00
kijai 7fc5ec305d Refactor model loading 2025-08-16 16:03:15 +03:00
kijai e5b0b53265 Update wanvideo_2_1_I2V_FantasyPortrait_example_01.json 2025-08-16 15:46:10 +03:00
kijai e3ef4dd27b bump version 2025-08-16 12:57:32 +03:00
kijai 8bf010a835 onnx device selector 2025-08-16 00:11:11 +03:00
kijai 0ced709f1b Fix multitalk scheduler
Don't know if it's of any relevance, but working again
2025-08-15 18:06:45 +03:00
kijai 9659589d67 Update nodes.py 2025-08-15 17:59:09 +03:00
kijai b60ab34db6 Update nodes_model_loading.py 2025-08-15 15:08:43 +03:00
kijai 9f3a5962e6 Update nodes_model_loading.py 2025-08-15 14:56:09 +03:00
kijai 2df76f3363 Update nodes_model_loading.py 2025-08-15 14:45:53 +03:00
kijai 0beb5a8911 Update nodes.py 2025-08-15 13:47:46 +03:00
kijai c05109ff82 expose adapter projection scaling for FantasyPortrait
Not sure how useful but seemingly can at least separate the effect on mouth and rest of the head
2025-08-15 13:45:41 +03:00
kijai 3de816004d Update readme.md 2025-08-15 13:34:53 +03:00
kijai 8d487f99b4 Add face det code license 2025-08-15 13:34:01 +03:00
kijai 7d2c0b58e2 Update wanvideo_2_1_I2V_FantasyPortrait_example_01.json 2025-08-15 13:17:44 +03:00
kijai d6da0f17bb Add landmark visualization 2025-08-15 13:08:45 +03:00
kijai 656f5b1132 Update nodes.py 2025-08-15 12:45:37 +03:00
kijai 8c4f2bb95e Update pdf.py 2025-08-15 12:31:50 +03:00
kijai d6bee68845 Update multitalk.py 2025-08-15 12:24:37 +03:00
kijai 4641665393 Update nodes.py 2025-08-15 12:24:33 +03:00
kijai 2e9e8c03ea Add trajectory example for Fun 2.2 2025-08-15 10:50:26 +03:00
kijai 03458ee955 Update nodes.py 2025-08-15 10:27:49 +03:00
kijai f51c284354 remove print 2025-08-15 10:18:40 +03:00
kijai f2fd411a7b Add option for manually setting rope offset for Stand-In reference
Experimental
2025-08-15 10:18:00 +03:00
kijai ef8c4aba1d Update __init__.py 2025-08-14 20:27:35 +03:00
kijai 5cfc629de5 Update face_utils.py 2025-08-14 20:18:33 +03:00
kijai 5e7c89bd00 Update wan_video_vae.py 2025-08-14 19:01:02 +03:00
kijai 994e6e3afb Create requirements.txt 2025-08-14 19:00:31 +03:00
kijai 0c4426fc2b FantasyPortrait reqs and cleanup
unsure if I want to add onnx in the requirements at this stage...
2025-08-14 18:53:08 +03:00
kijai 43447dbaf9 Update utils.py 2025-08-14 16:57:01 +03:00
kijai 2f49c7fcdc Update nodes.py 2025-08-14 16:23:23 +03:00
Jukka Seppänen 7bf10bc800 Merge pull request #1039 from Ordinator0/extract-lora-stem
correctly extract lora name
2025-08-14 16:06:33 +03:00
kijai 8d8f388a3c Workaround for Uni3C on latest diffusers 2025-08-14 15:56:58 +03:00
Ordinator0 e8145f2074 correctly extract lora name 2025-08-14 05:55:03 -07:00
kijai 24683f1dad Squashed commit of the following:
commit 6a01a8a1d80d36b5b8ac979a36069c8eb0c2f9a7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Aug 14 15:40:50 2025 +0300

    Update wanvideo_2_1_I2V_FantasyPortrait_example_01.json

commit e3cf4bf5bc13321e6d8f91fcb7ee92210a7adf01
Merge: bbf14ec f3d5f6b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Aug 14 15:08:53 2025 +0300

    Merge branch 'main' into fantasy_portrait

commit bbf14ec9e9965c1a8582eea02b50913e79a036d0
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Aug 14 02:16:11 2025 +0300

    update

commit 8192f9f4302b3933e48641a8ef313929e3e263aa
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Aug 14 01:46:15 2025 +0300

    progress bar, fix context windows

commit 39fab8ad4d950a974ba49bd177814f30f518b478
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Aug 14 01:29:15 2025 +0300

    Update nodes.py

commit 36f472c0134e6342ab8c2062a7e3f0b0c829003b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Aug 14 01:14:24 2025 +0300

    Add start/end percent

commit 16f5922c6bc575754412c9b473907377562b956c
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Thu Aug 14 00:58:57 2025 +0300

    init
2025-08-14 15:41:25 +03:00
kijai f3d5f6b3ab allow stand-in to work with chunked rope 2025-08-14 14:12:06 +03:00
kijai b29756a962 cleanup 2025-08-14 13:35:36 +03:00
kijai 885fc25cb6 Delete wanvideo_480p_I2V_endframe_example_01.json 2025-08-14 13:32:00 +03:00
kijai c9daa36ef7 Update wanvideo_Stand-In_reference_example_01.json 2025-08-14 13:31:53 +03:00
kijai 5691ae5314 Refactor Fun control input
Allows using reference only
2025-08-14 13:03:51 +03:00
kijai 64a0c03673 Fix Fun 2.2 reference image input
Turns out the reference image doesn't work in the initial fp8_scaled models I shared due to the ref_conv layer ending up scaled as well, I've fixed the models and added error indicating the issue when trying to use reference image on such bugged model.

No change in behaviour when using start image instead.
2025-08-14 11:13:03 +03:00
kijai e1c052ff43 Update attention.py 2025-08-14 00:11:02 +03:00
kijai 8db5521463 Possible bandaid for q k v dtype mismatch
Still no clue where or why it happens, but clearly is happening for some when using sage with Stand-In
2025-08-14 00:10:38 +03:00
kijai 54dbcacd57 Fix fun control 2025-08-13 21:20:05 +03:00
kijai 4805b70d3b Fix end_image only for 2.2 I2V 2025-08-13 21:04:08 +03:00
kijai 68e18eabb2 Create wanvideo_Fun2_2_control_camera_example_01.json 2025-08-13 19:33:25 +03:00
kijai db86a1b838 Update nodes.py 2025-08-13 19:26:48 +03:00
kijai 00bcd1ce8b Support Fun 2.2 Control-Camera
https://huggingface.co/alibaba-pai/Wan2.2-Fun-A14B-Control-Camera

https://huggingface.co/Kijai/WanVideo_comfy_fp8_scaled/upload/main/Fun
2025-08-13 18:20:00 +03:00
kijai b08f10ffd8 Re-implement Stand-in kv_cache 2025-08-13 11:58:22 +03:00
kijai d596901055 Allow stand-in with VACE
Don't know how well it can actually work, has some effect
2025-08-13 10:00:37 +03:00
kijai f933ad0840 Fix stand-in with GGUF 2025-08-13 01:05:22 +03:00
kijai 083237051f I suppose kv_cache isn't actually used for this... 2025-08-13 00:07:14 +03:00
kijai 8319c961c0 Update model.py 2025-08-12 22:10:57 +03:00
kijai 5420c5aad9 Update nodes.py 2025-08-12 22:10:50 +03:00
kijai 987e92cc43 Create wanvideo_Stand-In_reference_example_01.json 2025-08-12 22:10:46 +03:00
kijai a5ae70624d Support Stand-In LoRA
https://github.com/WeChatCV/Stand-In
2025-08-12 21:51:38 +03:00
kijai 2c854c53ee version 2025-08-12 21:48:15 +03:00
kijai 06863fe2d2 Allow doing T2V with the I2V encode node when no images are connected and allow using native set latent noise mask -node
Latent masking still doesn't work properly
2025-08-12 11:47:33 +03:00
kijai 97da701379 Fix DF sampler 2025-08-11 18:46:46 +03:00
Jukka Seppänen bf84e99dec Merge pull request #1019 from UaRuairc/patch-1
Fix res_multistep.step to set sigma_next correctly when ending early
2025-08-11 16:24:38 +03:00
kijai 057ee642f6 Add async prefetch as option for block swapping
Can speed up block swapping on some systems, added debug option as well so you can see if enabling the option gives you any benefit.
2025-08-11 16:17:16 +03:00
FB 1e091d181b Fix res_multistep.step to set sigma_next correctly when ending early
Use the next sigma from the schedule when sampler ends early, only assume `sigma_next = 0` when on the true final step.
2025-08-11 14:17:08 +01:00
kijai 7c93aea182 rename fp8 optimization in code for clarity as it does more than that with loras 2025-08-11 14:32:57 +03:00
kijai 341932b58f Fix unmerged lora application on unquantized models 2025-08-11 01:34:25 +03:00
kijai d6270ace0a Merge branch 'pr/1013' 2025-08-10 21:35:41 +03:00
kijai abd350bbf4 Actually catch interrupt to free memory 2025-08-10 21:07:13 +03:00
Mel Massadian 19e88b0293 Merge branch 'kijai:main' into feat/encode_progress 2025-08-10 15:37:02 +02:00
kijai 7e474b6b16 Fix pusa extension 2025-08-10 01:34:16 +03:00
kijai 53e740a1e6 Update nodes.py 2025-08-09 16:24:30 +03:00
kijai 8185e9b3cb Better fp8 linear layer patching with torch.compile 2025-08-09 16:18:33 +03:00
denk c4c06d2c94 Revert "renamed controlnet workflows"
This reverts commit 31ae7a6227.
2025-08-09 16:09:23 +03:00
denk 31ae7a6227 renamed controlnet workflows 2025-08-09 14:56:52 +03:00
denk af0b4d9a29 Fix controlnet bug for wan2.2 and add workflow examples 2025-08-09 13:38:19 +03:00
kijai cf8e403f88 Fix gguf 2025-08-09 13:28:20 +03:00
kijai 0da602f8e9 silence the img key error spam when using 2.1 I2V loras on 2.2 or T2V 2025-08-09 12:17:55 +03:00
kijai 99c946ae43 Fix lora not always working with compile 2025-08-09 11:20:30 +03:00
kijai f0e81d01c9 prints 2025-08-09 10:52:27 +03:00
kijai b57e245beb Update wanvideo_Fun_2_2_control_example_01.json 2025-08-09 10:49:52 +03:00
kijai ceedc9e493 Update nodes_model_loading.py 2025-08-09 10:49:42 +03:00
kijai 48fa904ad8 fp8 matmul for scaled models
Fp8 matmul (fp8_fast) doesn't seem feasible with unmerged LoRAs as you'd need to first upcast, then apply LoRA, then downcast back to fp8 and that is too slow. Direct adding in fp8 is also not possible since that's just not something fp8 dtypes support.
2025-08-09 10:17:11 +03:00
kijai 1757847e5f bump version 2025-08-09 10:13:05 +03:00
Mel Massadian a4871802ad feat: ✨ add progress for VACE_Encode
allows to cancel earlier
2025-08-08 16:25:31 +02:00
kijai a64d8bb0df Fix fun control 2.2 timestep scheduling 2025-08-08 17:03:01 +03:00
kijai f82d47273c Update wanvideo_Fun_2_2_control_example_01.json 2025-08-08 16:26:26 +03:00
kijai 9be0eb133c Fix lora scheduling in multilora loader 2025-08-08 13:47:08 +03:00
kijai db93edb665 Create wanvideo_Fun_2_2_control_example_01.json 2025-08-08 13:42:07 +03:00
kijai e1acc2c4e7 Support Fun 2.2 2025-08-08 11:49:04 +03:00
kijai d446e97309 Experimental RAAG (Ratio Aware Adaptive Guidance) implementation
https://arxiv.org/abs/2508.03442

Available in experimental_args, alpha > 0 = enabled, default value is 1.0
2025-08-08 00:17:32 +03:00
kijai ab06ac2f64 Closer to original EchoShot implementation, bugfixes
Turns out I never had the right way of using EchoShot and was always just accidentally using my existing rope splitting... which just worked with the EchoShot weights, this commit adds proper EchoShot and it's used when their prompt format is used, the previous behaviour can be restored by using my splitting method with " | " between the prompts. EchoShot example has been updated to reflect that.
2025-08-07 12:08:42 +03:00
kijai 6191451e2c support qwen 2.5 7b 2025-08-06 21:02:39 +03:00
kijai 2a21c4f6ae VAE cleanup and adjustments, add node to allow using native VAE 2025-08-06 17:28:50 +03:00
kijai 69dd4689bf Make prompt extender use sampling and add seed option 2025-08-05 21:17:11 +03:00
kijai 92eac2a9e1 typo 2025-08-05 18:25:42 +03:00
kijai 9129c8cc1d Revert dimension check until I learn how to math 2025-08-05 14:26:45 +03:00
kijai bacb2a2053 Update nodes_model_loading.py 2025-08-05 01:44:42 +03:00
kijai 8e71286b6d Allow torch.compile VAE decoder
Slight speedup especially for 5B...
2025-08-05 01:31:31 +03:00
kijai 0e9b1973b0 Add error indicating incompatible input size
More relevant now with 5B VAE requiring things to be divisible by 32
2025-08-05 00:35:52 +03:00
kijai 6ec463288f Fix LoRA scheduling with multiple LoRAs 2025-08-04 23:47:02 +03:00
kijai 1b73c10ed5 Update nodes.py 2025-08-04 23:18:50 +03:00
kijai fc7f504e85 Fix empty latents bug from previous commit 2025-08-04 23:16:25 +03:00
kijai fe7e124329 More 5B features 2025-08-04 21:27:36 +03:00
kijai 78a7e20838 Fix lora schedule with multilora node 2025-08-04 16:51:53 +03:00
kijai 7862749cfe Fix single text encode 2025-08-04 01:25:06 +03:00
kijai 267e86704b Update gguf.py 2025-08-04 00:49:53 +03:00
kijai 7a00a875e7 LoRA scheduling on GGUFs too 2025-08-04 00:49:29 +03:00
kijai 3552478d02 Update nodes_model_loading.py 2025-08-04 00:04:21 +03:00
kijai 801f543a7a Allow LoRA timestep scheduling
LoRA strength can now be a list of floats that represents the strength at any given step
2025-08-04 00:02:19 +03:00
kijai 88b5feb1ed Fix VRAM management (again) 2025-08-03 22:35:24 +03:00
Jukka Seppänen 1654741cdf Merge pull request #955 from SD-inst/error_propagation
Propagate errors to frontend
2025-08-03 22:03:30 +03:00
kijai 5cde6f2216 Fix T5 double RAM use 2025-08-03 12:28:25 +03:00
kijai 26c1911e41 Allow loading 5B controlnet
inference doesn't work yet
2025-08-03 12:17:10 +03:00
kijai 309a93c221 Update nodes_model_loading.py 2025-08-02 21:13:29 +03:00
rkfg 00875a65a6 Propagate errors to frontend 2025-08-02 20:29:55 +03:00
kijai c93e69fb9a Merge branch 'main' of https://github.com/kijai/ComfyUI-WanVideoWrapper 2025-08-02 19:53:26 +03:00
kijai b5199ec7e0 Fix control lora 2025-08-02 19:53:20 +03:00
Jukka Seppänen 158b62dafa Merge pull request #952 from SD-inst/tdqm-fix
Reduce tdqm spam
2025-08-02 19:46:40 +03:00
kijai e9b537f23c Update nodes.py 2025-08-02 19:35:02 +03:00
kijai 7074688ced Add prompt extension with local Qwen 2025-08-02 19:27:26 +03:00
rkfg ed247d78b6 Reduce tdqm spam 2025-08-02 16:38:27 +03:00
kijai ce4ac89659 Add shift to DummyComfyWanModelObject 2025-08-02 10:40:42 +03:00
kijai 8035353cc0 Update pyproject.toml 2025-08-02 02:09:54 +03:00
kijai 99d012230f Catch missing taew2_1.safetensors better 2025-08-02 02:09:42 +03:00
kijai a9925496d5 update some workflows 2025-08-02 02:03:01 +03:00
kijai 907d1b0b0f 5B + context windows 2025-08-02 00:34:10 +03:00
kijai 25fe9ea642 handle cancellation better 2025-08-01 21:57:08 +03:00
kijai 168f3ccbee Update nodes.py 2025-08-01 21:24:35 +03:00
kijai d9be7b9c26 Allow other schedulers for 5B 2025-08-01 21:21:55 +03:00
kijai e05864e751 Update nodes.py 2025-08-01 18:31:40 +03:00
kijai 83bb0e7f35 fp8 text encoder fixes 2025-08-01 17:42:32 +03:00
kijai 8321e19009 Fix single text encoder after previous change 2025-08-01 14:45:29 +03:00
kijai 5406a72f62 Improve text encoder cache and add alternative node which leaves nothing to memory after encoding
WanVideoTextEncodeCached -node will load T5 when the prompt is not found in the cache, then unload it completely.
2025-08-01 12:36:19 +03:00
kijai 9edab74562 Default non_blocking to False
too many RAM issues
2025-08-01 02:29:07 +03:00
kijai 7eebde487f typo 2025-08-01 00:54:54 +03:00
kijai 945700d56e 5B fixes 2025-07-31 22:41:56 +03:00
kijai 5bacc50088 Add safety check for context windows to avoid crash on tensor slicing 2025-07-31 17:09:05 +03:00
kijai b60814fb6b Fix split prompting, get rid of diffusers on model class 2025-07-31 15:11:47 +03:00
kijai 10b67de349 Fix loading TAEW from extra paths 2025-07-31 13:47:59 +03:00
kijai 4cf2ffc979 Merge branch 'main' of https://github.com/kijai/ComfyUI-WanVideoWrapper 2025-07-31 12:50:17 +03:00
kijai ef23a6213d return denoised samples as well for previewing unfinished latents 2025-07-31 12:50:10 +03:00
Jukka Seppänen ee1e762a7f Update readme.md 2025-07-31 12:37:16 +03:00
kijai 0cc4583653 flowedit fix 2025-07-31 12:21:40 +03:00
kijai e545b18627 vid2vid fix
Need to use add_noise_to_samples = True now for vid2vid, defaulting to False to not break 2.2 dual sampler workflows, but hardcode to always be True if denoise is used to not break old workflows
2025-07-30 19:41:51 +03:00
kijai 4eabdcce78 wrong way around 2025-07-30 18:45:19 +03:00
kijai 138b377d1d Fix loading fp8 models where everything is fp8... 2025-07-30 17:34:56 +03:00
kijai 9a5d752ba1 Support custom sigmas on more schedulers 2025-07-30 16:20:48 +03:00
kijai ec066008a8 Make start/end step work with vid2vid 2025-07-30 15:59:22 +03:00
kijai 406643f8ee Create wanvideo_2_2_5B_I2V_example_WIP.json 2025-07-30 00:46:46 +03:00
kijai 6ef3ee83ef Update nodes.py 2025-07-30 00:40:42 +03:00
kijai 3f4e8433bd Add node to allow getting BasicScheduler sigmas 2025-07-29 23:06:01 +03:00
kijai 1570287db1 Update nodes_model_loading.py 2025-07-29 18:10:02 +03:00
kijai 15a57d22f7 fix slightly increased memory use 2025-07-29 18:08:32 +03:00
kijai f835baa3d3 fix riflex 2025-07-29 10:52:39 +03:00
kijai 01e452c0d1 Update wanvideo2_2_I2V_A14B_example_WIP.json 2025-07-29 09:14:45 +03:00
kijai 234cffbf24 Fix comfy org scaled models giving black output 2025-07-29 08:05:48 +03:00
kijai de11d52348 Update wanvideo2_2_I2V_A14B_example_WIP.json 2025-07-29 02:46:01 +03:00
kijai 6875b99611 remove this for now 2025-07-29 02:33:20 +03:00
kijai e4fb931878 Update wanvideo2_2_I2V_A14B_example_WIP.json 2025-07-29 02:33:12 +03:00
kijai 5238f40fe3 Update nodes.py 2025-07-29 02:20:20 +03:00
kijai a49936b645 Create wanvideo2_2_I2V_A14B_example_WIP.json 2025-07-29 02:20:17 +03:00
kijai a3032e4363 clear lora patches first 2025-07-29 01:45:25 +03:00
kijai 7e290c67bf Fix comfy progress bar when not doing full steps 2025-07-28 23:16:06 +03:00
kijai fb7073259f detect 2.2 I2V
untested
2025-07-28 22:43:37 +03:00
kijai 498c6859fe Merge branch 'pr/883' 2025-07-28 22:16:48 +03:00
kijai 08d4569794 Update nodes.py 2025-07-28 21:42:35 +03:00
kijai d6e9af7090 2.2 stuff 2025-07-28 21:38:40 +03:00
kabachuha fd285a7e7d implement tangential guidance 2025-07-28 21:08:10 +03:00
kijai 998a69cc0a Update utils.py 2025-07-28 10:46:46 +03:00
kijai 3659c3f6e6 Progress bar for text encoding 2025-07-28 10:00:17 +03:00
kijai ee6e790d80 Update nodes.py 2025-07-28 09:50:50 +03:00
kijai 5e725ff8fb Allow running T5 no cpu
This is actually surprisingly fast, on AMD Ryzen 9 9900X it's faster to not move the T5 to GPU than to keep it on CPU all the way.
2025-07-28 09:39:36 +03:00
Jukka Seppänen 925ec31578 Merge pull request #876 from AustinMroz/main
Fix animated latent preview compatibility
2025-07-27 19:33:11 +03:00
Austin Mroz 28da83c9b6 Fix animated latent preview compatibility 2025-07-27 10:39:01 -05:00
kijai f61aae6158 minimax + context windows 2025-07-27 17:10:41 +03:00
kijai f86d2b84a0 Update nodes.py 2025-07-27 15:34:57 +03:00
kijai 3567202e75 Fix compile and support radial attn in the DF sampler 2025-07-27 14:19:31 +03:00
Jukka Seppänen c0a4e8aa87 Update readme.md 2025-07-27 01:43:32 +03:00
kijai bc68d35e5a Create wanvideo_1_3B_EchoShot_example.json 2025-07-26 19:25:04 +03:00
kijai cb36e1013e Fix 2025-07-26 19:01:58 +03:00
kijai df1476bf42 Support EchoShot
https://github.com/D2I-ai/EchoShot
2025-07-26 18:51:29 +03:00
kijai 9a1ab1c656 only show lora metadata if it actually exists 2025-07-26 14:32:53 +03:00
kijai 5c36a4b111 minor memory tweaks 2025-07-26 14:19:07 +03:00
kijai 3157e0389e Update nodes.py 2025-07-26 11:50:24 +03:00
kijai 34c2cc23a9 Update nodes_model_loading.py 2025-07-25 20:49:41 +03:00
kijai a35eb7dfa4 Basic sageattn3 support
Probably going to need further timestep scheduling as the quality of sage3 on it's own is pretty bad.
2025-07-25 20:48:34 +03:00
kijai 76750eec05 Allow loading aitoolkit/lycoris format 2025-07-25 17:24:43 +03:00
kijai b89cf970e1 Fix GGUF + LoRA + VACE
Always something part two...
2025-07-25 16:44:43 +03:00
kijai 794a38cc7e Further GGUF+LoRA fixes 2025-07-25 00:40:48 +03:00
kijai 73b7444baa remove print 2025-07-24 23:54:30 +03:00
kijai ea36decf07 Fix GGUF + LoRA + torch.compile
Always something...
2025-07-24 23:52:18 +03:00
kijai b3820055d0 Fix GGUF with LoRAs and support GGUF witgh SetLoras -node 2025-07-24 14:59:59 +03:00
kijai 35e637cbfd Allow merging LoRA to fp8 scaled models 2025-07-24 11:51:36 +03:00
kijai d6425cab02 Update nodes_model_loading.py 2025-07-24 00:08:39 +03:00
kijai 59cd0dc272 Update nodes.py 2025-07-23 18:11:15 +03:00
kijai 838803d053 Update nodes_model_loading.py 2025-07-23 17:58:28 +03:00
kijai bdce322e65 remove loras if SetLora bypassed/disconnected 2025-07-23 17:51:59 +03:00
kijai 66226a8a1b why was this ever this way... 2025-07-23 17:17:27 +03:00
kijai 3ed58330d9 better 2025-07-23 16:56:33 +03:00
kijai 65ac9fa6d9 cache positive and negative separately, add cache to single encode node as well 2025-07-23 16:33:11 +03:00
kijai 5e56e3649f fix control lora 2025-07-23 16:20:30 +03:00
kijai 0212ad7a2d cache only the prompts
so we can disconnect T5 completely if text embed exists in cache
2025-07-23 15:33:32 +03:00
kijai 1a1af02912 Allow using different init image for context windows beyond first one
This can be useful with models like MAGREF that behave differently when the init image is padded with white, idea is to use full init for first window and padded image for the rest, so that it works more like reference and doesn't force the window to snap back to the init.
2025-07-23 12:51:57 +03:00
kijai 9d9b188b0b Update fp8_optimization.py 2025-07-23 01:22:42 +03:00
kijai fa93f5b3c1 Add optional prompt disk caching
Caches text embeds to disk with unique hash based on T5 name, dtype and prompts. If found on disk skips encoding, persists after restarts.
2025-07-23 00:52:56 +03:00
kijai 729b6fdf7c support sparse sageattn2 for radial attn
Needs this installed:
https://github.com/Radioheading/Block-Sparse-SageAttention-2.0
2025-07-23 00:27:14 +03:00
kijai 9c45f3a97d Update nodes_model_loading.py 2025-07-23 00:10:14 +03:00
kijai eea896dace Update nodes_model_loading.py 2025-07-22 18:25:03 +03:00
kijai 70fcdff3c5 Update nodes_model_loading.py 2025-07-22 17:36:41 +03:00
kijai bfd6af8141 Update nodes_model_loading.py 2025-07-22 15:41:32 +03:00
kijai 98fa84bb17 Update nodes_model_loading.py 2025-07-22 13:58:43 +03:00
kijai 3988accdf3 Update nodes_model_loading.py 2025-07-21 23:44:24 +03:00
kijai 9911b5e6f0 oops 2025-07-21 22:49:05 +03:00
kijai 91515717d6 bump version 2025-07-21 22:46:50 +03:00
kijai 296baa30ce Add WanVideoSetLoRAs
Node to set the LoRA weights to use with the unmerged LoRA mode, not able to merge LoRAs but allows instant LoRA switching without any loading times. The effect of unmerged LoRAs is stronger and differs from merged LoRAs.
2025-07-21 22:39:17 +03:00
kijai 41d8bd9ec9 fix for context windows 2025-07-21 20:09:53 +03:00
kijai 4bf78e5316 Possibly reduce LoRA loading memory usage 2025-07-21 15:25:07 +03:00
kijai 6ef9224dfe Update model.py 2025-07-21 13:51:33 +03:00
kijai aeac12ed7a Only convert layers that actually have scaled weights... 2025-07-21 01:57:27 +03:00
kijai 29ce253bb6 Support fp8_scaled models and allow running LoRAs unmerged on other models as well 2025-07-21 01:37:30 +03:00
kijai 59add151eb Update attention.py 2025-07-20 16:22:27 +03:00
kijai e5b125cd89 Allow setting any block as dense for Radial attention
Also added WanVideoBlockList helper to to create list of ints that can be inserted to the dense_blocks -input.
2025-07-20 16:21:37 +03:00
kijai 0dba876979 Correct radial attn mask size when using context windows 2025-07-20 13:22:56 +03:00
kijai bb82d63453 Update nodes.py 2025-07-20 13:08:40 +03:00
kijai 768a4b2d90 Fix radial attention cache not resetting on resolution change, add more robust supported dimension check 2025-07-20 12:15:01 +03:00
kijai 9f75b7e051 Fix tiled_vae + endframe 2025-07-20 02:33:24 +03:00
kijai 1f25bba463 Merge branch 'pr/828' 2025-07-20 02:13:26 +03:00
kijai 47978275bb Update nodes.py 2025-07-20 02:13:07 +03:00
kijai edd9b20691 Add res_multistep 2025-07-20 01:41:17 +03:00
komikndr d3aba90119 adding 64 block_size 2025-07-19 20:05:19 +07:00
kijai 0f6fe96626 Update __init__.py 2025-07-18 23:01:03 +03:00
kijai feddafcea2 Update model.py 2025-07-18 18:52:32 +03:00
kijai 5a305a7f9f include sparse_sage_API code
https://github.com/jt-zhang/Sparse_SageAttention_API
2025-07-18 17:16:43 +03:00
kijai cf18ecc2da update i2v basic example 2025-07-18 16:42:02 +03:00
kijai c68f5f7439 move node 2025-07-18 16:41:49 +03:00
kijai c08296747f more refactoring 2025-07-18 16:31:34 +03:00
kijai 6b9565ed63 Code refactoring 2025-07-18 16:15:50 +03:00
kijai 8b8da7d8fe Update nodes.py 2025-07-18 14:41:38 +03:00
kijai 8538c517e6 Merge branch 'radial_attn' 2025-07-18 14:40:10 +03:00
kijai ed11e2bb3e Update nodes.py 2025-07-18 14:37:11 +03:00
kijai a1cc320c36 force using the set node 2025-07-18 14:26:45 +03:00
kijai 0e920e5b18 support VACE 2025-07-18 14:04:04 +03:00
kijai 55fc8a13e0 better mask cache, optimize 2025-07-18 13:55:06 +03:00
kijai 8919e65bab cache the mask 2025-07-18 01:43:45 +03:00
kijai a7166fc1a4 fix and rename args to be clearer 2025-07-18 01:28:01 +03:00
kijai 712eecf64b more compile friendly 2025-07-18 00:53:12 +03:00
kijai 605011a237 allow setting dense attention mode 2025-07-17 23:47:08 +03:00
kijai 9039d95721 cleanup and optimize 2025-07-17 22:40:26 +03:00
kijai 7c8020f8e8 works but slow 2025-07-17 21:49:14 +03:00
kijai 0417c8ea00 Update nodes.py 2025-07-17 20:48:16 +03:00
kijai 6ee7ad508e init (not working) 2025-07-17 20:47:02 +03:00
Jukka Seppänen 31b2c686cf Merge pull request #768 from Eikwang/fixunianimate
fix unianimate erro : OverflowError: cannot convert float infinity to…
2025-07-17 19:33:37 +03:00
kijai e011d031a8 Merge branch 'pr/799' 2025-07-17 19:33:10 +03:00
kijai 6e5b493b28 Create wanvideo_14B_pusa_I2V_example_01.json 2025-07-17 19:27:32 +03:00
kijai 8d7504144d Update nodes.py 2025-07-17 19:25:03 +03:00
kijai d8e90a94c1 more advanced input for Pusa 2025-07-17 19:10:24 +03:00
kijai 6bc53b771d refactor scheduler import 2025-07-17 15:56:55 +03:00
kijai 2263d02a0a basic Pusa support 2025-07-17 10:06:01 +03:00
Alexander Measure 88f50a8fd9 Update apply lora low_mem_load
Avoid creating 2 copies of WAN in VRAM when applying a LoRA by setting inplace_update=True on model.patch_weight_to_device.
2025-07-14 14:10:56 -04:00
kijai 17d48e3e45 rename VACE Model - VACE Module for clarity 2025-07-14 17:38:08 +03:00
kijai a37c0e6ac8 Update nodes.py 2025-07-14 17:28:17 +03:00
kijai 24d3de2126 Add EasyCache 2025-07-14 17:27:10 +03:00
kijai e1f5d185ee Use original Multitalk attention code with exactly 2 speakers 2025-07-14 17:06:48 +03:00
kijai f96128bf2c Make multitalk sampling actually use the seed... 2025-07-14 09:20:03 +03:00
kijai 8fd62b3e78 Multitalk wav2vec model doesn't actually need the .pt file
chinese-wav2vec2-base-fairseq-ckpt.pt can be safely deleted and won't be autodownloaded in the future
2025-07-14 09:13:41 +03:00
kijai 08669de287 Fix FantasyTalking when wav2vec loaded on offload_device 2025-07-13 23:52:19 +03:00
kijai b6d156afa5 Update nodes.py 2025-07-13 19:13:20 +03:00
kijai 1e4ed578ba Update pyproject.toml 2025-07-13 18:42:17 +03:00
kijai 365a1ce1d6 more chunking and refactoring 2025-07-12 20:18:52 +03:00
kijai 895c92febd Update model.py 2025-07-12 03:04:46 +03:00
kijai 039b211ef7 cleanup 2025-07-12 02:06:02 +03:00
kijai 127bfa2c99 refactor self attention 2025-07-11 23:00:59 +03:00
kijai 880194a1f9 bump version: 1.2.3 2025-07-11 01:13:25 +03:00
kijai 98395f94ce Update nodes_model_loading.py 2025-07-11 01:11:43 +03:00
kijai 1077d2323c fix max res selection 2025-07-10 23:41:54 +03:00
kijai 15f0e62e5f Add chunked RoPE option to reduce peak VRAM usage when not using torch.compile 2025-07-10 17:15:06 +03:00
kijai 49de5335c0 fix multi lora loader 2025-07-10 16:30:20 +03:00
kijai a9ba0e1cb4 small possible optimizations 2025-07-10 14:18:04 +03:00
kijai 5df3154112 Update requirements.txt 2025-07-10 12:53:37 +03:00
kijai 6e930d16cc Update nodes.py 2025-07-10 12:41:40 +03:00
kijai d71bed76af display possible lora metadata 2025-07-10 12:33:40 +03:00
kijai da2574d1bb WanVideoVACEStartToEndFrame fixes 2025-07-09 16:17:01 +03:00
kijai 456d04e318 fix vid2vid 2025-07-08 19:14:01 +03:00
kijai b65f161e8b Fix gguf + lora alpha scaling 2025-07-08 16:47:41 +03:00
kijai ce0bae1839 revert this 2025-07-08 16:23:29 +03:00
kijai 948805b6e5 remove print 2025-07-08 16:22:43 +03:00
astink b55ec9425b fix unianimate erro : OverflowError: cannot convert float infinity to integer 2025-07-08 20:14:22 +08:00
kijaiandkabachuha da98636599 FreeInit
https: //github.com/TianxingWu/FreeInit
Co-Authored-By: kabachuha <14872007+kabachuha@users.noreply.github.com>
2025-07-08 15:09:54 +03:00
kijai 8500514ef9 gguf + lora application fix 2025-07-08 12:53:29 +03:00
kijai 20e9914645 Update gguf.py 2025-07-08 12:33:50 +03:00
kijai 372e33bd5b possible fix for some gguf + torch.compile issues 2025-07-08 11:32:27 +03:00
kijai a57d7aa002 lora loading progressbar 2025-07-08 09:27:58 +03:00
kijai d16e2aa2ff Apply possible LoRA alpha with GGUF too 2025-07-08 09:27:22 +03:00
kijai b355f2b839 Update requirements.txt 2025-07-07 19:04:31 +03:00
Jukka Seppänen aa6fc5b12e Update readme.md 2025-07-07 17:22:06 +03:00
kijai 53ac085fc0 fix loop decode 2025-07-07 16:57:19 +03:00
kijai 5e1a5c6cf6 these won't work together 2025-07-07 01:44:06 +03:00
kijai 79ed360009 fix Q5 etc. 2025-07-07 01:31:27 +03:00
kijai ea183da9ba bump version 2025-07-06 23:58:22 +03:00
kijai d903ea7d3d small fixes, add comfy pbar for model loading 2025-07-06 23:55:37 +03:00
kijai a35142b006 Update nodes_model_loading.py 2025-07-06 23:20:04 +03:00
kijai 6ae5531bad Basic GGUF support (yes, really) 2025-07-06 21:37:45 +03:00
kijai ee1250c6bb fix cfg multitalk 2025-07-06 01:04:17 +03:00
Jukka Seppänen 20af8955c3 Merge pull request #687 from peteromallet/main
feat: Add ExtractStartFramesForContinuations node
2025-07-04 20:04:08 +03:00
kijai 79134368b4 bump version 2025-07-04 20:00:43 +03:00
kijai 78cc3050e0 update FLF2V example 2025-07-04 20:00:21 +03:00
kijai 490b0a8ead better context window progress bar, add VAE decode progress bar 2025-07-04 20:00:03 +03:00
kijai 32e5c4d200 update examples 2025-07-04 17:19:51 +03:00
kijai 3d7801cee4 Update nodes.py 2025-07-04 16:57:20 +03:00
kijai ca3c58bb1a Merge branch 'main' into multitalk 2025-07-04 16:43:53 +03:00
kijai 811dd32713 Update wanvideo_ATI_testing_01.json 2025-07-04 16:43:19 +03:00
kijai 74c6d25486 update WanVideoVACEStartToEndFrame 2025-07-04 16:43:13 +03:00
kijai 974dd656da some offloading for multitalk 2025-07-03 21:56:19 +03:00
kijai 4e6a85a148 Merge branch 'main' into multitalk 2025-07-03 21:56:06 +03:00
kijai 70e3c6b524 update 2025-07-03 21:41:44 +03:00
kijai eb6de2e8d0 Update nodes.py 2025-07-03 19:18:28 +03:00
kijaiandRudra-ai-coder 475f371016 Multiple talkers
Initial commit, works but needs more utility for the mask creation.

Based mostly on Rudra-ai-coder's modifications.

Co-Authored-By: Rudra-ai-coder <177262225+rudra-ai-coder@users.noreply.github.com>
2025-07-03 19:03:12 +03:00
kijai 06b932792f Update nodes.py 2025-07-02 17:30:35 +03:00
kijai 74a441ef93 Update nodes.py 2025-07-02 17:24:57 +03:00
kijai de279dd060 update 2025-07-02 17:15:18 +03:00
kijai ebe3068a8d progress bar 2025-07-02 16:59:45 +03:00
kijai 8e9fb0e11c Update nodes.py 2025-07-02 16:47:19 +03:00
kijai d687833c15 Create wanvideo_multitalk_test_02.json 2025-07-02 16:47:15 +03:00
kijai 0a11c67a0c MultiTalk sampling
New node (WanVideoImageToVideoMultiTalk) that also enables the continuous sampling method from the original code. Compared to context windows this works better for shorter clips, but degrades longer it goes.
2025-07-02 16:18:04 +03:00
kijai 49430f900b Fix for transformers 4.53.0 2025-07-01 12:00:28 +03:00
kijai f621d2a5fc Update model.py 2025-06-30 20:31:39 +03:00
POM e33d8c84ec feat: Add ExtractStartFramesForContinuations node 2025-06-23 00:35:45 +02:00
kijai 8479624614 fix split prompting 2025-06-20 17:31:01 +03:00
kijai 2da36b7fda context window progress bar
better than nothing, and allow cancelling mid context
2025-06-20 17:29:27 +03:00
kijai b605308687 Update nodes.py 2025-06-20 16:33:55 +03:00
kijai 547bd4e45a Update nodes.py 2025-06-20 16:18:07 +03:00
kijai 4334ed58b4 Update nodes.py 2025-06-20 15:26:13 +03:00
kijai c44eca13a9 Update nodes.py 2025-06-20 15:22:25 +03:00
kijai c3ab39b068 Update requirements.txt 2025-06-20 15:18:13 +03:00
kijai 157fbb1b7f Merge branch 'main' into multitalk 2025-06-20 13:53:31 +03:00
kijai 2b96d9cef2 WanVideoLoraSelectMulti 2025-06-20 13:53:13 +03:00
kijai 8387a64b15 bump version 2025-06-20 10:57:13 +03:00
kijai 38c98a6476 Create wanvideo_multitalk_test_01.json 2025-06-20 01:08:16 +03:00
kijai 263333a0d1 tiny vae adjustment 2025-06-19 20:38:23 +03:00
kijai 09b4c3a865 support context windows, torch compile fixes 2025-06-19 20:29:32 +03:00
kijai f3614e6720 Update nodes.py 2025-06-19 00:56:36 +03:00
kijai 96e8914b33 Fix audio cfg 2025-06-18 21:19:24 +03:00
kijai 1fe72d27aa loudness norm 2025-06-18 18:35:06 +03:00
kijai b5cae8bf1a Fix FLF2V model 2025-06-18 15:49:01 +03:00
kijai 58104b620f init 2025-06-18 15:41:53 +03:00
Jukka Seppänen 058286fc0f Merge pull request #657 from ObiLeek/main
refactor(decode): Replace min-max scaling with standard [-1, 1] conversion
2025-06-17 15:14:03 +03:00
tomas e8daa145e4 refactor(decode): Replace min-max scaling with standard [-1, 1] conversion 2025-06-17 13:30:22 +02:00
kijai dd1f3c3ecb fix 2025-06-17 12:14:56 +03:00
kijai b322192d93 Support minimaxremover
https://minimax-remover.github.io/
2025-06-17 09:45:12 +03:00
kijai c3ee35f3ec Update basic_flowmatch.py 2025-06-13 13:06:12 +03:00
kijai ad43eed0d0 Add MagCache and consolidate cache_args to support both
Not really getting great results yet with MagCache, but at least pretty much on bar with TeaCache so it can be an option that hopefully improves in time
2025-06-13 12:31:13 +03:00
kijai e7c39757f6 Allow compiling crossattn
I don't remember why I even disabled this... seems to work, big effect with NAG
2025-06-12 19:04:15 +03:00
kijai f6df3d75ef Merge branch 'pr/643' 2025-06-12 18:32:20 +03:00
kijai 34fecc3063 Update nodes.py 2025-06-12 18:19:52 +03:00
kijai d34cabc985 Cleanup, create own node for NAG 2025-06-12 18:18:28 +03:00
kijai 23796233df Some fixes and adjustments 2025-06-12 17:26:54 +03:00
kijai bac29263da Update nodes.py 2025-06-12 16:39:36 +03:00
kabachuha 803ba1cd72 rm commented out code 2025-06-12 15:34:11 +03:00
kabachuha e8d3f1cca5 fixup when context lens are none 2025-06-12 15:30:30 +03:00
kabachuha bdf7a3d02a pass nag scale into the model 2025-06-12 15:26:59 +03:00
kabachuha c4040426fc send negative embeds into nag 2025-06-12 14:50:12 +03:00
kabachuha 0ac366ba03 add nag to attention internal parts 2025-06-12 14:43:51 +03:00
202 changed files with 582310 additions and 14852 deletions
-1
View File
@@ -1 +0,0 @@
github: [kijai]
+2 -1
View File
@@ -10,4 +10,5 @@ logs/
tools/
.vscode/
convert_*
*.pt
*.pt
*.pth
+157
View File
@@ -0,0 +1,157 @@
from einops import rearrange
import torch
import torch.nn as nn
import torch.nn.functional as F
CACHE_T = 2
class RMS_norm(nn.Module):
def __init__(self, dim, channel_first=True, images=True, bias=False):
super().__init__()
broadcastable_dims = (1, 1, 1) if not images else (1, 1)
shape = (dim, *broadcastable_dims) if channel_first else (dim,)
self.channel_first = channel_first
self.scale = dim**0.5
self.gamma = nn.Parameter(torch.ones(shape))
self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.
def forward(self, x):
return F.normalize(
x, dim=(1 if self.channel_first else
-1)) * self.scale * self.gamma + self.bias
class CausalConv3d(nn.Conv3d):
"""
Causal 3d convolusion.
"""
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self._padding = (self.padding[2], self.padding[2], self.padding[1],
self.padding[1], 2 * self.padding[0], 0)
self.padding = (0, 0, 0)
def forward(self, x, cache_x=None):
padding = list(self._padding)
if cache_x is not None and self._padding[4] > 0:
cache_x = cache_x.to(x.device)
x = torch.cat([cache_x, x], dim=2)
padding[4] -= cache_x.shape[2]
x = F.pad(x, padding, mode='replicate')
return super().forward(x)
class PixelShuffle3d(nn.Module):
def __init__(self, ff, hh, ww):
super().__init__()
self.ff = ff
self.hh = hh
self.ww = ww
def forward(self, x):
# x: (B, C, F, H, W)
return rearrange(x,
'b c (f ff) (h hh) (w ww) -> b (c ff hh ww) f h w',
ff=self.ff, hh=self.hh, ww=self.ww)
class Buffer_LQ4x_Proj(nn.Module):
def __init__(self, in_dim, out_dim, layer_num=30):
super().__init__()
self.ff = 1
self.hh = 16
self.ww = 16
self.hidden_dim1 = 2048
self.hidden_dim2 = 3072
self.layer_num = layer_num
self.pixel_shuffle = PixelShuffle3d(self.ff, self.hh, self.ww)
self.conv1 = CausalConv3d(in_dim*self.ff*self.hh*self.ww, self.hidden_dim1, (4, 3, 3), stride=(2, 1, 1), padding=(1, 1, 1)) # f -> f/2 h -> h w -> w
self.norm1 = RMS_norm(self.hidden_dim1, images=False)
self.act1 = nn.SiLU()
self.conv2 = CausalConv3d(self.hidden_dim1, self.hidden_dim2, (4, 3, 3), stride=(2, 1, 1), padding=(1, 1, 1)) # f -> f/2 h -> h w -> w
self.norm2 = RMS_norm(self.hidden_dim2, images=False)
self.act2 = nn.SiLU()
self.linear_layers = nn.ModuleList([nn.Linear(self.hidden_dim2, out_dim) for _ in range(layer_num)])
self.clip_idx = 0
def forward(self, video):
self.clear_cache()
# x: (B, C, F, H, W)
t = video.shape[2]
iter_ = 1 + (t - 1) // 4
first_frame = video[:, :, :1, :, :].repeat(1, 1, 3, 1, 1)
video = torch.cat([first_frame, video], dim=2)
out_x = []
for i in range(iter_):
x = self.pixel_shuffle(video[:,:,i*4:(i+1)*4,:,:])
cache1_x = x[:, :, -CACHE_T:, :, :].clone()
self.cache['conv1'] = cache1_x
x = self.conv1(x, self.cache['conv1'])
x = self.norm1(x)
x = self.act1(x)
cache2_x = x[:, :, -CACHE_T:, :, :].clone()
self.cache['conv2'] = cache2_x
if i == 0:
continue
x = self.conv2(x, self.cache['conv2'])
x = self.norm2(x)
x = self.act2(x)
out_x.append(x)
out_x = torch.cat(out_x, dim = 2)
out_x = rearrange(out_x, 'b c f h w -> b (f h w) c')
outputs = []
for i in range(self.layer_num):
outputs.append(self.linear_layers[i](out_x))
self.clear_cache()
return outputs
def clear_cache(self):
self.cache = {}
self.cache['conv1'] = None
self.cache['conv2'] = None
self.clip_idx = 0
def stream_forward(self, video_clip):
if self.clip_idx == 0:
# self.clear_cache()
first_frame = video_clip[:, :, :1, :, :].repeat(1, 1, 3, 1, 1)
video_clip = torch.cat([first_frame, video_clip], dim=2)
x = self.pixel_shuffle(video_clip)
cache1_x = x[:, :, -CACHE_T:, :, :].clone()
self.cache['conv1'] = cache1_x
x = self.conv1(x, self.cache['conv1'])
x = self.norm1(x)
x = self.act1(x)
cache2_x = x[:, :, -CACHE_T:, :, :].clone()
self.cache['conv2'] = cache2_x
self.clip_idx += 1
return None
else:
x = self.pixel_shuffle(video_clip)
cache1_x = x[:, :, -CACHE_T:, :, :].clone()
self.cache['conv1'] = cache1_x
x = self.conv1(x, self.cache['conv1'])
x = self.norm1(x)
x = self.act1(x)
cache2_x = x[:, :, -CACHE_T:, :, :].clone()
self.cache['conv2'] = cache2_x
x = self.conv2(x, self.cache['conv2'])
x = self.norm2(x)
x = self.act2(x)
out_x = rearrange(x, 'b c f h w -> b (f h w) c')
outputs = []
for i in range(self.layer_num):
outputs.append(self.linear_layers[i](out_x))
self.clip_idx += 1
return outputs
+261
View File
@@ -0,0 +1,261 @@
"""
Tiny AutoEncoder for Hunyuan Video (Decoder-only, pruned)
- Encoder removed
- Transplant/widening helpers removed
- Deepening (IdentityConv2d+ReLU) is now built into the decoder structure itself
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from tqdm.auto import tqdm
from collections import namedtuple
from einops import rearrange
import torch.nn.init as init
DecoderResult = namedtuple("DecoderResult", ("frame", "memory"))
TWorkItem = namedtuple("TWorkItem", ("input_tensor", "block_index"))
# ----------------------------
# Utility / building blocks
# ----------------------------
class IdentityConv2d(nn.Conv2d):
"""Same-shape Conv2d initialized to identity (Dirac)."""
def __init__(self, C, kernel_size=3, bias=False):
pad = kernel_size // 2
super().__init__(C, C, kernel_size, padding=pad, bias=bias)
with torch.no_grad():
init.dirac_(self.weight)
if self.bias is not None:
self.bias.zero_()
def conv(n_in, n_out, **kwargs):
return nn.Conv2d(n_in, n_out, 3, padding=1, **kwargs)
class Clamp(nn.Module):
def forward(self, x):
return torch.tanh(x / 3) * 3
class MemBlock(nn.Module):
def __init__(self, n_in, n_out):
super().__init__()
self.conv = nn.Sequential(
conv(n_in * 2, n_out), nn.ReLU(inplace=True),
conv(n_out, n_out), nn.ReLU(inplace=True),
conv(n_out, n_out)
)
self.skip = nn.Conv2d(n_in, n_out, 1, bias=False) if n_in != n_out else nn.Identity()
self.act = nn.ReLU(inplace=True)
def forward(self, x, past):
return self.act(self.conv(torch.cat([x, past], 1)) + self.skip(x))
class TPool(nn.Module):
def __init__(self, n_f, stride):
super().__init__()
self.stride = stride
self.conv = nn.Conv2d(n_f*stride, n_f, 1, bias=False)
def forward(self, x):
_NT, C, H, W = x.shape
return self.conv(x.reshape(-1, self.stride * C, H, W))
class TGrow(nn.Module):
def __init__(self, n_f, stride):
super().__init__()
self.stride = stride
self.conv = nn.Conv2d(n_f, n_f*stride, 1, bias=False)
def forward(self, x):
_NT, C, H, W = x.shape
x = self.conv(x)
return x.reshape(-1, C, H, W)
class PixelShuffle3d(nn.Module):
def __init__(self, ff, hh, ww):
super().__init__()
self.ff = ff
self.hh = hh
self.ww = ww
def forward(self, x):
# x: (B, C, F, H, W)
B, C, F, H, W = x.shape
if F % self.ff != 0:
first_frame = x[:, :, 0:1, :, :].repeat(1, 1, self.ff - F % self.ff, 1, 1)
x = torch.cat([first_frame, x], dim=2)
return rearrange(
x,
'b c (f ff) (h hh) (w ww) -> b (c ff hh ww) f h w',
ff=self.ff, hh=self.hh, ww=self.ww
).transpose(1, 2)
# ----------------------------
# Generic NTCHW graph executor (kept; used by decoder)
# ----------------------------
def apply_model_with_memblocks(model, x, parallel, show_progress_bar, mem=None):
"""
Apply a sequential model with memblocks to the given input.
Args:
- model: nn.Sequential of blocks to apply
- x: input data, of dimensions NTCHW
- parallel: if True, parallelize over timesteps (fast but uses O(T) memory)
if False, each timestep will be processed sequentially (slow but uses O(1) memory)
- show_progress_bar: if True, enables tqdm progressbar display
Returns NTCHW tensor of output data.
"""
assert x.ndim == 5, f"TAEHV operates on NTCHW tensors, but got {x.ndim}-dim tensor"
N, T, C, H, W = x.shape
if parallel:
x = x.reshape(N*T, C, H, W)
for b in tqdm(model, disable=not show_progress_bar):
if isinstance(b, MemBlock):
NT, C, H, W = x.shape
T = NT // N
_x = x.reshape(N, T, C, H, W)
mem = F.pad(_x, (0,0,0,0,0,0,1,0), value=0)[:,:T].reshape(x.shape)
x = b(x, mem)
else:
x = b(x)
NT, C, H, W = x.shape
T = NT // N
x = x.view(N, T, C, H, W)
else:
out = []
work_queue = [TWorkItem(xt, 0) for t, xt in enumerate(x.reshape(N, T * C, H, W).chunk(T, dim=1))]
progress_bar = tqdm(range(T), disable=not show_progress_bar)
while work_queue:
xt, i = work_queue.pop(0)
if i == 0:
progress_bar.update(1)
if i == len(model):
out.append(xt)
else:
b = model[i]
if isinstance(b, MemBlock):
if mem[i] is None:
xt_new = b(xt, xt * 0)
mem[i] = xt
else:
xt_new = b(xt, mem[i])
mem[i].copy_(xt)
work_queue.insert(0, TWorkItem(xt_new, i+1))
elif isinstance(b, TPool):
if mem[i] is None:
mem[i] = []
mem[i].append(xt)
if len(mem[i]) > b.stride:
raise ValueError("TPool internal state invalid.")
elif len(mem[i]) == b.stride:
N_, C_, H_, W_ = xt.shape
xt = b(torch.cat(mem[i], 1).view(N_*b.stride, C_, H_, W_))
mem[i] = []
work_queue.insert(0, TWorkItem(xt, i+1))
elif isinstance(b, TGrow):
xt = b(xt)
NT, C_, H_, W_ = xt.shape
for xt_next in reversed(xt.view(N, b.stride*C_, H_, W_).chunk(b.stride, 1)):
work_queue.insert(0, TWorkItem(xt_next, i+1))
else:
xt = b(xt)
work_queue.insert(0, TWorkItem(xt, i+1))
progress_bar.close()
x = torch.stack(out, 1)
return x, mem
# ----------------------------
# Decoder-only TAEHV
# ----------------------------
class TAEHV(nn.Module):
image_channels = 3
def __init__(
self,
decoder_time_upscale=(True, True),
decoder_space_upscale=(True, True, True),
channels = [256, 128, 64, 64],
latent_channels = 16,
dtype=torch.float32
):
"""Initialize TAEHV (decoder-only) with built-in deepening after every ReLU.
Deepening config: how_many_each=1, k=3 (fixed as requested).
"""
super().__init__()
self.dtype = dtype
self.latent_channels = latent_channels
n_f = channels
self.frames_to_trim = 2**sum(decoder_time_upscale) - 1
# Build the decoder "skeleton"
base_decoder = nn.Sequential(
Clamp(), conv(self.latent_channels, n_f[0]), nn.ReLU(inplace=True),
MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]),
nn.Upsample(scale_factor=2 if decoder_space_upscale[0] else 1),
TGrow(n_f[0], 1),
conv(n_f[0], n_f[1], bias=False),
MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]),
nn.Upsample(scale_factor=2 if decoder_space_upscale[1] else 1),
TGrow(n_f[1], 2 if decoder_time_upscale[0] else 1),
conv(n_f[1], n_f[2], bias=False),
MemBlock(n_f[2], n_f[2]), MemBlock(n_f[2], n_f[2]), MemBlock(n_f[2], n_f[2]),
nn.Upsample(scale_factor=2 if decoder_space_upscale[2] else 1),
TGrow(n_f[2], 2 if decoder_time_upscale[1] else 1),
conv(n_f[2], n_f[3], bias=False),
nn.ReLU(inplace=True), conv(n_f[3], TAEHV.image_channels),
)
# Inline deepening: insert (IdentityConv2d(k=3) + ReLU) after every ReLU
self.decoder = self._apply_identity_deepen(base_decoder, how_many_each=1, k=3)
self.pixel_shuffle = PixelShuffle3d(4, 8, 8)
# Initialize decoder mem state
self.clean_mem()
@staticmethod
def _apply_identity_deepen(decoder: nn.Sequential, how_many_each=1, k=3) -> nn.Sequential:
"""Return a new Sequential where every nn.ReLU is followed by how_many_each*(IdentityConv2d(k)+ReLU)."""
new_layers = []
for b in decoder:
new_layers.append(b)
if isinstance(b, nn.ReLU):
# Deduce channel count from preceding layer
C = None
if len(new_layers) >= 2 and isinstance(new_layers[-2], nn.Conv2d):
C = new_layers[-2].out_channels
elif len(new_layers) >= 2 and isinstance(new_layers[-2], MemBlock):
C = new_layers[-2].conv[-1].out_channels
if C is not None:
for _ in range(how_many_each):
new_layers.append(IdentityConv2d(C, kernel_size=k, bias=False))
new_layers.append(nn.ReLU(inplace=True))
return nn.Sequential(*new_layers)
def decode_video(self, x, parallel=False, show_progress_bar=False, cond=None):
"""Decode a sequence of frames from latents.
x: NTCHW latent tensor; returns NTCHW RGB in ~[0, 1].
"""
trim_flag = self.mem[-8] is None # keeps original relative check
if cond is not None:
shuffled = self.pixel_shuffle(cond.to(x))
x = torch.cat([shuffled[:, :x.shape[1]], x], dim=2)
x, self.mem = apply_model_with_memblocks(self.decoder, x, parallel, show_progress_bar, mem=self.mem)
self.clean_mem()
if trim_flag:
return x[:, self.frames_to_trim:]
return x
def clean_mem(self):
self.mem = [None] * len(self.decoder)
def build_tcdecoder(new_channels = [512, 256, 128, 128], device="cuda", dtype=torch.bfloat16, new_latent_channels=None):
big = TAEHV(channels=new_channels, latent_channels=new_latent_channels, dtype=dtype).to(device).to(dtype)
return big
+71
View File
@@ -0,0 +1,71 @@
import folder_paths
import torch
from comfy.utils import load_torch_file
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoAddFlashVSRInput:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"images": ("IMAGE", {"tooltip": "Low-res video frames to enhance"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.01, "tooltip": "Strength to apply the FlashVSR latent"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, images, strength):
updated = dict(embeds)
updated["flashvsr_LQ_images"] = images
updated["flashvsr_strength"] = strength
return (updated,)
class WanVideoFlashVSRDecoderLoader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model_name": (folder_paths.get_filename_list("vae"), {"tooltip": "These models are loaded from 'ComfyUI/models/vae'"}),
},
"optional": {
"precision": (["fp16", "fp32", "bf16"],
{"default": "bf16"}
),
}
}
RETURN_TYPES = ("WANVAE",)
RETURN_NAMES = ("vae", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Loads Wan VAE model from 'ComfyUI/models/vae'"
def loadmodel(self, model_name, precision):
from .TCDecoder import build_tcdecoder
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
model_path = folder_paths.get_full_path("vae", model_name)
sd = load_torch_file(model_path, safe_load=True)
TCDecoder = build_tcdecoder(new_channels=[512, 256, 128, 128], new_latent_channels=16+768, dtype=dtype)
TCDecoder.load_state_dict(sd, strict=True)
TCDecoder.to(dtype)
return (TCDecoder,)
NODE_CLASS_MAPPINGS = {
"WanVideoAddFlashVSRInput": WanVideoAddFlashVSRInput,
"WanVideoFlashVSRDecoderLoader": WanVideoFlashVSRDecoderLoader,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddFlashVSRInput": "WanVideo Add FlashVSR Input",
"WanVideoFlashVSRDecoderLoader": "WanVideo FlashVSR Decoder Loader",
}
+87
View File
@@ -0,0 +1,87 @@
import torch
from einops import rearrange
from torch import nn
from einops import rearrange
class WanRMSNorm(nn.Module):
def __init__(self, dim, eps=1e-5):
super().__init__()
self.dim = dim
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
r"""
Args:
x(Tensor): Shape [B, L, C]
"""
return self._norm(x.to(self.weight.dtype)) * self.weight
def _norm(self, x):
return x * (torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)).to(x.dtype)
class DummyAdapterLayer(nn.Module):
def __init__(self, layer):
super().__init__()
self.layer = layer
def forward(self, *args, **kwargs):
return self.layer(*args, **kwargs)
class AudioProjModel(nn.Module):
def __init__(
self,
seq_len=5,
blocks=13, # add a new parameter blocks
channels=768, # add a new parameter channels
intermediate_dim=512,
output_dim=1536,
context_tokens=16,
):
super().__init__()
self.seq_len = seq_len
self.blocks = blocks
self.channels = channels
self.input_dim = seq_len * blocks * channels # update input_dim to be the product of blocks and channels.
self.intermediate_dim = intermediate_dim
self.context_tokens = context_tokens
self.output_dim = output_dim
# define multiple linear layers
self.audio_proj_glob_1 = DummyAdapterLayer(nn.Linear(self.input_dim, intermediate_dim))
self.audio_proj_glob_2 = DummyAdapterLayer(nn.Linear(intermediate_dim, intermediate_dim))
self.audio_proj_glob_3 = DummyAdapterLayer(nn.Linear(intermediate_dim, context_tokens * output_dim))
self.audio_proj_glob_norm = DummyAdapterLayer(nn.LayerNorm(output_dim))
self.initialize_weights()
def initialize_weights(self):
# Initialize transformer layers:
def _basic_init(module):
if isinstance(module, nn.Linear):
torch.nn.init.xavier_uniform_(module.weight)
if module.bias is not None:
nn.init.constant_(module.bias, 0)
self.apply(_basic_init)
def forward(self, audio_embeds):
video_length = audio_embeds.shape[1]
audio_embeds = rearrange(audio_embeds, "bz f w b c -> (bz f) w b c")
batch_size, window_size, blocks, channels = audio_embeds.shape
audio_embeds = audio_embeds.view(batch_size, window_size * blocks * channels)
audio_embeds = torch.relu(self.audio_proj_glob_1(audio_embeds))
audio_embeds = torch.relu(self.audio_proj_glob_2(audio_embeds))
context_tokens = self.audio_proj_glob_3(audio_embeds).reshape(batch_size, self.context_tokens, self.output_dim)
context_tokens = self.audio_proj_glob_norm(context_tokens.to(self.audio_proj_glob_norm.layer.weight.dtype)).to(audio_embeds.dtype)
context_tokens = rearrange(context_tokens, "(bz f) m c -> bz f m c", f=video_length)
return context_tokens
+287
View File
@@ -0,0 +1,287 @@
import folder_paths
import torch
import torch.nn.functional as F
import os
import json
import torchaudio
from comfy.utils import load_torch_file, common_upscale
import comfy.model_management as mm
from accelerate import init_empty_weights
from ..utils import set_module_tensor_to_device, log
from ..nodes import WanVideoEncodeLatentBatch
script_directory = os.path.dirname(os.path.abspath(__file__))
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
def linear_interpolation_fps(features, input_fps, output_fps, output_len=None):
features = features.transpose(1, 2) # [1, C, T]
seq_len = features.shape[2] / float(input_fps)
if output_len is None:
output_len = int(seq_len * output_fps)
output_features = F.interpolate(features, size=output_len, align_corners=True, mode='linear')
return output_features.transpose(1, 2)
def get_audio_emb_window(audio_emb, frame_num, frame0_idx, audio_shift=2):
zero_audio_embed = torch.zeros((audio_emb.shape[1], audio_emb.shape[2]), dtype=audio_emb.dtype, device=audio_emb.device)
zero_audio_embed_3 = torch.zeros((3, audio_emb.shape[1], audio_emb.shape[2]), dtype=audio_emb.dtype, device=audio_emb.device)
iter_ = 1 + (frame_num - 1) // 4
audio_emb_wind = []
for lt_i in range(iter_):
if lt_i == 0:
st = frame0_idx + lt_i - 2
ed = frame0_idx + lt_i + 3
wind_feat = torch.stack([
audio_emb[i] if (0 <= i < audio_emb.shape[0]) else zero_audio_embed
for i in range(st, ed)
], dim=0)
wind_feat = torch.cat((zero_audio_embed_3, wind_feat), dim=0)
else:
st = frame0_idx + 1 + 4 * (lt_i - 1) - audio_shift
ed = frame0_idx + 1 + 4 * lt_i + audio_shift
wind_feat = torch.stack([
audio_emb[i] if (0 <= i < audio_emb.shape[0]) else zero_audio_embed
for i in range(st, ed)
], dim=0)
audio_emb_wind.append(wind_feat)
audio_emb_wind = torch.stack(audio_emb_wind, dim=0)
return audio_emb_wind, ed - audio_shift
class WhisperModelLoader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": (folder_paths.get_filename_list("audio_encoders"), {"tooltip": "These models are loaded from the 'ComfyUI/models/audio_encoders' folder",}),
"base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}),
"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"}),
},
}
RETURN_TYPES = ("WHISPERMODEL",)
RETURN_NAMES = ("whisper_model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, model, base_precision, load_device):
from transformers import WhisperConfig, WhisperModel, WhisperFeatureExtractor
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]
if load_device == "offload_device":
transformer_load_device = offload_device
else:
transformer_load_device = device
config_path = os.path.join(script_directory, "whisper_config.json")
whisper_config = WhisperConfig(**json.load(open(config_path)))
with init_empty_weights():
whisper = WhisperModel(whisper_config).eval()
whisper.decoder = None # we only need the encoder
feature_extractor_config = {
"chunk_length": 30,
"feature_extractor_type": "WhisperFeatureExtractor",
"feature_size": 128,
"hop_length": 160,
"n_fft": 400,
"n_samples": 480000,
"nb_max_frames": 3000,
"padding_side": "right",
"padding_value": 0.0,
"processor_class": "WhisperProcessor",
"return_attention_mask": False,
"sampling_rate": 16000
}
feature_extractor = WhisperFeatureExtractor(**feature_extractor_config)
model_path = folder_paths.get_full_path_or_raise("audio_encoders", model)
sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True)
for name, param in whisper.named_parameters():
key = "model." + name
value=sd[key]
set_module_tensor_to_device(whisper, name, device=offload_device, dtype=base_dtype, value=value)
whisper_model = {
"feature_extractor": feature_extractor,
"model": whisper,
"dtype": base_dtype,
}
return (whisper_model,)
class HuMoEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"num_frames": ("INT", {"default": 81, "min": -1, "max": 10000, "step": 1, "tooltip": "The total frame count to generate."}),
"width": ("INT", {"default": 832, "min": 64, "max": 4096, "step": 16}),
"height": ("INT", {"default": 480, "min": 64, "max": 4096, "step": 16}),
"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: slower inference but more motion is allowed"}),
"audio_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "The percent of the video to start applying audio conditioning"}),
"audio_end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "The percent of the video to stop applying audio conditioning"})
},
"optional" : {
"whisper_model": ("WHISPERMODEL",),
"vae": ("WANVAE", ),
"reference_images": ("IMAGE", {"tooltip": "reference images for the humo model"}),
"audio": ("AUDIO",),
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", )
RETURN_NAMES = ("image_embeds", )
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, num_frames, width, height, audio_scale, audio_cfg_scale, audio_start_percent, audio_end_percent, whisper_model=None, vae=None, reference_images=None, audio=None, tiled_vae=False):
if reference_images is not None and vae is None:
raise ValueError("VAE is required when reference images are provided")
if whisper_model is None and audio is not None:
raise ValueError("Whisper model is required when audio is provided")
model = whisper_model["model"]
feature_extractor = whisper_model["feature_extractor"]
dtype = whisper_model["dtype"]
sampling_rate = 16000
if audio is not None:
audio_input = audio["waveform"][0]
sample_rate = audio["sample_rate"]
if sample_rate != sampling_rate:
audio_input = torchaudio.functional.resample(audio_input, sample_rate, sampling_rate)
if audio_input.shape[1] == 2:
audio_input = audio_input.mean(dim=0, keepdim=False)
else:
audio_input = audio_input[0]
model.to(device)
audio_len = len(audio_input) // 640
# feature extraction
audio_features = []
window = 750*640
for i in range(0, len(audio_input), window):
audio_feature = feature_extractor(audio_input[i:i+window], sampling_rate=sampling_rate, return_tensors="pt").input_features
audio_features.append(audio_feature)
audio_features = torch.cat(audio_features, dim=-1).to(device, dtype)
# preprocess
window = 3000
audio_prompts = []
for i in range(0, audio_features.shape[-1], window):
audio_prompt = model.encoder(audio_features[:,:,i:i+window], output_hidden_states=True).hidden_states
audio_prompt = torch.stack(audio_prompt, dim=2)
audio_prompts.append(audio_prompt)
model.to(offload_device)
audio_prompts = torch.cat(audio_prompts, dim=1)
audio_prompts = audio_prompts[:,:audio_len*2]
feat0 = linear_interpolation_fps(audio_prompts[:, :, 0: 8].mean(dim=2), 50, 25)
feat1 = linear_interpolation_fps(audio_prompts[:, :, 8: 16].mean(dim=2), 50, 25)
feat2 = linear_interpolation_fps(audio_prompts[:, :, 16: 24].mean(dim=2), 50, 25)
feat3 = linear_interpolation_fps(audio_prompts[:, :, 24: 32].mean(dim=2), 50, 25)
feat4 = linear_interpolation_fps(audio_prompts[:, :, 32], 50, 25)
audio_emb = torch.stack([feat0, feat1, feat2, feat3, feat4], dim=2)[0] # [T, 5, 1280]
else:
audio_emb = torch.zeros(num_frames, 5, 1280, device=device)
audio_len = num_frames
pixel_frame_num = num_frames if num_frames != -1 else audio_len
pixel_frame_num = 4 * ((pixel_frame_num - 1) // 4) + 1
latent_frame_num = (pixel_frame_num - 1) // 4 + 1
log.info(f"HuMo set to generate {pixel_frame_num} frames")
#audio_emb, _ = get_audio_emb_window(audio_emb, pixel_frame_num, frame0_idx=0)
num_refs = 0
if reference_images is not None:
if reference_images.shape[1] != height or reference_images.shape[2] != width:
reference_images_in = common_upscale(reference_images.movedim(-1, 1), width, height, "lanczos", "disabled").movedim(1, -1)
else:
reference_images_in = reference_images
samples, = WanVideoEncodeLatentBatch.encode(self, vae, reference_images_in, tiled_vae, None, None, None, None)
samples = samples["samples"].transpose(0, 2).squeeze(0)
num_refs = samples.shape[1]
vae.to(device)
zero_frames = torch.zeros(1, 3, pixel_frame_num + 4*num_refs, height, width, device=device, dtype=vae.dtype)
zero_latents = vae.encode(zero_frames, device=device, tiled=tiled_vae)[0].to(offload_device)
vae.to(offload_device)
mm.soft_empty_cache()
target_shape = (16, latent_frame_num + num_refs, height // 8, width // 8)
mask = torch.ones(4, target_shape[1], target_shape[2], target_shape[3], device=offload_device, dtype=vae.dtype)
if reference_images is not None:
mask[:,:-num_refs] = 0
image_cond = torch.cat([zero_latents[:, :(target_shape[1]-num_refs)], samples], dim=1)
#zero_audio_pad = torch.zeros(num_refs, *audio_emb.shape[1:]).to(audio_emb.device)
#audio_emb = torch.cat([audio_emb, zero_audio_pad], dim=0)
else:
image_cond = zero_latents
mask = torch.zeros_like(mask)
image_cond = torch.cat([mask, image_cond], dim=0)
image_cond_neg = torch.cat([mask, zero_latents], dim=0)
embeds = {
"humo_audio_emb": audio_emb,
"humo_audio_emb_neg": torch.zeros_like(audio_emb, dtype=audio_emb.dtype, device=audio_emb.device),
"humo_image_cond": image_cond,
"humo_image_cond_neg": image_cond_neg,
"humo_reference_count": num_refs,
"target_shape": target_shape,
"num_frames": pixel_frame_num,
"humo_audio_scale": audio_scale,
"humo_audio_cfg_scale": audio_cfg_scale,
"humo_start_percent": audio_start_percent,
"humo_end_percent": audio_end_percent,
}
return (embeds, )
class WanVideoCombineEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds_1": ("WANVIDIMAGE_EMBEDS",),
"embeds_2": ("WANVIDIMAGE_EMBEDS",),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
def add(self, embeds_1, embeds_2):
# Combine the two sets of embeds
combined = {**embeds_1, **embeds_2}
return (combined,)
NODE_CLASS_MAPPINGS = {
"WhisperModelLoader": WhisperModelLoader,
"HuMoEmbeds": HuMoEmbeds,
"WanVideoCombineEmbeds": WanVideoCombineEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WhisperModelLoader": "Whisper Model Loader",
"HuMoEmbeds": "HuMo Embeds",
"WanVideoCombineEmbeds": "WanVideo Combine Embeds",
}
+50
View File
@@ -0,0 +1,50 @@
{
"_name_or_path": "openai/whisper-large-v3",
"activation_dropout": 0.0,
"activation_function": "gelu",
"apply_spec_augment": false,
"architectures": [
"WhisperForConditionalGeneration"
],
"attention_dropout": 0.0,
"begin_suppress_tokens": [
220,
50257
],
"bos_token_id": 50257,
"classifier_proj_size": 256,
"d_model": 1280,
"decoder_attention_heads": 20,
"decoder_ffn_dim": 5120,
"decoder_layerdrop": 0.0,
"decoder_layers": 32,
"decoder_start_token_id": 50258,
"dropout": 0.0,
"encoder_attention_heads": 20,
"encoder_ffn_dim": 5120,
"encoder_layerdrop": 0.0,
"encoder_layers": 32,
"eos_token_id": 50257,
"init_std": 0.02,
"is_encoder_decoder": true,
"mask_feature_length": 10,
"mask_feature_min_masks": 0,
"mask_feature_prob": 0.0,
"mask_time_length": 10,
"mask_time_min_masks": 2,
"mask_time_prob": 0.05,
"max_length": 448,
"max_source_positions": 1500,
"max_target_positions": 448,
"median_filter_width": 7,
"model_type": "whisper",
"num_hidden_layers": 32,
"num_mel_bins": 128,
"pad_token_id": 50256,
"scale_embedding": false,
"torch_dtype": "float16",
"transformers_version": "4.36.0.dev0",
"use_cache": true,
"use_weighted_layer_sum": false,
"vocab_size": 51866
}
+212
View File
@@ -0,0 +1,212 @@
import torch.nn as nn
import torch.nn.functional as F
import torch
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):
def __init__(
self,
dim: int,
hidden_dim: int,
multiple_of: int = 256,
):
super().__init__()
hidden_dim = int(2 * hidden_dim / 3)
hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)
self.dim = dim
self.hidden_dim = hidden_dim
self.w1 = nn.Linear(dim, hidden_dim, bias=False)
self.w2 = nn.Linear(hidden_dim, dim, bias=False)
self.w3 = nn.Linear(dim, hidden_dim, bias=False)
def forward(self, x):
return self.w2(F.silu(self.w1(x)) * self.w3(x))
class TimestepEmbedder(nn.Module):
"""
Embeds scalar timesteps into vector representations.
"""
def __init__(self, t_embed_dim, frequency_embedding_size=256):
super().__init__()
self.t_embed_dim = t_embed_dim
self.frequency_embedding_size = frequency_embedding_size
self.mlp = nn.Sequential(
nn.Linear(frequency_embedding_size, t_embed_dim, bias=True),
nn.SiLU(),
nn.Linear(t_embed_dim, t_embed_dim, bias=True),
)
@staticmethod
def timestep_embedding(t, dim, max_period=10000):
"""
Create sinusoidal timestep embeddings.
:param t: a 1-D Tensor of N indices, one per batch element.
These may be fractional.
:param dim: the dimension of the output.
:param max_period: controls the minimum frequency of the embeddings.
:return: an (N, D) Tensor of positional embeddings.
"""
half = dim // 2
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half)
freqs = freqs.to(device=t.device)
args = t[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
return embedding
def forward(self, t, dtype):
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
if t_freq.dtype != dtype:
t_freq = t_freq.to(dtype)
t_emb = self.mlp(t_freq)
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
File diff suppressed because it is too large Load Diff
+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",
}
BIN
View File
Binary file not shown.
BIN
View File
Binary file not shown.
+213
View File
@@ -0,0 +1,213 @@
import cv2
import math
import torch
import numpy as np
from PIL import Image
from torchvision import transforms
def intrinsic_matrix_from_field_of_view(imshape, fov_degrees:float =55 ): # nlf default fov_degrees 55
imshape = np.array(imshape)
fov_radians = fov_degrees * np.array(np.pi / 180)
larger_side = np.max(imshape)
focal_length = larger_side / (np.tan(fov_radians / 2) * 2)
# intrinsic_matrix 3*3
return np.array([
[focal_length, 0, imshape[1] / 2],
[0, focal_length, imshape[0] / 2],
[0, 0, 1],
])
def p3d_to_p2d(point_3d, height, width): # point3d n*1024*3
camera_matrix = intrinsic_matrix_from_field_of_view((height,width))
camera_matrix = np.expand_dims(camera_matrix, axis=0)
camera_matrix = np.expand_dims(camera_matrix, axis=0) # 1*1*3*3
point_3d = np.expand_dims(point_3d,axis=-1) # n*1024*3*1
point_2d = (camera_matrix@point_3d).squeeze(-1)
point_2d[:,:,:2] = point_2d[:,:,:2]/point_2d[:,:,2:3]
return point_2d[:,:,:] # n*1024*2
def get_pose_images(smpl_data, offset):
pose_images = []
for data in smpl_data:
if isinstance(data, np.ndarray):
joints3d = data
else:
joints3d = data.numpy()
canvas = np.zeros(shape=(offset[0], offset[1], 3), dtype=np.uint8)
joints3d = p3d_to_p2d(joints3d, offset[0], offset[1])
canvas = draw_3d_points(canvas, joints3d[0], stickwidth=int(offset[1]/350))
pose_images.append(Image.fromarray(canvas))
return pose_images
def get_control_conditions(poses, h, w, stick_width=1.0, point_radius=2, style="original"):
control_images = []
for idx, pose in enumerate(poses):
canvas = np.zeros(shape=(h, w, 3), dtype=np.uint8)
try:
joints3d = p3d_to_p2d(pose, h, w)
if style == "original":
canvas = draw_3d_points(
canvas,
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))
# Image.fromarray(resized_canvas).save(f'tmp/{idx}_pose.jpg')
control_images.append(resized_canvas)
except Exception:
control_images.append(Image.fromarray(canvas))
control_pixel_values = np.array(control_images)
control_pixel_values = torch.from_numpy(control_pixel_values).contiguous() / 255.
return control_pixel_values
def draw_3d_points(canvas, points, stickwidth=2, r=2, draw_line=True):
colors = [
[255, 0, 0], # 0
[0, 255, 0], # 1
[0, 0, 255], # 2
[255, 0, 255], # 3
[255, 255, 0], # 4
[85, 255, 0], # 5
[0, 75, 255], # 6
[0, 255, 85], # 7
[0, 255, 170], # 8
[170, 0, 255], # 9
[85, 0, 255], # 10
[0, 85, 255], # 11
[0, 255, 255], # 12
[85, 0, 255], # 13
[170, 0, 255], # 14
[255, 0, 255], # 15
[255, 0, 170], # 16
[255, 0, 85], # 17
]
connetions = [
[15,12],[12, 16],[16, 18],[18, 20],[20, 22],
[12,17],[17,19],[19,21],
[21,23],[12,9],[9,6],
[6,3],[3,0],[0,1],
[1,4],[4,7],[7,10],[0,2],[2,5],[5,8],[8,11]
]
connection_colors = [
[255, 0, 0], # 0
[0, 255, 0], # 1
[0, 0, 255], # 2
[255, 255, 0], # 3
[255, 0, 255], # 4
[0, 255, 0], # 5
[0, 85, 255], # 6
[255, 175, 0], # 7
[0, 0, 255], # 8
[255, 85, 0], # 9
[0, 255, 85], # 10
[255, 0, 255], # 11
[255, 0, 0], # 12
[0, 175, 255], # 13
[255, 255, 0], # 14
[0, 0, 255], # 15
[0, 255, 0], # 16
]
# draw point
for i in range(len(points)):
x,y = points[i][0:2]
x,y = int(x),int(y)
if i==13 or i == 14:
continue
cv2.circle(canvas, (x, y), r, colors[i%17], thickness=-1)
# draw line
if draw_line:
for i in range(len(connetions)):
point1_idx,point2_idx = connetions[i][0:2]
point1 = points[point1_idx]
point2 = points[point2_idx]
Y = [point2[0],point1[0]]
X = [point2[1],point1[1]]
mX = int(np.mean(X))
mY = int(np.mean(Y))
length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5
angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1]))
polygon = cv2.ellipse2Poly((mY, mX), (int(length / 2), stickwidth), int(angle), 0, 360, 1)
cv2.fillConvexPoly(canvas, polygon, connection_colors[i%17])
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
+1
View File
@@ -0,0 +1 @@
from .vqvae import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder
+329
View File
@@ -0,0 +1,329 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
class Encoder(nn.Module):
def __init__(
self,
in_channels=3,
mid_channels=[128, 512],
out_channels=3072,
downsample_time=[1, 1],
downsample_joint=[1, 1],
num_attention_heads=8,
attention_head_dim=64,
dim=3072,
):
super(Encoder, self).__init__()
self.conv_in = nn.Conv2d(in_channels, mid_channels[0], kernel_size=3, stride=1, padding=1)
self.resnet1 = nn.ModuleList([ResBlock(mid_channels[0], mid_channels[0]) for _ in range(3)])
self.downsample1 = Downsample(mid_channels[0], mid_channels[0], downsample_time[0], downsample_joint[0])
self.resnet2 = ResBlock(mid_channels[0], mid_channels[1])
self.resnet3 = nn.ModuleList([ResBlock(mid_channels[1], mid_channels[1]) for _ in range(3)])
self.downsample2 = Downsample(mid_channels[1], mid_channels[1], downsample_time[1], downsample_joint[1])
self.conv_out = nn.Conv2d(mid_channels[-1], out_channels, kernel_size=3, stride=1, padding=1)
def forward(self, x):
x = self.conv_in(x)
for resnet in self.resnet1:
x = resnet(x)
x = self.downsample1(x)
x = self.resnet2(x)
for resnet in self.resnet3:
x = resnet(x)
x = self.downsample2(x)
x = self.conv_out(x)
return x
class VectorQuantizer(nn.Module):
def __init__(self, nb_code, code_dim):
super().__init__()
self.nb_code = nb_code
self.code_dim = code_dim
self.mu = 0.99
self.reset_codebook()
self.reset_count = 0
self.usage = torch.zeros((self.nb_code, 1))
def reset_codebook(self):
self.init = False
self.code_sum = None
self.code_count = None
self.register_buffer('codebook', torch.zeros(self.nb_code, self.code_dim).cuda())
def _tile(self, x):
nb_code_x, code_dim = x.shape
if nb_code_x < self.nb_code:
n_repeats = (self.nb_code + nb_code_x - 1) // nb_code_x
std = 0.01 / np.sqrt(code_dim)
out = x.repeat(n_repeats, 1)
out = out + torch.randn_like(out) * std
else:
out = x
return out
def preprocess(self, x):
# [bs, c, f, j] -> [bs * f * j, c]
x = x.permute(0, 2, 3, 1).contiguous()
x = x.view(-1, x.shape[-1])
return x
def quantize(self, x):
# [bs * f * j, dim=3072]
# Calculate latent code x_l
k_w = self.codebook.t()
distance = torch.sum(x ** 2, dim=-1, keepdim=True) - 2 * torch.matmul(x, k_w) + torch.sum(k_w ** 2, dim=0, keepdim=True)
_, code_idx = torch.min(distance, dim=-1)
return code_idx
def dequantize(self, code_idx):
x = F.embedding(code_idx, self.codebook) # indexing: [bs * f * j, 32]
return x
def forward(self, x, return_vq=False):
bs, c, f, j = x.shape # SMPL data frames: [bs, 3072, f, j]
# Preprocess
x = self.preprocess(x)
# return x.view(bs, f*j, c).contiguous(), None
assert x.shape[-1] == self.code_dim
# quantize and dequantize through bottleneck
code_idx = self.quantize(x)
x_d = self.dequantize(code_idx)
# Loss
commit_loss = F.mse_loss(x, x_d.detach())
# Passthrough
x_d = x + (x_d - x).detach()
if return_vq:
return x_d.view(bs, f*j, c).contiguous(), commit_loss
# return (x_d, x_d.view(bs, f, j, c).permute(0, 3, 1, 2).contiguous()), commit_loss, perplexity
# Postprocess
x_d = x_d.view(bs, f, j, c).permute(0, 3, 1, 2).contiguous()
return x_d, commit_loss
class Decoder(nn.Module):
def __init__(
self,
in_channels=3072,
mid_channels=[512, 128],
out_channels=3,
upsample_rate=None,
frame_upsample_rate=[1.0, 1.0],
joint_upsample_rate=[1.0, 1.0],
dim=128,
attention_head_dim=64,
num_attention_heads=8,
):
super(Decoder, self).__init__()
self.conv_in = nn.Conv2d(in_channels, mid_channels[0], kernel_size=3, stride=1, padding=1)
self.resnet1 = nn.ModuleList([ResBlock(mid_channels[0], mid_channels[0]) for _ in range(3)])
self.upsample1 = Upsample(mid_channels[0], mid_channels[0], frame_upsample_rate=frame_upsample_rate[0], joint_upsample_rate=joint_upsample_rate[0])
self.resnet2 = ResBlock(mid_channels[0], mid_channels[1])
self.resnet3 = nn.ModuleList([ResBlock(mid_channels[1], mid_channels[1]) for _ in range(3)])
self.upsample2 = Upsample(mid_channels[1], mid_channels[1], frame_upsample_rate=frame_upsample_rate[1], joint_upsample_rate=joint_upsample_rate[1])
self.conv_out = nn.Conv2d(mid_channels[-1], out_channels, kernel_size=3, stride=1, padding=1)
def forward(self, x):
x = self.conv_in(x)
for resnet in self.resnet1:
x = resnet(x)
x = self.upsample1(x)
x = self.resnet2(x)
for resnet in self.resnet3:
x = resnet(x)
x = self.upsample2(x)
x = self.conv_out(x)
return x
class Upsample(nn.Module):
def __init__(
self,
in_channels,
out_channels,
upsample_rate=None,
frame_upsample_rate=None,
joint_upsample_rate=None,
):
super(Upsample, self).__init__()
self.upsampler = nn.Conv1d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
self.upsample_rate = upsample_rate
self.frame_upsample_rate = frame_upsample_rate
self.joint_upsample_rate = joint_upsample_rate
self.upsample_rate = upsample_rate
def forward(self, inputs):
if inputs.shape[2] > 1 and inputs.shape[2] % 2 == 1:
# split first frame
x_first, x_rest = inputs[:, :, 0], inputs[:, :, 1:]
if self.upsample_rate is not None:
# import pdb; pdb.set_trace()
x_first = F.interpolate(x_first, scale_factor=self.upsample_rate)
x_rest = F.interpolate(x_rest, scale_factor=self.upsample_rate)
else:
# import pdb; pdb.set_trace()
# x_first = F.interpolate(x_first, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="bilinear", align_corners=True)
x_rest = F.interpolate(x_rest, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="bilinear", align_corners=True)
x_first = x_first[:, :, None, :]
inputs = torch.cat([x_first, x_rest], dim=2)
elif inputs.shape[2] > 1:
if self.upsample_rate is not None:
inputs = F.interpolate(inputs, scale_factor=self.upsample_rate)
else:
inputs = F.interpolate(inputs, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="bilinear", align_corners=True)
else:
inputs = inputs.squeeze(2)
if self.upsample_rate is not None:
inputs = F.interpolate(inputs, scale_factor=self.upsample_rate)
else:
inputs = F.interpolate(inputs, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="linear", align_corners=True)
inputs = inputs[:, :, None, :, :]
b, c, t, j = inputs.shape
inputs = inputs.permute(0, 2, 1, 3).reshape(b * t, c, j)
inputs = self.upsampler(inputs)
inputs = inputs.reshape(b, t, *inputs.shape[1:]).permute(0, 2, 1, 3)
return inputs
class Downsample(nn.Module):
def __init__(
self,
in_channels,
out_channels,
frame_downsample_rate,
joint_downsample_rate
):
super(Downsample, self).__init__()
self.frame_downsample_rate = frame_downsample_rate
self.joint_downsample_rate = joint_downsample_rate
self.joint_downsample = nn.Conv1d(in_channels, out_channels, kernel_size=3, stride=self.joint_downsample_rate, padding=1)
def forward(self, x):
# (batch_size, channels, frames, joints) -> (batch_size * joints, channels, frames)
if self.frame_downsample_rate > 1:
batch_size, channels, frames, joints = x.shape
x = x.permute(0, 3, 1, 2).reshape(batch_size * joints, channels, frames)
if x.shape[-1] % 2 == 1:
x_first, x_rest = x[..., 0], x[..., 1:]
if x_rest.shape[-1] > 0:
# (batch_size * height * width, channels, frames - 1) -> (batch_size * height * width, channels, (frames - 1) // 2)
x_rest = F.avg_pool1d(x_rest, kernel_size=self.frame_downsample_rate, stride=self.frame_downsample_rate)
x = torch.cat([x_first[..., None], x_rest], dim=-1)
# (batch_size * joints, channels, (frames // 2) + 1) -> (batch_size, channels, (frames // 2) + 1, joints)
x = x.reshape(batch_size, joints, channels, x.shape[-1]).permute(0, 2, 3, 1)
else:
# (batch_size * joints, channels, frames) -> (batch_size * joints, channels, frames // 2)
x = F.avg_pool1d(x, kernel_size=2, stride=2)
# (batch_size * joints, channels, frames // 2) -> (batch_size, height, width, channels, frames // 2) -> (batch_size, channels, frames // 2, height, width)
x = x.reshape(batch_size, joints, channels, x.shape[-1]).permute(0, 2, 3, 1)
# Pad the tensor
# pad = (0, 1)
# x = F.pad(x, pad, mode="constant", value=0)
batch_size, channels, frames, joints = x.shape
# (batch_size, channels, frames, joints) -> (batch_size * frames, channels, joints)
x = x.permute(0, 2, 1, 3).reshape(batch_size * frames, channels, joints)
x = self.joint_downsample(x)
# (batch_size * frames, channels, joints) -> (batch_size, channels, frames, joints)
x = x.reshape(batch_size, frames, x.shape[1], x.shape[2]).permute(0, 2, 1, 3)
return x
class ResBlock(nn.Module):
def __init__(self,
in_channels,
out_channels,
group_num=32,
max_channels=512):
super(ResBlock, self).__init__()
skip = max(1, max_channels // out_channels - 1)
self.block = nn.Sequential(
nn.GroupNorm(group_num, in_channels, eps=1e-06, affine=True),
nn.SiLU(),
nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=skip, dilation=skip),
nn.GroupNorm(group_num, out_channels, eps=1e-06, affine=True),
nn.SiLU(),
nn.Conv2d(out_channels, out_channels, kernel_size=1, stride=1, padding=0),
)
self.conv_short = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0) if in_channels != out_channels else nn.Identity()
def forward(self, x):
hidden_states = self.block(x)
if hidden_states.shape != x.shape:
x = self.conv_short(x)
x = x + hidden_states
return x
class SMPL_VQVAE(nn.Module):
def __init__(self, encoder, decoder, vq):
super(SMPL_VQVAE, self).__init__()
self.encoder = encoder
self.decoder = decoder
self.vq = vq
def to(self, device):
self.encoder = self.encoder.to(device)
self.decoder = self.decoder.to(device)
self.vq = self.vq.to(device)
self.device = device
return self
def encdec_slice_frames(self, x, frame_batch_size, encdec, return_vq):
num_frames = x.shape[2]
remaining_frames = num_frames % frame_batch_size
x_output = []
for i in range(num_frames // frame_batch_size):
remaining_frames = num_frames % frame_batch_size
start_frame = frame_batch_size * i + (0 if i == 0 else remaining_frames)
end_frame = frame_batch_size * (i + 1) + remaining_frames
x_intermediate = x[:, :, start_frame:end_frame]
x_intermediate = encdec(x_intermediate)
x_output.append(x_intermediate)
if encdec == self.encoder and self.vq is not None:
x_output, loss = self.vq(torch.cat(x_output, dim=2), return_vq=return_vq)
return x_output, loss
else:
return torch.cat(x_output, dim=2), None, None
def forward(self, x, return_vq=False):
x = x.permute(0, 3, 1, 2)
x, loss = self.encdec_slice_frames(x, frame_batch_size=8, encdec=self.encoder, return_vq=return_vq)
if return_vq:
return x, loss
x, _, _ = self.encdec_slice_frames(x, frame_batch_size=2, encdec=self.decoder, return_vq=return_vq)
x = x.permute(0, 2, 3, 1)
return x, loss
+193
View File
@@ -0,0 +1,193 @@
import torch
import numpy as np
from typing import Union, Tuple
def get_1d_rotary_pos_embed(
dim: int,
pos: Union[np.ndarray, int],
theta: float = 10000.0,
use_real=False,
linear_factor=1.0,
ntk_factor=1.0,
repeat_interleave_real=True,
freqs_dtype=torch.float32, # torch.float32, torch.float64 (flux)
):
"""
Precompute the frequency tensor for complex exponentials (cis) with given dimensions.
This function calculates a frequency tensor with complex exponentials using the given dimension 'dim' and the end
index 'end'. The 'theta' parameter scales the frequencies. The returned tensor contains complex values in complex64
data type.
Args:
dim (`int`): Dimension of the frequency tensor.
pos (`np.ndarray` or `int`): Position indices for the frequency tensor. [S] or scalar
theta (`float`, *optional*, defaults to 10000.0):
Scaling factor for frequency computation. Defaults to 10000.0.
use_real (`bool`, *optional*):
If True, return real part and imaginary part separately. Otherwise, return complex numbers.
linear_factor (`float`, *optional*, defaults to 1.0):
Scaling factor for the context extrapolation. Defaults to 1.0.
ntk_factor (`float`, *optional*, defaults to 1.0):
Scaling factor for the NTK-Aware RoPE. Defaults to 1.0.
repeat_interleave_real (`bool`, *optional*, defaults to `True`):
If `True` and `use_real`, real part and imaginary part are each interleaved with themselves to reach `dim`.
Otherwise, they are concateanted with themselves.
freqs_dtype (`torch.float32` or `torch.float64`, *optional*, defaults to `torch.float32`):
the dtype of the frequency tensor.
Returns:
`torch.Tensor`: Precomputed frequency tensor with complex exponentials. [S, D/2]
"""
assert dim % 2 == 0
if isinstance(pos, int):
pos = torch.arange(pos)
if isinstance(pos, np.ndarray):
pos = torch.from_numpy(pos) # type: ignore # [S]
theta = theta * ntk_factor
freqs = (
1.0
/ (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=pos.device)[: (dim // 2)] / dim))
/ linear_factor
) # [D/2]
freqs = torch.outer(pos, freqs) # type: ignore # [S, D/2]
if use_real and repeat_interleave_real:
freqs_cos = freqs.cos().repeat_interleave(2, dim=1).float() # [S, D]
freqs_sin = freqs.sin().repeat_interleave(2, dim=1).float() # [S, D]
return freqs_cos, freqs_sin
elif use_real:
freqs_cos = torch.cat([freqs.cos(), freqs.cos()], dim=-1).float() # [S, D]
freqs_sin = torch.cat([freqs.sin(), freqs.sin()], dim=-1).float() # [S, D]
return freqs_cos, freqs_sin
else:
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 # [S, D/2]
return freqs_cis
def get_3d_rotary_pos_embed(
embed_dim, crops_coords, grid_size, temporal_size, theta: int = 10000, use_real: bool = True
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""
RoPE for video tokens with 3D structure.
Args:
embed_dim: (`int`):
The embedding dimension size, corresponding to hidden_size_head.
crops_coords (`Tuple[int]`):
The top-left and bottom-right coordinates of the crop.
grid_size (`Tuple[int]`):
The grid size of the spatial positional embedding (height, width).
temporal_size (`int`):
The size of the temporal dimension.
theta (`float`):
Scaling factor for frequency computation.
Returns:
`torch.Tensor`: positional embedding with shape `(temporal_size * grid_size[0] * grid_size[1], embed_dim/2)`.
"""
if use_real is not True:
raise ValueError(" `use_real = False` is not currently supported for get_3d_rotary_pos_embed")
start, stop = crops_coords
grid_size_h, grid_size_w = grid_size
grid_h = np.linspace(start[0], stop[0], grid_size_h, endpoint=False, dtype=np.float32)
grid_w = np.linspace(start[1], stop[1], grid_size_w, endpoint=False, dtype=np.float32)
grid_t = np.linspace(0, temporal_size, temporal_size, endpoint=False, dtype=np.float32)
# Compute dimensions for each axis
dim_t = embed_dim // 4
dim_h = embed_dim // 8 * 3
dim_w = embed_dim // 8 * 3
# Temporal frequencies
freqs_t = get_1d_rotary_pos_embed(dim_t, grid_t, use_real=True)
# Spatial frequencies for height and width
freqs_h = get_1d_rotary_pos_embed(dim_h, grid_h, use_real=True)
freqs_w = get_1d_rotary_pos_embed(dim_w, grid_w, use_real=True)
# BroadCast and concatenate temporal and spaial frequencie (height and width) into a 3d tensor
def combine_time_height_width(freqs_t, freqs_h, freqs_w):
freqs_t = freqs_t[:, None, None, :].expand(
-1, grid_size_h, grid_size_w, -1
) # temporal_size, grid_size_h, grid_size_w, dim_t
freqs_h = freqs_h[None, :, None, :].expand(
temporal_size, -1, grid_size_w, -1
) # temporal_size, grid_size_h, grid_size_2, dim_h
freqs_w = freqs_w[None, None, :, :].expand(
temporal_size, grid_size_h, -1, -1
) # temporal_size, grid_size_h, grid_size_2, dim_w
freqs = torch.cat(
[freqs_t, freqs_h, freqs_w], dim=-1
) # temporal_size, grid_size_h, grid_size_w, (dim_t + dim_h + dim_w)
freqs = freqs.view(
temporal_size * grid_size_h * grid_size_w, -1
) # (temporal_size * grid_size_h * grid_size_w), (dim_t + dim_h + dim_w)
return freqs
t_cos, t_sin = freqs_t # both t_cos and t_sin has shape: temporal_size, dim_t
h_cos, h_sin = freqs_h # both h_cos and h_sin has shape: grid_size_h, dim_h
w_cos, w_sin = freqs_w # both w_cos and w_sin has shape: grid_size_w, dim_w
cos = combine_time_height_width(t_cos, h_cos, w_cos)
sin = combine_time_height_width(t_sin, h_sin, w_sin)
return cos, sin
def get_3d_motion_spatial_embed(
embed_dim: int, num_joints: int, joints_mean: np.ndarray, joints_std: np.ndarray, theta: float = 10000.0
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
assert embed_dim % 2 == 0 and embed_dim % 3 == 0
def create_rope_pe(dim, pos, freqs_dtype=torch.float32):
if isinstance(pos, np.ndarray):
pos = torch.from_numpy(pos)
freqs = (
1.0
/ (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=pos.device)[: (dim // 2)] / dim))
) # [D/2]
freqs = torch.outer(pos, freqs) # type: ignore # [S, D/2]
freqs_cos = freqs.cos().repeat_interleave(2, dim=1).float() # [S, D]
freqs_sin = freqs.sin().repeat_interleave(2, dim=1).float() # [S, D]
return freqs_cos, freqs_sin
pos_x = joints_mean[:, 0]
pos_y = joints_mean[:, 1]
pos_z = joints_mean[:, 2]
normalized_pos_x = (pos_x - pos_x.mean())
normalized_pos_y = (pos_y - pos_y.mean())
normalized_pos_z = (pos_z - pos_z.mean())
freqs_cos_x, freqs_sin_x = create_rope_pe(embed_dim // 3, normalized_pos_x)
freqs_cos_y, freqs_sin_y = create_rope_pe(embed_dim // 3, normalized_pos_y)
freqs_cos_z, freqs_sin_z = create_rope_pe(embed_dim // 3, normalized_pos_z)
freqs_cos = torch.cat([freqs_cos_x, freqs_cos_y, freqs_cos_z], dim=-1)
freqs_sin = torch.cat([freqs_sin_x, freqs_sin_y, freqs_sin_z], dim=-1)
return freqs_cos, freqs_sin
def prepare_motion_embeddings(num_frames, num_joints, joints_mean, joints_std, theta=10000, device='cuda'):
time_embed = get_1d_rotary_pos_embed(44, num_frames, theta, use_real=True)
time_embed_cos = time_embed[0][:, None, :].expand(-1, num_joints, -1).reshape(num_frames*num_joints, -1)
time_embed_sin = time_embed[1][:, None, :].expand(-1, num_joints, -1).reshape(num_frames*num_joints, -1)
spatial_motion_embed = get_3d_motion_spatial_embed(84, num_joints, joints_mean, joints_std, theta)
spatial_embed_cos = spatial_motion_embed[0][None, :, :].expand(num_frames, -1, -1).reshape(num_frames*num_joints, -1)
spatial_embed_sin = spatial_motion_embed[1][None, :, :].expand(num_frames, -1, -1).reshape(num_frames*num_joints, -1)
motion_embed_cos = torch.cat([time_embed_cos, spatial_embed_cos], dim=-1).to(device=device)
motion_embed_sin = torch.cat([time_embed_sin, spatial_embed_sin], dim=-1).to(device=device)
return motion_embed_cos, motion_embed_sin
def apply_rotary_emb(x, freqs_cis):
cos, sin = freqs_cis # [S, D]
cos = cos[None, None]
sin = sin[None, None]
cos, sin = cos.to(x.device), sin.to(x.device)
x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2]
x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3)
out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype)
return out
+340
View File
@@ -0,0 +1,340 @@
import os
import torch
from ..utils import log
import numpy as np
import comfy.model_management as mm
from comfy.utils import load_torch_file
import folder_paths
script_directory = os.path.dirname(os.path.abspath(__file__))
device = mm.get_torch_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")
folder_paths.add_model_folder_path("nlf", os.path.join(folder_paths.models_dir, "nlf"))
from .motion4d import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder
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:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"url": (model_list, {"default": "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"}),
},
}
RETURN_TYPES = ("NLFMODEL",)
RETURN_NAMES = ("nlf_model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
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):
log.info(f"Downloading NLF model to: {local_model_path}")
import requests
os.makedirs(os.path.dirname(local_model_path), exist_ok=True)
response = requests.get(url)
if response.status_code == 200:
with open(local_model_path, "wb") as f:
f.write(response.content)
else:
print("Failed to download file:", response.status_code)
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,)
class LoadNLFModel:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"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_NAMES = ("nlf_model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, nlf_model, warmup=True):
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,
class LoadVQVAE:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model_name": (folder_paths.get_filename_list("vae"), {"tooltip": "These models are loaded from 'ComfyUI/models/vae'"}),
},
}
RETURN_TYPES = ("VQVAE",)
RETURN_NAMES = ("vqvae", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, model_name):
model_path = folder_paths.get_full_path("vae", model_name)
vae_sd = load_torch_file(model_path, safe_load=True)
# Get motion tokenizer
motion_encoder = Encoder(
in_channels=3,
mid_channels=[128, 512],
out_channels=3072,
downsample_time=[2, 2],
downsample_joint=[1, 1]
)
motion_quant = VectorQuantizer(nb_code=8192, code_dim=3072)
motion_decoder = Decoder(
in_channels=3072,
mid_channels=[512, 128],
out_channels=3,
upsample_rate=2.0,
frame_upsample_rate=[2.0, 2.0],
joint_upsample_rate=[1.0, 1.0]
)
vqvae = SMPL_VQVAE(motion_encoder, motion_decoder, motion_quant).to(device)
vqvae.load_state_dict(vae_sd, strict=True)
return vqvae,
class MTVCrafterEncodePoses:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"vqvae": ("VQVAE", {"tooltip": "VQVAE model"}),
"poses": ("NLFPRED", {"tooltip": "Input poses for the model"}),
},
}
RETURN_TYPES = ("MTVCRAFTERMOTION", "NLFPRED")
RETURN_NAMES = ("mtvcrafter_motion", "pose_results")
FUNCTION = "encode"
CATEGORY = "WanVideoWrapper"
def encode(self, vqvae, poses):
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"))
smpl_poses = []
for pose in poses['joints3d_nonparam'][0]:
smpl_poses.append(pose[0].cpu().numpy())
smpl_poses = np.array(smpl_poses)
norm_poses = torch.tensor((smpl_poses - global_mean) / global_std).unsqueeze(0)
print(f"norm_poses shape: {norm_poses.shape}, dtype: {norm_poses.dtype}")
vqvae.to(device)
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
vqvae.to(offload_device)
poses_dict = {
'mtv_motion_tokens': motion_tokens,
'global_mean': global_mean,
'global_std': global_std
}
return poses_dict, recon_motion
class NLFPredict:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": ("NLFMODEL",),
"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", "BBOX",)
RETURN_NAMES = ("pose_results", "bboxes")
FUNCTION = "predict"
CATEGORY = "WanVideoWrapper"
def predict(self, model, images, per_batch=-1):
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 = {
'joints3d_nonparam': [all_joints3d_nonparam],
}
# Convert bboxes to list format: [x_min, y_min, x_max, y_max] for each detection
# Each box tensor is shape (1, 5) with [x_min, y_min, x_max, y_max, confidence]
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:
# Extract first 4 values (x_min, y_min, x_max, y_max), drop confidence
bbox_values = box[0, :4].cpu().tolist()
formatted_boxes.append(bbox_values)
return (pose_results, formatted_boxes)
class DrawNLFPoses:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"poses": ("NLFPRED", {"tooltip": "Input poses for the model"}),
"width": ("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_NAMES = ("image",)
FUNCTION = "predict"
CATEGORY = "WanVideoWrapper"
def predict(self, poses, width, height, stick_width=1.0, point_radius=2, style="original"):
from .draw_pose import get_control_conditions
if isinstance(poses, dict):
pose_input = poses['joints3d_nonparam'][0] if 'joints3d_nonparam' in poses else poses
else:
pose_input = poses
control_conditions = get_control_conditions(pose_input, height, width, stick_width=stick_width, point_radius=point_radius, style=style)
return (control_conditions,)
NODE_CLASS_MAPPINGS = {
"LoadNLFModel": LoadNLFModel,
"DownloadAndLoadNLFModel": DownloadAndLoadNLFModel,
"NLFPredict": NLFPredict,
"DrawNLFPoses": DrawNLFPoses,
"LoadVQVAE": LoadVQVAE,
"MTVCrafterEncodePoses": MTVCrafterEncodePoses
}
NODE_DISPLAY_NAME_MAPPINGS = {
"LoadNLFModel": "Load NLF Model",
"DownloadAndLoadNLFModel": "(Download)Load NLF Model",
"NLFPredict": "NLF Predict",
"DrawNLFPoses": "Draw NLF Poses",
"LoadVQVAE": "Load VQVAE",
"MTVCrafterEncodePoses": "MTV Crafter Encode Poses"
}
+48
View File
@@ -0,0 +1,48 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
class ChannelLastConv1d(nn.Conv1d):
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = x.permute(0, 2, 1)
x = super().forward(x)
x = x.permute(0, 2, 1)
return x
class ConvMLP(nn.Module):
def __init__(
self,
dim: int,
hidden_dim: int,
multiple_of: int = 256,
kernel_size: int = 3,
padding: int = 1,
):
"""
Initialize the FeedForward module.
Args:
dim (int): Input dimension.
hidden_dim (int): Hidden dimension of the feedforward layer.
multiple_of (int): Value to ensure hidden dimension is a multiple of this value.
Attributes:
w1 (ColumnParallelLinear): Linear transformation for the first layer.
w2 (RowParallelLinear): Linear transformation for the second layer.
w3 (ColumnParallelLinear): Linear transformation for the third layer.
"""
super().__init__()
hidden_dim = int(2 * hidden_dim / 3)
hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)
self.w1 = ChannelLastConv1d(dim, hidden_dim, bias=False, kernel_size=kernel_size, padding=padding)
self.w2 = ChannelLastConv1d(hidden_dim, dim, bias=False, kernel_size=kernel_size, padding=padding)
self.w3 = ChannelLastConv1d(dim, hidden_dim, bias=False, kernel_size=kernel_size, padding=padding)
def forward(self, x):
return self.w2(F.silu(self.w1(x)) * self.w3(x))
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2022 NVIDIA CORPORATION.
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+1
View File
@@ -0,0 +1 @@
from .bigvgan import BigVGAN
+120
View File
@@ -0,0 +1,120 @@
# Implementation adapted from https://github.com/EdwardDixon/snake under the MIT license.
# LICENSE is in incl_licenses directory.
import torch
from torch import nn, sin, pow
from torch.nn import Parameter
class Snake(nn.Module):
'''
Implementation of a sine-based periodic activation function
Shape:
- Input: (B, C, T)
- Output: (B, C, T), same shape as the input
Parameters:
- alpha - trainable parameter
References:
- This activation function is from this paper by Liu Ziyin, Tilman Hartwig, Masahito Ueda:
https://arxiv.org/abs/2006.08195
Examples:
>>> a1 = snake(256)
>>> x = torch.randn(256)
>>> x = a1(x)
'''
def __init__(self, in_features, alpha=1.0, alpha_trainable=True, alpha_logscale=False):
'''
Initialization.
INPUT:
- in_features: shape of the input
- alpha: trainable parameter
alpha is initialized to 1 by default, higher values = higher-frequency.
alpha will be trained along with the rest of your model.
'''
super(Snake, self).__init__()
self.in_features = in_features
# initialize alpha
self.alpha_logscale = alpha_logscale
if self.alpha_logscale: # log scale alphas initialized to zeros
self.alpha = Parameter(torch.zeros(in_features) * alpha)
else: # linear scale alphas initialized to ones
self.alpha = Parameter(torch.ones(in_features) * alpha)
self.alpha.requires_grad = alpha_trainable
self.no_div_by_zero = 0.000000001
def forward(self, x):
'''
Forward pass of the function.
Applies the function to the input elementwise.
Snake ∶= x + 1/a * sin^2 (xa)
'''
alpha = self.alpha.unsqueeze(0).unsqueeze(-1) # line up with x to [B, C, T]
if self.alpha_logscale:
alpha = torch.exp(alpha)
x = x + (1.0 / (alpha + self.no_div_by_zero)) * pow(sin(x * alpha), 2)
return x
class SnakeBeta(nn.Module):
'''
A modified Snake function which uses separate parameters for the magnitude of the periodic components
Shape:
- Input: (B, C, T)
- Output: (B, C, T), same shape as the input
Parameters:
- alpha - trainable parameter that controls frequency
- beta - trainable parameter that controls magnitude
References:
- This activation function is a modified version based on this paper by Liu Ziyin, Tilman Hartwig, Masahito Ueda:
https://arxiv.org/abs/2006.08195
Examples:
>>> a1 = snakebeta(256)
>>> x = torch.randn(256)
>>> x = a1(x)
'''
def __init__(self, in_features, alpha=1.0, alpha_trainable=True, alpha_logscale=False):
'''
Initialization.
INPUT:
- in_features: shape of the input
- alpha - trainable parameter that controls frequency
- beta - trainable parameter that controls magnitude
alpha is initialized to 1 by default, higher values = higher-frequency.
beta is initialized to 1 by default, higher values = higher-magnitude.
alpha will be trained along with the rest of your model.
'''
super(SnakeBeta, self).__init__()
self.in_features = in_features
# initialize alpha
self.alpha_logscale = alpha_logscale
if self.alpha_logscale: # log scale alphas initialized to zeros
self.alpha = Parameter(torch.zeros(in_features) * alpha)
self.beta = Parameter(torch.zeros(in_features) * alpha)
else: # linear scale alphas initialized to ones
self.alpha = Parameter(torch.ones(in_features) * alpha)
self.beta = Parameter(torch.ones(in_features) * alpha)
self.alpha.requires_grad = alpha_trainable
self.beta.requires_grad = alpha_trainable
self.no_div_by_zero = 0.000000001
def forward(self, x):
'''
Forward pass of the function.
Applies the function to the input elementwise.
SnakeBeta ∶= x + 1/b * sin^2 (xa)
'''
alpha = self.alpha.unsqueeze(0).unsqueeze(-1) # line up with x to [B, C, T]
beta = self.beta.unsqueeze(0).unsqueeze(-1)
if self.alpha_logscale:
alpha = torch.exp(alpha)
beta = torch.exp(beta)
x = x + (1.0 / (beta + self.no_div_by_zero)) * pow(sin(x * alpha), 2)
return x
+6
View File
@@ -0,0 +1,6 @@
# Adapted from https://github.com/junjun3518/alias-free-torch under the Apache License 2.0
# LICENSE is in incl_licenses directory.
from .filter import *
from .resample import *
from .act import *
+28
View File
@@ -0,0 +1,28 @@
# Adapted from https://github.com/junjun3518/alias-free-torch under the Apache License 2.0
# LICENSE is in incl_licenses directory.
import torch.nn as nn
from .resample import UpSample1d, DownSample1d
class Activation1d(nn.Module):
def __init__(self,
activation,
up_ratio: int = 2,
down_ratio: int = 2,
up_kernel_size: int = 12,
down_kernel_size: int = 12):
super().__init__()
self.up_ratio = up_ratio
self.down_ratio = down_ratio
self.act = activation
self.upsample = UpSample1d(up_ratio, up_kernel_size)
self.downsample = DownSample1d(down_ratio, down_kernel_size)
# x: [B,C,T]
def forward(self, x):
x = self.upsample(x)
x = self.act(x)
x = self.downsample(x)
return x
+95
View File
@@ -0,0 +1,95 @@
# Adapted from https://github.com/junjun3518/alias-free-torch under the Apache License 2.0
# LICENSE is in incl_licenses directory.
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
if 'sinc' in dir(torch):
sinc = torch.sinc
else:
# This code is adopted from adefossez's julius.core.sinc under the MIT License
# https://adefossez.github.io/julius/julius/core.html
# LICENSE is in incl_licenses directory.
def sinc(x: torch.Tensor):
"""
Implementation of sinc, i.e. sin(pi * x) / (pi * x)
__Warning__: Different to julius.sinc, the input is multiplied by `pi`!
"""
return torch.where(x == 0,
torch.tensor(1., device=x.device, dtype=x.dtype),
torch.sin(math.pi * x) / math.pi / x)
# This code is adopted from adefossez's julius.lowpass.LowPassFilters under the MIT License
# https://adefossez.github.io/julius/julius/lowpass.html
# LICENSE is in incl_licenses directory.
def kaiser_sinc_filter1d(cutoff, half_width, kernel_size): # return filter [1,1,kernel_size]
even = (kernel_size % 2 == 0)
half_size = kernel_size // 2
#For kaiser window
delta_f = 4 * half_width
A = 2.285 * (half_size - 1) * math.pi * delta_f + 7.95
if A > 50.:
beta = 0.1102 * (A - 8.7)
elif A >= 21.:
beta = 0.5842 * (A - 21)**0.4 + 0.07886 * (A - 21.)
else:
beta = 0.
window = torch.kaiser_window(kernel_size, beta=beta, periodic=False)
# ratio = 0.5/cutoff -> 2 * cutoff = 1 / ratio
if even:
time = (torch.arange(-half_size, half_size) + 0.5)
else:
time = torch.arange(kernel_size) - half_size
if cutoff == 0:
filter_ = torch.zeros_like(time)
else:
filter_ = 2 * cutoff * window * sinc(2 * cutoff * time)
# Normalize filter to have sum = 1, otherwise we will have a small leakage
# of the constant component in the input signal.
filter_ /= filter_.sum()
filter = filter_.view(1, 1, kernel_size)
return filter
class LowPassFilter1d(nn.Module):
def __init__(self,
cutoff=0.5,
half_width=0.6,
stride: int = 1,
padding: bool = True,
padding_mode: str = 'replicate',
kernel_size: int = 12):
# kernel_size should be even number for stylegan3 setup,
# in this implementation, odd number is also possible.
super().__init__()
if cutoff < -0.:
raise ValueError("Minimum cutoff must be larger than zero.")
if cutoff > 0.5:
raise ValueError("A cutoff above 0.5 does not make sense.")
self.kernel_size = kernel_size
self.even = (kernel_size % 2 == 0)
self.pad_left = kernel_size // 2 - int(self.even)
self.pad_right = kernel_size // 2
self.stride = stride
self.padding = padding
self.padding_mode = padding_mode
filter = kaiser_sinc_filter1d(cutoff, half_width, kernel_size)
self.register_buffer("filter", filter)
#input [B, C, T]
def forward(self, x):
_, C, _ = x.shape
if self.padding:
x = F.pad(x, (self.pad_left, self.pad_right),
mode=self.padding_mode)
out = F.conv1d(x, self.filter.expand(C, -1, -1),
stride=self.stride, groups=C)
return out
+49
View File
@@ -0,0 +1,49 @@
# Adapted from https://github.com/junjun3518/alias-free-torch under the Apache License 2.0
# LICENSE is in incl_licenses directory.
import torch.nn as nn
from torch.nn import functional as F
from .filter import LowPassFilter1d
from .filter import kaiser_sinc_filter1d
class UpSample1d(nn.Module):
def __init__(self, ratio=2, kernel_size=None):
super().__init__()
self.ratio = ratio
self.kernel_size = int(6 * ratio // 2) * 2 if kernel_size is None else kernel_size
self.stride = ratio
self.pad = self.kernel_size // ratio - 1
self.pad_left = self.pad * self.stride + (self.kernel_size - self.stride) // 2
self.pad_right = self.pad * self.stride + (self.kernel_size - self.stride + 1) // 2
filter = kaiser_sinc_filter1d(cutoff=0.5 / ratio,
half_width=0.6 / ratio,
kernel_size=self.kernel_size)
self.register_buffer("filter", filter)
# x: [B, C, T]
def forward(self, x):
_, C, _ = x.shape
x = F.pad(x, (self.pad, self.pad), mode='replicate')
x = self.ratio * F.conv_transpose1d(
x, self.filter.expand(C, -1, -1), stride=self.stride, groups=C)
x = x[..., self.pad_left:-self.pad_right]
return x
class DownSample1d(nn.Module):
def __init__(self, ratio=2, kernel_size=None):
super().__init__()
self.ratio = ratio
self.kernel_size = int(6 * ratio // 2) * 2 if kernel_size is None else kernel_size
self.lowpass = LowPassFilter1d(cutoff=0.5 / ratio,
half_width=0.6 / ratio,
stride=ratio,
kernel_size=self.kernel_size)
def forward(self, x):
xx = self.lowpass(x)
return xx
+62
View File
@@ -0,0 +1,62 @@
import torch
import torch.nn as nn
from types import SimpleNamespace
from .models import BigVGANVocoder
from comfy.utils import load_torch_file
# BigVGAN vocoder configuration
_bigvgan_vocoder_config = {
'resblock': '1',
'num_gpus': 0,
'batch_size': 64,
'num_mels': 80,
'learning_rate': 0.0001,
'adam_b1': 0.8,
'adam_b2': 0.99,
'lr_decay': 0.999,
'seed': 1234,
'upsample_rates': [4, 4, 2, 2, 2, 2],
'upsample_kernel_sizes': [8, 8, 4, 4, 4, 4],
'upsample_initial_channel': 1536,
'resblock_kernel_sizes': [3, 7, 11],
'resblock_dilation_sizes': [
[1, 3, 5],
[1, 3, 5],
[1, 3, 5]
],
'activation': 'snakebeta',
'snake_logscale': True,
'resolutions': [
[1024, 120, 600],
[2048, 240, 1200],
[512, 50, 240]
],
'mpd_reshapes': [2, 3, 5, 7, 11],
'use_spectral_norm': False,
'discriminator_channel_mult': 1,
}
class BigVGAN(nn.Module):
def __init__(self, ckpt_path):
super().__init__()
# Convert dictionary to namespace object for attribute access
vocoder_cfg = SimpleNamespace(**_bigvgan_vocoder_config)
self.vocoder = BigVGANVocoder(vocoder_cfg).eval()
vocoder_ckpt = load_torch_file(ckpt_path)
self.vocoder.load_state_dict(vocoder_ckpt)
self.weight_norm_removed = False
self.remove_weight_norm()
@torch.inference_mode()
def forward(self, x):
assert self.weight_norm_removed, 'call remove_weight_norm() before inference'
return self.vocoder(x)
def remove_weight_norm(self):
self.vocoder.remove_weight_norm()
self.weight_norm_removed = True
return self
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2020 Jungil Kong
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+21
View File
@@ -0,0 +1,21 @@
MIT License
Copyright (c) 2020 Edward Dixon
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+201
View File
@@ -0,0 +1,201 @@
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
1. Definitions.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
"Licensor" shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
"Legal Entity" shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
"control" means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
"You" (or "Your") shall mean an individual or Legal Entity
exercising permissions granted by this License.
"Source" form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
"Object" form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.
"Work" shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
"Derivative Works" shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
of this License, Derivative Works shall not include works that remain
separable from, or merely link (or bind by name) to the interfaces of,
the Work and Derivative Works thereof.
"Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, "submitted"
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
"Contributor" shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:
(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and
(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a "NOTICE" text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS
APPENDIX: How to apply the Apache License to your work.
To apply the Apache License to your work, attach the following
boilerplate notice, with the fields enclosed by brackets "[]"
replaced with your own identifying information. (Don't include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
Copyright [yyyy] [name of copyright owner]
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.
+29
View File
@@ -0,0 +1,29 @@
BSD 3-Clause License
Copyright (c) 2019, Seungwon Park 박승원
All rights reserved.
Redistribution and use in source and binary forms, with or without
modification, are permitted provided that the following conditions are met:
1. Redistributions of source code must retain the above copyright notice, this
list of conditions and the following disclaimer.
2. Redistributions in binary form must reproduce the above copyright notice,
this list of conditions and the following disclaimer in the documentation
and/or other materials provided with the distribution.
3. Neither the name of the copyright holder nor the names of its
contributors may be used to endorse or promote products derived from
this software without specific prior written permission.
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
+16
View File
@@ -0,0 +1,16 @@
Copyright 2020 Alexandre Défossez
Permission is hereby granted, free of charge, to any person obtaining a copy of this software and
associated documentation files (the "Software"), to deal in the Software without restriction,
including without limitation the rights to use, copy, modify, merge, publish, distribute,
sublicense, and/or sell copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all copies or
substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT
NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND
NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM,
DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
+255
View File
@@ -0,0 +1,255 @@
# Copyright (c) 2022 NVIDIA CORPORATION.
# Licensed under the MIT license.
# Adapted from https://github.com/jik876/hifi-gan under the MIT license.
# LICENSE is in incl_licenses directory.
import torch
import torch.nn as nn
from torch.nn import Conv1d, ConvTranspose1d
from torch.nn.utils.parametrizations import weight_norm
from torch.nn.utils.parametrize import remove_parametrizations
from . import activations
from .alias_free_torch import *
from .utils import get_padding, init_weights
LRELU_SLOPE = 0.1
class AMPBlock1(torch.nn.Module):
def __init__(self, h, channels, kernel_size=3, dilation=(1, 3, 5), activation=None):
super(AMPBlock1, self).__init__()
self.h = h
self.convs1 = nn.ModuleList([
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=dilation[0],
padding=get_padding(kernel_size, dilation[0]))),
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=dilation[1],
padding=get_padding(kernel_size, dilation[1]))),
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=dilation[2],
padding=get_padding(kernel_size, dilation[2])))
])
self.convs1.apply(init_weights)
self.convs2 = nn.ModuleList([
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=1,
padding=get_padding(kernel_size, 1))),
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=1,
padding=get_padding(kernel_size, 1))),
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=1,
padding=get_padding(kernel_size, 1)))
])
self.convs2.apply(init_weights)
self.num_layers = len(self.convs1) + len(self.convs2) # total number of conv layers
if activation == 'snake': # periodic nonlinearity with snake function and anti-aliasing
self.activations = nn.ModuleList([
Activation1d(
activation=activations.Snake(channels, alpha_logscale=h.snake_logscale))
for _ in range(self.num_layers)
])
elif activation == 'snakebeta': # periodic nonlinearity with snakebeta function and anti-aliasing
self.activations = nn.ModuleList([
Activation1d(
activation=activations.SnakeBeta(channels, alpha_logscale=h.snake_logscale))
for _ in range(self.num_layers)
])
else:
raise NotImplementedError(
"activation incorrectly specified. check the config file and look for 'activation'."
)
def forward(self, x):
acts1, acts2 = self.activations[::2], self.activations[1::2]
for c1, c2, a1, a2 in zip(self.convs1, self.convs2, acts1, acts2):
xt = a1(x)
xt = c1(xt)
xt = a2(xt)
xt = c2(xt)
x = xt + x
return x
def remove_weight_norm(self):
for l in self.convs1:
remove_parametrizations(l, 'weight')
for l in self.convs2:
remove_parametrizations(l, 'weight')
class AMPBlock2(torch.nn.Module):
def __init__(self, h, channels, kernel_size=3, dilation=(1, 3), activation=None):
super(AMPBlock2, self).__init__()
self.h = h
self.convs = nn.ModuleList([
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=dilation[0],
padding=get_padding(kernel_size, dilation[0]))),
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=dilation[1],
padding=get_padding(kernel_size, dilation[1])))
])
self.convs.apply(init_weights)
self.num_layers = len(self.convs) # total number of conv layers
if activation == 'snake': # periodic nonlinearity with snake function and anti-aliasing
self.activations = nn.ModuleList([
Activation1d(
activation=activations.Snake(channels, alpha_logscale=h.snake_logscale))
for _ in range(self.num_layers)
])
elif activation == 'snakebeta': # periodic nonlinearity with snakebeta function and anti-aliasing
self.activations = nn.ModuleList([
Activation1d(
activation=activations.SnakeBeta(channels, alpha_logscale=h.snake_logscale))
for _ in range(self.num_layers)
])
else:
raise NotImplementedError(
"activation incorrectly specified. check the config file and look for 'activation'."
)
def forward(self, x):
for c, a in zip(self.convs, self.activations):
xt = a(x)
xt = c(xt)
x = xt + x
return x
def remove_weight_norm(self):
for l in self.convs:
remove_parametrizations(l, 'weight')
class BigVGANVocoder(torch.nn.Module):
# this is our main BigVGAN model. Applies anti-aliased periodic activation for resblocks.
def __init__(self, h):
super().__init__()
self.h = h
self.num_kernels = len(h.resblock_kernel_sizes)
self.num_upsamples = len(h.upsample_rates)
# pre conv
self.conv_pre = weight_norm(Conv1d(h.num_mels, h.upsample_initial_channel, 7, 1, padding=3))
# define which AMPBlock to use. BigVGAN uses AMPBlock1 as default
resblock = AMPBlock1 if h.resblock == '1' else AMPBlock2
# transposed conv-based upsamplers. does not apply anti-aliasing
self.ups = nn.ModuleList()
for i, (u, k) in enumerate(zip(h.upsample_rates, h.upsample_kernel_sizes)):
self.ups.append(
nn.ModuleList([
weight_norm(
ConvTranspose1d(h.upsample_initial_channel // (2**i),
h.upsample_initial_channel // (2**(i + 1)),
k,
u,
padding=(k - u) // 2))
]))
# residual blocks using anti-aliased multi-periodicity composition modules (AMP)
self.resblocks = nn.ModuleList()
for i in range(len(self.ups)):
ch = h.upsample_initial_channel // (2**(i + 1))
for j, (k, d) in enumerate(zip(h.resblock_kernel_sizes, h.resblock_dilation_sizes)):
self.resblocks.append(resblock(h, ch, k, d, activation=h.activation))
# post conv
if h.activation == "snake": # periodic nonlinearity with snake function and anti-aliasing
activation_post = activations.Snake(ch, alpha_logscale=h.snake_logscale)
self.activation_post = Activation1d(activation=activation_post)
elif h.activation == "snakebeta": # periodic nonlinearity with snakebeta function and anti-aliasing
activation_post = activations.SnakeBeta(ch, alpha_logscale=h.snake_logscale)
self.activation_post = Activation1d(activation=activation_post)
else:
raise NotImplementedError(
"activation incorrectly specified. check the config file and look for 'activation'."
)
self.conv_post = weight_norm(Conv1d(ch, 1, 7, 1, padding=3))
# weight initialization
for i in range(len(self.ups)):
self.ups[i].apply(init_weights)
self.conv_post.apply(init_weights)
def forward(self, x):
# pre conv
x = self.conv_pre(x)
for i in range(self.num_upsamples):
# upsampling
for i_up in range(len(self.ups[i])):
x = self.ups[i][i_up](x)
# AMP blocks
xs = None
for j in range(self.num_kernels):
if xs is None:
xs = self.resblocks[i * self.num_kernels + j](x)
else:
xs += self.resblocks[i * self.num_kernels + j](x)
x = xs / self.num_kernels
# post conv
x = self.activation_post(x)
x = self.conv_post(x)
x = torch.tanh(x)
return x
def remove_weight_norm(self):
print('Removing weight norm...')
for l in self.ups:
for l_i in l:
remove_parametrizations(l_i, 'weight')
for l in self.resblocks:
l.remove_weight_norm()
remove_parametrizations(self.conv_pre, 'weight')
remove_parametrizations(self.conv_post, 'weight')
+20
View File
@@ -0,0 +1,20 @@
# Adapted from https://github.com/jik876/hifi-gan under the MIT license.
# LICENSE is in incl_licenses directory.
from torch.nn.utils.parametrizations import weight_norm
def init_weights(m, mean=0.0, std=0.01):
classname = m.__class__.__name__
if classname.find("Conv") != -1:
m.weight.data.normal_(mean, std)
def apply_weight_norm(m):
classname = m.__class__.__name__
if classname.find("Conv") != -1:
weight_norm(m)
def get_padding(kernel_size, dilation=1):
return int((kernel_size * dilation - dilation) / 2)
+212
View File
@@ -0,0 +1,212 @@
# Reference: # https://github.com/bytedance/Make-An-Audio-2
from typing import Literal
import torch
import torch.nn as nn
import numpy as np
# following is from librosa
def hz_to_mel(frequencies, *, htk = False):
frequencies = np.asanyarray(frequencies)
if htk:
mels: np.ndarray = 2595.0 * np.log10(1.0 + frequencies / 700.0)
return mels
# Fill in the linear part
f_min = 0.0
f_sp = 200.0 / 3
mels = (frequencies - f_min) / f_sp
# Fill in the log-scale part
min_log_hz = 1000.0 # beginning of log region (Hz)
min_log_mel = (min_log_hz - f_min) / f_sp # same (Mels)
logstep = np.log(6.4) / 27.0 # step size for log region
if frequencies.ndim:
# If we have array data, vectorize
log_t = frequencies >= min_log_hz
mels[log_t] = min_log_mel + np.log(frequencies[log_t] / min_log_hz) / logstep
elif frequencies >= min_log_hz:
# If we have scalar data, heck directly
mels = min_log_mel + np.log(frequencies / min_log_hz) / logstep
return mels
def mel_to_hz(mels, *, htk = False):
mels = np.asanyarray(mels)
if htk:
return 700.0 * (10.0 ** (mels / 2595.0) - 1.0)
# Fill in the linear scale
f_min = 0.0
f_sp = 200.0 / 3
freqs = f_min + f_sp * mels
# And now the nonlinear scale
min_log_hz = 1000.0 # beginning of log region (Hz)
min_log_mel = (min_log_hz - f_min) / f_sp # same (Mels)
logstep = np.log(6.4) / 27.0 # step size for log region
if mels.ndim:
# If we have vector data, vectorize
log_t = mels >= min_log_mel
freqs[log_t] = min_log_hz * np.exp(logstep * (mels[log_t] - min_log_mel))
elif mels >= min_log_mel:
# If we have scalar data, check directly
freqs = min_log_hz * np.exp(logstep * (mels - min_log_mel))
return freqs
def mel_frequencies(n_mels = 128, *, fmin = 0.0, fmax = 11025.0, htk = False):
min_mel = hz_to_mel(fmin, htk=htk)
max_mel = hz_to_mel(fmax, htk=htk)
mels = np.linspace(min_mel, max_mel, n_mels)
hz: np.ndarray = mel_to_hz(mels, htk=htk)
return hz
def librosa_mel_fn(
*,
sr: float,
n_fft: int,
n_mels: int = 128,
fmin: float = 0.0,
fmax = None,
htk = False,
norm = "slaney",
dtype = np.float32,
) -> np.ndarray:
if fmax is None:
fmax = float(sr) / 2
# Initialize the weights
n_mels = int(n_mels)
weights = np.zeros((n_mels, int(1 + n_fft // 2)), dtype=dtype)
# Center freqs of each FFT bin
fftfreqs = np.fft.rfftfreq(n=n_fft, d=1.0 / sr)
# 'Center freqs' of mel bands - uniformly spaced between limits
mel_f = mel_frequencies(n_mels + 2, fmin=fmin, fmax=fmax, htk=htk)
fdiff = np.diff(mel_f)
ramps = np.subtract.outer(mel_f, fftfreqs)
for i in range(n_mels):
# lower and upper slopes for all bins
lower = -ramps[i] / fdiff[i]
upper = ramps[i + 2] / fdiff[i + 1]
# .. then intersect them with each other and zero
weights[i] = np.maximum(0, np.minimum(lower, upper))
# Slaney-style mel is scaled to be approx constant energy per channel
enorm = 2.0 / (mel_f[2 : n_mels + 2] - mel_f[:n_mels])
weights *= enorm[:, np.newaxis]
return weights
def dynamic_range_compression_torch(x, C=1, clip_val=1e-5, *, norm_fn):
return norm_fn(torch.clamp(x, min=clip_val) * C)
def spectral_normalize_torch(magnitudes, norm_fn):
output = dynamic_range_compression_torch(magnitudes, norm_fn=norm_fn)
return output
class MelConverter(nn.Module):
def __init__(
self,
*,
sampling_rate: float,
n_fft: int,
num_mels: int,
hop_size: int,
win_size: int,
fmin: float,
fmax: float,
norm_fn,
):
super().__init__()
self.sampling_rate = sampling_rate
self.n_fft = n_fft
self.num_mels = num_mels
self.hop_size = hop_size
self.win_size = win_size
self.fmin = fmin
self.fmax = fmax
self.norm_fn = norm_fn
mel = librosa_mel_fn(sr=self.sampling_rate,
n_fft=self.n_fft,
n_mels=self.num_mels,
fmin=self.fmin,
fmax=self.fmax)
mel_basis = torch.from_numpy(mel).float()
hann_window = torch.hann_window(self.win_size)
self.register_buffer('mel_basis', mel_basis)
self.register_buffer('hann_window', hann_window)
@property
def device(self):
return self.mel_basis.device
def forward(self, waveform: torch.Tensor, center: bool = False) -> torch.Tensor:
waveform = waveform.clamp(min=-1., max=1.).to(self.device)
waveform = torch.nn.functional.pad(
waveform.unsqueeze(1),
[int((self.n_fft - self.hop_size) / 2),
int((self.n_fft - self.hop_size) / 2)],
mode='reflect')
waveform = waveform.squeeze(1)
spec = torch.stft(waveform,
self.n_fft,
hop_length=self.hop_size,
win_length=self.win_size,
window=self.hann_window,
center=center,
pad_mode='reflect',
normalized=False,
onesided=True,
return_complex=True)
spec = torch.view_as_real(spec)
spec = torch.sqrt(spec.pow(2).sum(-1) + (1e-9)).float()
spec = torch.matmul(self.mel_basis, spec)
spec = spectral_normalize_torch(spec, self.norm_fn)
return spec
def get_mel_converter(mode: Literal['16k', '44k']) -> MelConverter:
if mode == '16k':
return MelConverter(sampling_rate=16_000,
n_fft=1024,
num_mels=80,
hop_size=256,
win_size=1024,
fmin=0,
fmax=8_000,
norm_fn=torch.log10)
elif mode == '44k':
return MelConverter(sampling_rate=44_100,
n_fft=2048,
num_mels=128,
hop_size=512,
win_size=2048,
fmin=0,
fmax=44100 / 2,
norm_fn=torch.log)
else:
raise ValueError(f'Unknown mode: {mode}')
+267
View File
@@ -0,0 +1,267 @@
import torch
import torch.nn as nn
import folder_paths
import os
from .mel_converter import get_mel_converter
from .vae.autoencoder import AutoEncoderModule
from .vae.distributions import DiagonalGaussianDistribution
import torchaudio
from ..utils import log
from comfy import model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class FeaturesUtils(nn.Module):
def __init__(
self,
*,
tod_vae_ckpt: str,
bigvgan_vocoder_ckpt = None,
mode=['16k', '44k'],
need_vae_encoder: bool = True,
):
super().__init__()
self.mel_converter = get_mel_converter(mode)
self.tod = AutoEncoderModule(vae_ckpt_path=tod_vae_ckpt,
vocoder_ckpt_path=bigvgan_vocoder_ckpt,
mode=mode,
need_vae_encoder=need_vae_encoder)
def encode_audio(self, x) -> DiagonalGaussianDistribution:
assert self.tod is not None, 'VAE is not loaded'
# x: (B * L)
mel = self.mel_converter(x)
dist = self.tod.encode(mel)
return dist
def vocode(self, mel: torch.Tensor) -> torch.Tensor:
assert self.tod is not None, 'VAE is not loaded'
return self.tod.vocode(mel)
def decode(self, z: torch.Tensor) -> torch.Tensor:
assert self.tod is not None, 'VAE is not loaded'
return self.tod.decode(z)
@property
def device(self):
return next(self.parameters()).device
@property
def dtype(self):
return next(self.parameters()).dtype
def wrapped_decode(self, z):
with torch.amp.autocast('cuda', dtype=self.dtype):
mel_decoded = self.decode(z)
audio = self.vocode(mel_decoded)
return audio
def wrapped_encode(self, audio):
with torch.amp.autocast('cuda', dtype=self.dtype):
dist = self.encode_audio(audio)
return dist.mean
if not "mmaudio" in folder_paths.folder_names_and_paths:
folder_paths.add_model_folder_path("mmaudio", os.path.join(folder_paths.models_dir, "mmaudio"))
class OviMMAudioVAELoader:
"""Loads MMAudio VAE for audio encoding/decoding in Ovi"""
@classmethod
def INPUT_TYPES(s):
s.vae_files = folder_paths.get_filename_list("vae")
s.mmaudio_files = folder_paths.get_filename_list("mmaudio")
s.all_files = s.vae_files + s.mmaudio_files
return {
"required": {
"vae": (s.all_files, {"tooltip": "MMAudio VAE 16k (v1-16.pth) model from models/vae or models/mmaudio"}),
"vocoder": (s.all_files, {"tooltip": "BigVGAN vocoder (best_netG.pt) from models/vae or models/mmaudio"}),
"precision": (["bf16", "fp16", "fp32"], {"default": "bf16"}),
}
}
RETURN_TYPES = ("MMAUDIOVAE",)
RETURN_NAMES = ("mmaudio_vae",)
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper/Ovi"
DESCRIPTION = "Loads MMAudio VAE for Ovi audio generation"
def loadmodel(self, vae, vocoder, precision):
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
vae_path = folder_paths.get_full_path("vae", vae) if vae in self.vae_files else folder_paths.get_full_path("mmaudio", vae)
vocoder_path = folder_paths.get_full_path("vae", vocoder) if vocoder in self.vae_files else folder_paths.get_full_path("mmaudio", vocoder)
vae = FeaturesUtils(
tod_vae_ckpt=vae_path,
bigvgan_vocoder_ckpt=vocoder_path,
mode='16k',
need_vae_encoder=True
)
vae.to(device=offload_device, dtype=dtype)
vae.eval()
return (vae,)
class WanVideoDecodeOviAudio:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"mmaudio_vae": ("MMAUDIOVAE",),
"samples": ("LATENT",),
}
}
RETURN_TYPES = ("AUDIO",)
RETURN_NAMES = ("audio",)
FUNCTION = "decode"
CATEGORY = "WanVideoWrapper/Ovi"
def decode(self, mmaudio_vae, samples):
mm.soft_empty_cache()
audio_latents = samples.get("latent_ovi_audio", None)
if audio_latents is None:
raise ValueError("No Ovi audio latents found in input samples")
mmaudio_vae.to(device)
waveform = mmaudio_vae.wrapped_decode(audio_latents.to(device=device, dtype=mmaudio_vae.dtype))
audio = {"waveform": waveform.cpu().float(), "sample_rate": 16000}
mmaudio_vae.to(offload_device)
mm.soft_empty_cache()
return (audio,)
class WanVideoEncodeOviAudio:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"mmaudio_vae": ("MMAUDIOVAE",),
"audio": ("AUDIO",),
}
}
RETURN_TYPES = ("LATENT",)
RETURN_NAMES = ("samples",)
FUNCTION = "decode"
CATEGORY = "WanVideoWrapper/Ovi"
def decode(self, mmaudio_vae, audio):
mmaudio_vae.to(device)
waveform = audio.get("waveform", None)
sample_rate = audio.get("sample_rate", None)
if sample_rate != 16000:
waveform = torchaudio.functional.resample(waveform, sample_rate, 16000)
waveform = waveform.to(device=device, dtype=mmaudio_vae.dtype)[0][0].unsqueeze(0)
samples = mmaudio_vae.wrapped_encode(waveform)
mmaudio_vae.to(offload_device)
mm.soft_empty_cache()
return ({"latent_ovi_audio": samples},)
class WanVideoAddOviAudioToLatents:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"original_samples": ("LATENT",),
"audio_samples": ("LATENT",),
}
}
RETURN_TYPES = ("LATENT",)
RETURN_NAMES = ("samples",)
FUNCTION = "decode"
CATEGORY = "WanVideoWrapper/Ovi"
def decode(self, original_samples, audio_samples):
samples = original_samples.copy()
samples.update(audio_samples)
return (samples,)
class WanVideoEmptyMMAudioLatents:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"length": ("INT", {"default": 157, "min": 1, "max": 10000, "step": 1, "tooltip": "Length of the audio latent sequence"}),
}
}
RETURN_TYPES = ("LATENT",)
RETURN_NAMES = ("samples",)
FUNCTION = "decode"
CATEGORY = "WanVideoWrapper/Ovi"
def decode(self, length):
audio_latents = torch.zeros((length, 20), device=torch.device("cpu"), dtype=torch.float32) # 1, l c -> l, c
return ({"latent_ovi_audio": audio_latents},)
class WanVideoOviCFG:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"original_text_embeds": ("WANVIDEOTEXTEMBEDS",),
"ovi_audio_cfg": ("FLOAT", {"default": 3.0, "min": 0.0, "max": 100.0, "step": 0.01}),
},
"optional": {
"ovi_negative_text_embeds": ("WANVIDEOTEXTEMBEDS",),
}
}
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", )
RETURN_NAMES = ("text_embeds",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper/Ovi"
DESCRIPTION = "Adds Ovi negative text embeddings and audio CFG scale to the text embeddings dictionary"
def process(self, original_text_embeds, ovi_audio_cfg, ovi_negative_text_embeds=None):
negative_text_embeds = None
if ovi_negative_text_embeds is not None:
negative_text_embeds = ovi_negative_text_embeds.get("prompt_embeds", None)
if negative_text_embeds is None:
negative_text_embeds = original_text_embeds["prompt_embeds"]
log.info("WanVideoOviCFG: Ovi negative text embeddings not provided, using original prompt embeddings as negative embeddings")
else:
log.info("WanVideoOviCFG: Using provided Ovi audio negative text embeddings")
log.info("WanVideoOviCFG: negative text embedding shape: {}".format(negative_text_embeds[0].shape))
prompt_embeds_dict_copy = original_text_embeds.copy()
prompt_embeds_dict_copy.update({
"ovi_negative_prompt_embeds": negative_text_embeds,
"ovi_audio_cfg": ovi_audio_cfg,
})
return (prompt_embeds_dict_copy,)
NODE_CLASS_MAPPINGS = {
"OviMMAudioVAELoader": OviMMAudioVAELoader,
"WanVideoDecodeOviAudio": WanVideoDecodeOviAudio,
"WanVideoEncodeOviAudio": WanVideoEncodeOviAudio,
"WanVideoOviCFG": WanVideoOviCFG,
"WanVideoAddOviAudioToLatents": WanVideoAddOviAudioToLatents,
"WanVideoEmptyMMAudioLatents": WanVideoEmptyMMAudioLatents,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"OviMMAudioVAELoader": "Ovi MMAudio VAE Loader",
"WanVideoDecodeOviAudio": "WanVideo Decode Ovi Audio",
"WanVideoEncodeOviAudio": "WanVideo Encode Ovi Audio",
"WanVideoOviCFG": "WanVideo Ovi CFG",
"WanVideoAddOviAudioToLatents": "WanVideo Add MMAudio To Latents",
"WanVideoEmptyMMAudioLatents": "WanVideo Empty MMAudio Latents",
}
+54
View File
@@ -0,0 +1,54 @@
from typing import Literal, Optional
import torch
import torch.nn as nn
from .vae import VAE, get_my_vae
from .distributions import DiagonalGaussianDistribution
from ..bigvgan import BigVGAN
from comfy.utils import load_torch_file
class AutoEncoderModule(nn.Module):
def __init__(self,
*,
vae_ckpt_path,
vocoder_ckpt_path: Optional[str] = None,
mode: Literal['16k', '44k'],
need_vae_encoder: bool = True):
super().__init__()
self.vae: VAE = get_my_vae(mode).eval()
#vae_state_dict = torch.load(vae_ckpt_path, weights_only=True, map_location='cpu')'
vae_state_dict = load_torch_file(vae_ckpt_path)
self.vae.load_state_dict(vae_state_dict)
self.vae.remove_weight_norm()
if mode == '16k':
assert vocoder_ckpt_path is not None
self.vocoder = BigVGAN(vocoder_ckpt_path).eval()
elif mode == '44k':
raise NotImplementedError("44k mode requires BigVGANv2 which is not currently supported in this environment.")
self.vocoder = BigVGANv2.from_pretrained('nvidia/bigvgan_v2_44khz_128band_512x',
use_cuda_kernel=False)
self.vocoder.remove_weight_norm()
else:
raise ValueError(f'Unknown mode: {mode}')
for param in self.parameters():
param.requires_grad = False
if not need_vae_encoder:
del self.vae.encoder
@torch.inference_mode()
def encode(self, x: torch.Tensor) -> DiagonalGaussianDistribution:
return self.vae.encode(x)
@torch.inference_mode()
def decode(self, z: torch.Tensor) -> torch.Tensor:
return self.vae.decode(z)
@torch.inference_mode()
def vocode(self, spec: torch.Tensor) -> torch.Tensor:
return self.vocoder(spec)
+45
View File
@@ -0,0 +1,45 @@
from typing import Optional
import torch
import numpy as np
class DiagonalGaussianDistribution:
def __init__(self, parameters, deterministic=False):
self.parameters = parameters
self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
self.deterministic = deterministic
self.std = torch.exp(0.5 * self.logvar)
self.var = torch.exp(self.logvar)
if self.deterministic:
self.var = self.std = torch.zeros_like(self.mean).to(device=self.parameters.device)
def sample(self, rng: Optional[torch.Generator] = None):
# x = self.mean + self.std * torch.randn(self.mean.shape).to(device=self.parameters.device)
r = torch.empty_like(self.mean).normal_(generator=rng)
x = self.mean + self.std * r
return x
def kl(self, other=None):
if self.deterministic:
return torch.Tensor([0.])
else:
if other is None:
return 0.5 * torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar
else:
return 0.5 * (torch.pow(self.mean - other.mean, 2) / other.var +
self.var / other.var - 1.0 - self.logvar + other.logvar)
def nll(self, sample, dims=[1, 2, 3]):
if self.deterministic:
return torch.Tensor([0.])
logtwopi = np.log(2.0 * np.pi)
return 0.5 * torch.sum(logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
dim=dims)
def mode(self):
return self.mean
+168
View File
@@ -0,0 +1,168 @@
# Copyright (c) 2024, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# This work is licensed under a Creative Commons
# Attribution-NonCommercial-ShareAlike 4.0 International License.
# You should have received a copy of the license along with this
# work. If not, see http://creativecommons.org/licenses/by-nc-sa/4.0/
"""Improved diffusion model architecture proposed in the paper
"Analyzing and Improving the Training Dynamics of Diffusion Models"."""
import numpy as np
import torch
#----------------------------------------------------------------------------
# Variant of constant() that inherits dtype and device from the given
# reference tensor by default.
_constant_cache = dict()
def constant(value, shape=None, dtype=None, device=None, memory_format=None):
value = np.asarray(value)
if shape is not None:
shape = tuple(shape)
if dtype is None:
dtype = torch.get_default_dtype()
if device is None:
device = torch.device('cpu')
if memory_format is None:
memory_format = torch.contiguous_format
key = (value.shape, value.dtype, value.tobytes(), shape, dtype, device, memory_format)
tensor = _constant_cache.get(key, None)
if tensor is None:
tensor = torch.as_tensor(value.copy(), dtype=dtype, device=device)
if shape is not None:
tensor, _ = torch.broadcast_tensors(tensor, torch.empty(shape))
tensor = tensor.contiguous(memory_format=memory_format)
_constant_cache[key] = tensor
return tensor
def const_like(ref, value, shape=None, dtype=None, device=None, memory_format=None):
if dtype is None:
dtype = ref.dtype
if device is None:
device = ref.device
return constant(value, shape=shape, dtype=dtype, device=device, memory_format=memory_format)
#----------------------------------------------------------------------------
# Normalize given tensor to unit magnitude with respect to the given
# dimensions. Default = all dimensions except the first.
def normalize(x, dim=None, eps=1e-4):
if dim is None:
dim = list(range(1, x.ndim))
norm = torch.linalg.vector_norm(x, dim=dim, keepdim=True, dtype=torch.float32)
norm = torch.add(eps, norm, alpha=np.sqrt(norm.numel() / x.numel()))
return x / norm.to(x.dtype)
class Normalize(torch.nn.Module):
def __init__(self, dim=None, eps=1e-4):
super().__init__()
self.dim = dim
self.eps = eps
def forward(self, x):
return normalize(x, dim=self.dim, eps=self.eps)
#----------------------------------------------------------------------------
# Upsample or downsample the given tensor with the given filter,
# or keep it as is.
def resample(x, f=[1, 1], mode='keep'):
if mode == 'keep':
return x
f = np.float32(f)
assert f.ndim == 1 and len(f) % 2 == 0
pad = (len(f) - 1) // 2
f = f / f.sum()
f = np.outer(f, f)[np.newaxis, np.newaxis, :, :]
f = const_like(x, f)
c = x.shape[1]
if mode == 'down':
return torch.nn.functional.conv2d(x,
f.tile([c, 1, 1, 1]),
groups=c,
stride=2,
padding=(pad, ))
assert mode == 'up'
return torch.nn.functional.conv_transpose2d(x, (f * 4).tile([c, 1, 1, 1]),
groups=c,
stride=2,
padding=(pad, ))
#----------------------------------------------------------------------------
# Magnitude-preserving SiLU (Equation 81).
def mp_silu(x):
return torch.nn.functional.silu(x) / 0.596
class MPSiLU(torch.nn.Module):
def forward(self, x):
return mp_silu(x)
#----------------------------------------------------------------------------
# Magnitude-preserving sum (Equation 88).
def mp_sum(a, b, t=0.5):
return a.lerp(b, t) / np.sqrt((1 - t)**2 + t**2)
#----------------------------------------------------------------------------
# Magnitude-preserving concatenation (Equation 103).
def mp_cat(a, b, dim=1, t=0.5):
Na = a.shape[dim]
Nb = b.shape[dim]
C = np.sqrt((Na + Nb) / ((1 - t)**2 + t**2))
wa = C / np.sqrt(Na) * (1 - t)
wb = C / np.sqrt(Nb) * t
return torch.cat([wa * a, wb * b], dim=dim)
#----------------------------------------------------------------------------
# Magnitude-preserving convolution or fully-connected layer (Equation 47)
# with force weight normalization (Equation 66).
class MPConv1D(torch.nn.Module):
def __init__(self, in_channels, out_channels, kernel_size):
super().__init__()
self.out_channels = out_channels
self.weight = torch.nn.Parameter(torch.randn(out_channels, in_channels, kernel_size))
self.weight_norm_removed = False
def forward(self, x, gain=1):
assert self.weight_norm_removed, 'call remove_weight_norm() before inference'
w = self.weight * gain
if w.ndim == 2:
return x @ w.t()
assert w.ndim == 3
return torch.nn.functional.conv1d(x, w, padding=(w.shape[-1] // 2, ))
def remove_weight_norm(self):
w = self.weight.to(torch.float32)
w = normalize(w) # traditional weight normalization
w = w / np.sqrt(w[0].numel())
w = w.to(self.weight.dtype)
self.weight.data.copy_(w)
self.weight_norm_removed = True
return self
+376
View File
@@ -0,0 +1,376 @@
import logging
from typing import Optional
import torch
import torch.nn as nn
from .edm2_utils import MPConv1D
from .vae_modules import (AttnBlock1D, Downsample1D, ResnetBlock1D,
Upsample1D, nonlinearity)
from .distributions import DiagonalGaussianDistribution
log = logging.getLogger()
DATA_MEAN_80D = [
-1.6058, -1.3676, -1.2520, -1.2453, -1.2078, -1.2224, -1.2419, -1.2439, -1.2922, -1.2927,
-1.3170, -1.3543, -1.3401, -1.3836, -1.3907, -1.3912, -1.4313, -1.4152, -1.4527, -1.4728,
-1.4568, -1.5101, -1.5051, -1.5172, -1.5623, -1.5373, -1.5746, -1.5687, -1.6032, -1.6131,
-1.6081, -1.6331, -1.6489, -1.6489, -1.6700, -1.6738, -1.6953, -1.6969, -1.7048, -1.7280,
-1.7361, -1.7495, -1.7658, -1.7814, -1.7889, -1.8064, -1.8221, -1.8377, -1.8417, -1.8643,
-1.8857, -1.8929, -1.9173, -1.9379, -1.9531, -1.9673, -1.9824, -2.0042, -2.0215, -2.0436,
-2.0766, -2.1064, -2.1418, -2.1855, -2.2319, -2.2767, -2.3161, -2.3572, -2.3954, -2.4282,
-2.4659, -2.5072, -2.5552, -2.6074, -2.6584, -2.7107, -2.7634, -2.8266, -2.8981, -2.9673
]
DATA_STD_80D = [
1.0291, 1.0411, 1.0043, 0.9820, 0.9677, 0.9543, 0.9450, 0.9392, 0.9343, 0.9297, 0.9276, 0.9263,
0.9242, 0.9254, 0.9232, 0.9281, 0.9263, 0.9315, 0.9274, 0.9247, 0.9277, 0.9199, 0.9188, 0.9194,
0.9160, 0.9161, 0.9146, 0.9161, 0.9100, 0.9095, 0.9145, 0.9076, 0.9066, 0.9095, 0.9032, 0.9043,
0.9038, 0.9011, 0.9019, 0.9010, 0.8984, 0.8983, 0.8986, 0.8961, 0.8962, 0.8978, 0.8962, 0.8973,
0.8993, 0.8976, 0.8995, 0.9016, 0.8982, 0.8972, 0.8974, 0.8949, 0.8940, 0.8947, 0.8936, 0.8939,
0.8951, 0.8956, 0.9017, 0.9167, 0.9436, 0.9690, 1.0003, 1.0225, 1.0381, 1.0491, 1.0545, 1.0604,
1.0761, 1.0929, 1.1089, 1.1196, 1.1176, 1.1156, 1.1117, 1.1070
]
DATA_MEAN_128D = [
-3.3462, -2.6723, -2.4893, -2.3143, -2.2664, -2.3317, -2.1802, -2.4006, -2.2357, -2.4597,
-2.3717, -2.4690, -2.5142, -2.4919, -2.6610, -2.5047, -2.7483, -2.5926, -2.7462, -2.7033,
-2.7386, -2.8112, -2.7502, -2.9594, -2.7473, -3.0035, -2.8891, -2.9922, -2.9856, -3.0157,
-3.1191, -2.9893, -3.1718, -3.0745, -3.1879, -3.2310, -3.1424, -3.2296, -3.2791, -3.2782,
-3.2756, -3.3134, -3.3509, -3.3750, -3.3951, -3.3698, -3.4505, -3.4509, -3.5089, -3.4647,
-3.5536, -3.5788, -3.5867, -3.6036, -3.6400, -3.6747, -3.7072, -3.7279, -3.7283, -3.7795,
-3.8259, -3.8447, -3.8663, -3.9182, -3.9605, -3.9861, -4.0105, -4.0373, -4.0762, -4.1121,
-4.1488, -4.1874, -4.2461, -4.3170, -4.3639, -4.4452, -4.5282, -4.6297, -4.7019, -4.7960,
-4.8700, -4.9507, -5.0303, -5.0866, -5.1634, -5.2342, -5.3242, -5.4053, -5.4927, -5.5712,
-5.6464, -5.7052, -5.7619, -5.8410, -5.9188, -6.0103, -6.0955, -6.1673, -6.2362, -6.3120,
-6.3926, -6.4797, -6.5565, -6.6511, -6.8130, -6.9961, -7.1275, -7.2457, -7.3576, -7.4663,
-7.6136, -7.7469, -7.8815, -8.0132, -8.1515, -8.3071, -8.4722, -8.7418, -9.3975, -9.6628,
-9.7671, -9.8863, -9.9992, -10.0860, -10.1709, -10.5418, -11.2795, -11.3861
]
DATA_STD_128D = [
2.3804, 2.4368, 2.3772, 2.3145, 2.2803, 2.2510, 2.2316, 2.2083, 2.1996, 2.1835, 2.1769, 2.1659,
2.1631, 2.1618, 2.1540, 2.1606, 2.1571, 2.1567, 2.1612, 2.1579, 2.1679, 2.1683, 2.1634, 2.1557,
2.1668, 2.1518, 2.1415, 2.1449, 2.1406, 2.1350, 2.1313, 2.1415, 2.1281, 2.1352, 2.1219, 2.1182,
2.1327, 2.1195, 2.1137, 2.1080, 2.1179, 2.1036, 2.1087, 2.1036, 2.1015, 2.1068, 2.0975, 2.0991,
2.0902, 2.1015, 2.0857, 2.0920, 2.0893, 2.0897, 2.0910, 2.0881, 2.0925, 2.0873, 2.0960, 2.0900,
2.0957, 2.0958, 2.0978, 2.0936, 2.0886, 2.0905, 2.0845, 2.0855, 2.0796, 2.0840, 2.0813, 2.0817,
2.0838, 2.0840, 2.0917, 2.1061, 2.1431, 2.1976, 2.2482, 2.3055, 2.3700, 2.4088, 2.4372, 2.4609,
2.4731, 2.4847, 2.5072, 2.5451, 2.5772, 2.6147, 2.6529, 2.6596, 2.6645, 2.6726, 2.6803, 2.6812,
2.6899, 2.6916, 2.6931, 2.6998, 2.7062, 2.7262, 2.7222, 2.7158, 2.7041, 2.7485, 2.7491, 2.7451,
2.7485, 2.7233, 2.7297, 2.7233, 2.7145, 2.6958, 2.6788, 2.6439, 2.6007, 2.4786, 2.2469, 2.1877,
2.1392, 2.0717, 2.0107, 1.9676, 1.9140, 1.7102, 0.9101, 0.7164
]
class VAE(nn.Module):
def __init__(
self,
*,
data_dim: int,
embed_dim: int,
hidden_dim: int,
):
super().__init__()
if data_dim == 80:
data_mean = torch.tensor(DATA_MEAN_80D, dtype=torch.float32)
data_std = torch.tensor(DATA_STD_80D, dtype=torch.float32)
elif data_dim == 128:
data_mean = torch.tensor(DATA_MEAN_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")
# match old shape: (1, channels, 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(
dim=hidden_dim,
ch_mult=(1, 2, 4),
num_res_blocks=2,
attn_layers=[3],
down_layers=[0],
in_dim=data_dim,
embed_dim=embed_dim,
)
self.decoder = Decoder1D(
dim=hidden_dim,
ch_mult=(1, 2, 4),
num_res_blocks=2,
attn_layers=[3],
down_layers=[0],
in_dim=data_dim,
out_dim=data_dim,
embed_dim=embed_dim,
)
self.embed_dim = embed_dim
# self.quant_conv = nn.Conv1d(2 * embed_dim, 2 * embed_dim, 1)
# self.post_quant_conv = nn.Conv1d(embed_dim, embed_dim, 1)
self.initialize_weights()
def initialize_weights(self):
pass
def encode(self, x: torch.Tensor, normalize: bool = True) -> DiagonalGaussianDistribution:
if normalize:
x = self.normalize(x)
moments = self.encoder(x)
posterior = DiagonalGaussianDistribution(moments)
return posterior
def decode(self, z: torch.Tensor, unnormalize: bool = True) -> torch.Tensor:
dec = self.decoder(z)
if unnormalize:
dec = self.unnormalize(dec)
return dec
def normalize(self, x: torch.Tensor) -> torch.Tensor:
return (x - self.data_mean) / self.data_std
def unnormalize(self, x: torch.Tensor) -> torch.Tensor:
return x * self.data_std + self.data_mean
def forward(
self,
x: torch.Tensor,
sample_posterior: bool = True,
rng: Optional[torch.Generator] = None,
normalize: bool = True,
unnormalize: bool = True,
) -> tuple[torch.Tensor, DiagonalGaussianDistribution]:
posterior = self.encode(x, normalize=normalize)
if sample_posterior:
z = posterior.sample(rng)
else:
z = posterior.mode()
dec = self.decode(z, unnormalize=unnormalize)
return dec, posterior
def load_weights(self, src_dict) -> None:
self.load_state_dict(src_dict, strict=True)
@property
def device(self) -> torch.device:
return next(self.parameters()).device
def get_last_layer(self):
return self.decoder.conv_out.weight
def remove_weight_norm(self):
for name, m in self.named_modules():
if isinstance(m, MPConv1D):
m.remove_weight_norm()
log.debug(f"Removed weight norm from {name}")
return self
class Encoder1D(nn.Module):
def __init__(self,
*,
dim: int,
ch_mult: tuple[int] = (1, 2, 4, 8),
num_res_blocks: int,
attn_layers: list[int] = [],
down_layers: list[int] = [],
resamp_with_conv: bool = True,
in_dim: int,
embed_dim: int,
double_z: bool = True,
kernel_size: int = 3,
clip_act: float = 256.0):
super().__init__()
self.dim = dim
self.num_layers = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.in_channels = in_dim
self.clip_act = clip_act
self.down_layers = down_layers
self.attn_layers = attn_layers
self.conv_in = MPConv1D(in_dim, self.dim, kernel_size=kernel_size)
in_ch_mult = (1, ) + tuple(ch_mult)
self.in_ch_mult = in_ch_mult
# downsampling
self.down = nn.ModuleList()
for i_level in range(self.num_layers):
block = nn.ModuleList()
attn = nn.ModuleList()
block_in = dim * in_ch_mult[i_level]
block_out = dim * ch_mult[i_level]
for i_block in range(self.num_res_blocks):
block.append(
ResnetBlock1D(in_dim=block_in,
out_dim=block_out,
kernel_size=kernel_size,
use_norm=True))
block_in = block_out
if i_level in attn_layers:
attn.append(AttnBlock1D(block_in))
down = nn.Module()
down.block = block
down.attn = attn
if i_level in down_layers:
down.downsample = Downsample1D(block_in, resamp_with_conv)
self.down.append(down)
# middle
self.mid = nn.Module()
self.mid.block_1 = ResnetBlock1D(in_dim=block_in,
out_dim=block_in,
kernel_size=kernel_size,
use_norm=True)
self.mid.attn_1 = AttnBlock1D(block_in)
self.mid.block_2 = ResnetBlock1D(in_dim=block_in,
out_dim=block_in,
kernel_size=kernel_size,
use_norm=True)
# end
self.conv_out = MPConv1D(block_in,
2 * embed_dim if double_z else embed_dim,
kernel_size=kernel_size)
self.learnable_gain = nn.Parameter(torch.zeros([]))
def forward(self, x):
# downsampling
hs = [self.conv_in(x)]
for i_level in range(self.num_layers):
for i_block in range(self.num_res_blocks):
h = self.down[i_level].block[i_block](hs[-1])
if len(self.down[i_level].attn) > 0:
h = self.down[i_level].attn[i_block](h)
h = h.clamp(-self.clip_act, self.clip_act)
hs.append(h)
if i_level in self.down_layers:
hs.append(self.down[i_level].downsample(hs[-1]))
# middle
h = hs[-1]
h = self.mid.block_1(h)
h = self.mid.attn_1(h)
h = self.mid.block_2(h)
h = h.clamp(-self.clip_act, self.clip_act)
# end
h = nonlinearity(h)
h = self.conv_out(h, gain=(self.learnable_gain + 1))
return h
class Decoder1D(nn.Module):
def __init__(self,
*,
dim: int,
out_dim: int,
ch_mult: tuple[int] = (1, 2, 4, 8),
num_res_blocks: int,
attn_layers: list[int] = [],
down_layers: list[int] = [],
kernel_size: int = 3,
resamp_with_conv: bool = True,
in_dim: int,
embed_dim: int,
clip_act: float = 256.0):
super().__init__()
self.ch = dim
self.num_layers = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.in_channels = in_dim
self.clip_act = clip_act
self.down_layers = [i + 1 for i in down_layers] # each downlayer add one
# compute in_ch_mult, block_in and curr_res at lowest res
block_in = dim * ch_mult[self.num_layers - 1]
# z to block_in
self.conv_in = MPConv1D(embed_dim, block_in, kernel_size=kernel_size)
# middle
self.mid = nn.Module()
self.mid.block_1 = ResnetBlock1D(in_dim=block_in, out_dim=block_in, use_norm=True)
self.mid.attn_1 = AttnBlock1D(block_in)
self.mid.block_2 = ResnetBlock1D(in_dim=block_in, out_dim=block_in, use_norm=True)
# upsampling
self.up = nn.ModuleList()
for i_level in reversed(range(self.num_layers)):
block = nn.ModuleList()
attn = nn.ModuleList()
block_out = dim * ch_mult[i_level]
for i_block in range(self.num_res_blocks + 1):
block.append(ResnetBlock1D(in_dim=block_in, out_dim=block_out, use_norm=True))
block_in = block_out
if i_level in attn_layers:
attn.append(AttnBlock1D(block_in))
up = nn.Module()
up.block = block
up.attn = attn
if i_level in self.down_layers:
up.upsample = Upsample1D(block_in, resamp_with_conv)
self.up.insert(0, up) # prepend to get consistent order
# end
self.conv_out = MPConv1D(block_in, out_dim, kernel_size=kernel_size)
self.learnable_gain = nn.Parameter(torch.zeros([]))
def forward(self, z):
# z to block_in
h = self.conv_in(z)
# middle
h = self.mid.block_1(h)
h = self.mid.attn_1(h)
h = self.mid.block_2(h)
h = h.clamp(-self.clip_act, self.clip_act)
# upsampling
for i_level in reversed(range(self.num_layers)):
for i_block in range(self.num_res_blocks + 1):
h = self.up[i_level].block[i_block](h)
if len(self.up[i_level].attn) > 0:
h = self.up[i_level].attn[i_block](h)
h = h.clamp(-self.clip_act, self.clip_act)
if i_level in self.down_layers:
h = self.up[i_level].upsample(h)
h = nonlinearity(h)
h = self.conv_out(h, gain=(self.learnable_gain + 1))
return h
def VAE_16k(**kwargs) -> VAE:
return VAE(data_dim=80, embed_dim=20, hidden_dim=384, **kwargs)
def VAE_44k(**kwargs) -> VAE:
return VAE(data_dim=128, embed_dim=40, hidden_dim=512, **kwargs)
def get_my_vae(name: str, **kwargs) -> VAE:
if name == '16k':
return VAE_16k(**kwargs)
if name == '44k':
return VAE_44k(**kwargs)
raise ValueError(f'Unknown model: {name}')
if __name__ == '__main__':
network = get_my_vae('standard')
# print the number of parameters in terms of millions
num_params = sum(p.numel() for p in network.parameters()) / 1e6
print(f'Number of parameters: {num_params:.2f}M')
+117
View File
@@ -0,0 +1,117 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from .edm2_utils import (MPConv1D, mp_silu, mp_sum, normalize)
def nonlinearity(x):
# swish
return mp_silu(x)
class ResnetBlock1D(nn.Module):
def __init__(self, *, in_dim, out_dim=None, conv_shortcut=False, kernel_size=3, use_norm=True):
super().__init__()
self.in_dim = in_dim
out_dim = in_dim if out_dim is None else out_dim
self.out_dim = out_dim
self.use_conv_shortcut = conv_shortcut
self.use_norm = use_norm
self.conv1 = MPConv1D(in_dim, out_dim, kernel_size=kernel_size)
self.conv2 = MPConv1D(out_dim, out_dim, kernel_size=kernel_size)
if self.in_dim != self.out_dim:
if self.use_conv_shortcut:
self.conv_shortcut = MPConv1D(in_dim, out_dim, kernel_size=kernel_size)
else:
self.nin_shortcut = MPConv1D(in_dim, out_dim, kernel_size=1)
def forward(self, x: torch.Tensor) -> torch.Tensor:
# pixel norm
if self.use_norm:
x = normalize(x, dim=1)
h = x
h = nonlinearity(h)
h = self.conv1(h)
h = nonlinearity(h)
h = self.conv2(h)
if self.in_dim != self.out_dim:
if self.use_conv_shortcut:
x = self.conv_shortcut(x)
else:
x = self.nin_shortcut(x)
return mp_sum(x, h, t=0.3)
class AttnBlock1D(nn.Module):
def __init__(self, in_channels, num_heads=1):
super().__init__()
self.in_channels = in_channels
self.num_heads = num_heads
self.qkv = MPConv1D(in_channels, in_channels * 3, kernel_size=1)
self.proj_out = MPConv1D(in_channels, in_channels, kernel_size=1)
def forward(self, x):
h = x
y = self.qkv(h)
y = y.reshape(y.shape[0], self.num_heads, -1, 3, y.shape[-1])
q, k, v = normalize(y, dim=2).unbind(3)
q = rearrange(q, 'b h c l -> b h l c')
k = rearrange(k, 'b h c l -> b h l c')
v = rearrange(v, 'b h c l -> b h l c')
h = F.scaled_dot_product_attention(q, k, v)
h = rearrange(h, 'b h l c -> b (h c) l')
h = self.proj_out(h)
return mp_sum(x, h, t=0.3)
class Upsample1D(nn.Module):
def __init__(self, in_channels, with_conv):
super().__init__()
self.with_conv = with_conv
if self.with_conv:
self.conv = MPConv1D(in_channels, in_channels, kernel_size=3)
def forward(self, x):
x = F.interpolate(x, scale_factor=2.0, mode='nearest-exact') # support 3D tensor(B,C,T)
if self.with_conv:
x = self.conv(x)
return x
class Downsample1D(nn.Module):
def __init__(self, in_channels, with_conv):
super().__init__()
self.with_conv = with_conv
if self.with_conv:
# no asymmetric padding in torch conv, must do it ourselves
self.conv1 = MPConv1D(in_channels, in_channels, kernel_size=1)
self.conv2 = MPConv1D(in_channels, in_channels, kernel_size=1)
def forward(self, x):
if self.with_conv:
x = self.conv1(x)
x = F.avg_pool1d(x, kernel_size=2, stride=2)
if self.with_conv:
x = self.conv2(x)
return x
+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
+69 -29
View File
@@ -1,35 +1,75 @@
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
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 .unianimate.nodes import NODE_CLASS_MAPPINGS as UNIANIMATE_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNIANIMATE_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 .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
try:
from .utils import check_duplicate_nodes, log, color_text
duplicate_dirs = check_duplicate_nodes()
if duplicate_dirs:
warning_msg = f"WARNING: Found {len(duplicate_dirs)} other WanVideoWrapper directories:\n"
for dir_path in duplicate_dirs:
warning_msg += f" - {color_text(dir_path, 'yellow')}\n"
log.warning(color_text(warning_msg + "Please remove duplicates to avoid possible conflicts.", "red"))
except Exception:
pass
from .causvid.nodes import NODE_CLASS_MAPPINGS as CAUSVID_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as CAUSVID_NODE_DISPLAY_NAME_MAPPINGS
from .utils import log
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(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 = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
NODE_CLASS_MAPPINGS.update(CAUSVID_NODE_CLASS_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"),
]
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(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)
# Optional modules (will warn on import failure)
OPTIONAL_MODULES = [
(".nodes_deprecated", "Deprecated"),
(".s2v.nodes", "S2V"),
(".FlashVSR.flashvsr_nodes", "FlashVSR"),
(".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"),
]
NODE_DISPLAY_NAME_MAPPINGS.update(CAUSVID_NODE_DISPLAY_NAME_MAPPINGS)
def register_nodes(module_path: str, name: str, optional: bool) -> None:
"""Import and register nodes from a module."""
try:
import importlib
module = importlib.import_module(module_path, package=__package__)
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
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
# Register all node modules
for module_path, name in REQUIRED_MODULES:
register_nodes(module_path, name, optional=False)
for module_path, name in OPTIONAL_MODULES:
register_nodes(module_path, name, optional=True)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+159
View File
@@ -0,0 +1,159 @@
from ..utils import log
import torch
def set_transformer_cache_method(transformer, timesteps, cache_args=None):
transformer.cache_device = cache_args["cache_device"]
if cache_args["cache_type"] == "TeaCache":
log.info(f"TeaCache: Using cache device: {transformer.cache_device}")
transformer.teacache_state.clear_all()
transformer.enable_teacache = True
transformer.rel_l1_thresh = cache_args["rel_l1_thresh"]
transformer.teacache_start_step = cache_args["start_step"]
transformer.teacache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"]
transformer.teacache_use_coefficients = cache_args["use_coefficients"]
transformer.teacache_mode = cache_args["mode"]
elif cache_args["cache_type"] == "MagCache":
log.info(f"MagCache: Using cache device: {transformer.cache_device}")
transformer.magcache_state.clear_all()
transformer.enable_magcache = True
transformer.magcache_start_step = cache_args["start_step"]
transformer.magcache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"]
transformer.magcache_thresh = cache_args["magcache_thresh"]
transformer.magcache_K = cache_args["magcache_K"]
elif cache_args["cache_type"] == "EasyCache":
log.info(f"EasyCache: Using cache device: {transformer.cache_device}")
transformer.easycache_state.clear_all()
transformer.enable_easycache = True
transformer.easycache_start_step = cache_args["start_step"]
transformer.easycache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"]
transformer.easycache_thresh = cache_args["easycache_thresh"]
return transformer
class TeaCacheState:
def __init__(self, cache_device='cpu'):
self.cache_device = cache_device
self.states = {}
self._next_pred_id = 0
def new_prediction(self, cache_device='cpu'):
"""Create new prediction state and return its ID"""
self.cache_device = cache_device
pred_id = self._next_pred_id
self._next_pred_id += 1
self.states[pred_id] = {
'previous_residual': None,
'accumulated_rel_l1_distance': 0,
'previous_modulated_input': None,
'skipped_steps': [],
}
return pred_id
def update(self, pred_id, **kwargs):
"""Update state for specific prediction"""
if pred_id not in self.states:
return None
for key, value in kwargs.items():
self.states[pred_id][key] = value
def get(self, pred_id):
return self.states.get(pred_id, {})
def clear_all(self):
self.states = {}
self._next_pred_id = 0
class MagCacheState:
def __init__(self, cache_device='cpu'):
self.cache_device = cache_device
self.states = {}
self._next_pred_id = 0
def new_prediction(self, cache_device='cpu'):
"""Create new prediction state and return its ID"""
self.cache_device = cache_device
pred_id = self._next_pred_id
self._next_pred_id += 1
self.states[pred_id] = {
'residual_cache': None,
'accumulated_ratio': 1.0,
'accumulated_steps': 0,
'accumulated_err': 0,
'skipped_steps': [],
}
return pred_id
def update(self, pred_id, **kwargs):
"""Update state for specific prediction"""
if pred_id not in self.states:
return None
for key, value in kwargs.items():
self.states[pred_id][key] = value
def get(self, pred_id):
return self.states.get(pred_id, {})
def clear_all(self):
self.states = {}
self._next_pred_id = 0
class EasyCacheState:
def __init__(self, cache_device='cpu'):
self.cache_device = cache_device
self.states = {}
self._next_pred_id = 0
def new_prediction(self, cache_device='cpu'):
"""Create a new prediction state and return its ID."""
self.cache_device = cache_device
pred_id = self._next_pred_id
self._next_pred_id += 1
self.states[pred_id] = {
'previous_raw_input': None,
'previous_raw_output': None,
'cache': None,
'accumulated_error': 0.0,
'skipped_steps': [],
'cache_ovi': None,
}
return pred_id
def update(self, pred_id, **kwargs):
"""Update state for a specific prediction."""
if pred_id not in self.states:
return None
for key, value in kwargs.items():
self.states[pred_id][key] = value
def get(self, pred_id):
return self.states.get(pred_id, {})
def clear_all(self):
self.states = {}
self._next_pred_id = 0
def relative_l1_distance(last_tensor, current_tensor):
l1_distance = torch.abs(last_tensor.to(current_tensor.device) - current_tensor).mean()
norm = torch.abs(last_tensor).mean()
relative_l1_distance = l1_distance / norm
return relative_l1_distance.to(torch.float32).to(current_tensor.device)
def cache_report(transformer, cache_args):
cache_type = cache_args["cache_type"]
states = (
transformer.teacache_state.states if cache_type == "TeaCache" else
transformer.magcache_state.states if cache_type == "MagCache" else
transformer.easycache_state.states if cache_type == "EasyCache" else
None
)
state_names = {
0: "conditional",
1: "unconditional"
}
for pred_id, state in states.items():
name = state_names.get(pred_id, f"prediction_{pred_id}")
if 'skipped_steps' in state:
log.info(f"{cache_type} skipped: {len(state['skipped_steps'])} {name} steps: {state['skipped_steps']}")
transformer.teacache_state.clear_all()
transformer.magcache_state.clear_all()
transformer.easycache_state.clear_all()
del states
+128
View File
@@ -0,0 +1,128 @@
from comfy import model_management as mm
class WanVideoTeaCache:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"rel_l1_thresh": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.001,
"tooltip": "Higher values will make TeaCache more aggressive, faster, but may cause artifacts. Good value range for 1.3B: 0.05 - 0.08, for other models 0.15-0.30"}),
"start_step": ("INT", {"default": 1, "min": 0, "max": 9999, "step": 1, "tooltip": "Start percentage of the steps to apply TeaCache"}),
"end_step": ("INT", {"default": -1, "min": -1, "max": 9999, "step": 1, "tooltip": "End steps to apply TeaCache"}),
"cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}),
"use_coefficients": ("BOOLEAN", {"default": True, "tooltip": "Use calculated coefficients for more accuracy. When enabled therel_l1_thresh should be about 10 times higher than without"}),
},
"optional": {
"mode": (["e", "e0"], {"default": "e", "tooltip": "Choice between using e (time embeds, default) or e0 (modulated time embeds)"}),
},
}
RETURN_TYPES = ("CACHEARGS",)
RETURN_NAMES = ("cache_args",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = """
Patch WanVideo model to use TeaCache. Speeds up inference by caching the output and
applying it instead of doing the step. Best results are achieved by choosing the
appropriate coefficients for the model. Early steps should never be skipped, with too
aggressive values this can happen and the motion suffers. Starting later can help with that too.
When NOT using coefficients, the threshold value should be
about 10 times smaller than the value used with coefficients.
Official recommended values https://github.com/ali-vilab/TeaCache/tree/main/TeaCache4Wan2.1
"""
def process(self, rel_l1_thresh, start_step, end_step, cache_device, use_coefficients, mode="e"):
if cache_device == "main_device":
cache_device = mm.get_torch_device()
else:
cache_device = mm.unet_offload_device()
cache_args = {
"cache_type": "TeaCache",
"rel_l1_thresh": rel_l1_thresh,
"start_step": start_step,
"end_step": end_step,
"cache_device": cache_device,
"use_coefficients": use_coefficients,
"mode": mode,
}
return (cache_args,)
class WanVideoMagCache:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"magcache_thresh": ("FLOAT", {"default": 0.02, "min": 0.0, "max": 0.3, "step": 0.001, "tooltip": "How strongly to cache the output of diffusion model. This value must be non-negative."}),
"magcache_K": ("INT", {"default": 4, "min": 0, "max": 6, "step": 1, "tooltip": "The maxium skip steps of MagCache."}),
"start_step": ("INT", {"default": 1, "min": 0, "max": 9999, "step": 1, "tooltip": "Step to start applying MagCache"}),
"end_step": ("INT", {"default": -1, "min": -1, "max": 9999, "step": 1, "tooltip": "Step to end applying MagCache"}),
"cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}),
},
}
RETURN_TYPES = ("CACHEARGS",)
RETURN_NAMES = ("cache_args",)
FUNCTION = "setargs"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
DESCRIPTION = "MagCache for WanVideoWrapper, source https://github.com/Zehong-Ma/MagCache"
def setargs(self, magcache_thresh, magcache_K, start_step, end_step, cache_device):
if cache_device == "main_device":
cache_device = mm.get_torch_device()
else:
cache_device = mm.unet_offload_device()
cache_args = {
"cache_type": "MagCache",
"magcache_thresh": magcache_thresh,
"magcache_K": magcache_K,
"start_step": start_step,
"end_step": end_step,
"cache_device": cache_device,
}
return (cache_args,)
class WanVideoEasyCache:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"easycache_thresh": ("FLOAT", {"default": 0.015, "min": 0.0, "max": 1.0, "step": 0.001, "tooltip": "How strongly to cache the output of diffusion model. This value must be non-negative."}),
"start_step": ("INT", {"default": 10, "min": 0, "max": 9999, "step": 1, "tooltip": "Step to start applying EasyCache"}),
"end_step": ("INT", {"default": -1, "min": -1, "max": 9999, "step": 1, "tooltip": "Step to end applying EasyCache"}),
"cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}),
},
}
RETURN_TYPES = ("CACHEARGS",)
RETURN_NAMES = ("cache_args",)
FUNCTION = "setargs"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
DESCRIPTION = "EasyCache for WanVideoWrapper, source https://github.com/H-EmbodVis/EasyCache"
def setargs(self, easycache_thresh, start_step, end_step, cache_device):
if cache_device == "main_device":
cache_device = mm.get_torch_device()
else:
cache_device = mm.unet_offload_device()
cache_args = {
"cache_type": "EasyCache",
"easycache_thresh": easycache_thresh,
"start_step": start_step,
"end_step": end_step,
"cache_device": cache_device,
}
return (cache_args,)
NODE_CLASS_MAPPINGS = {
"WanVideoTeaCache": WanVideoTeaCache,
"WanVideoMagCache": WanVideoMagCache,
"WanVideoEasyCache": WanVideoEasyCache,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoTeaCache": "WanVideo TeaCache",
"WanVideoMagCache": "WanVideo MagCache",
"WanVideoEasyCache": "WanVideo EasyCache"
}
-612
View File
@@ -1,612 +0,0 @@
import os
import torch
import gc
from ..utils import log, print_memory, fourier_filter
import math
from tqdm import tqdm
from ..wanvideo.modules.model import rope_params
from ..wanvideo.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from ..wanvideo.utils.scheduling_flow_match_lcm import FlowMatchLCMScheduler
from ..wanvideo.utils.basic_flowmatch import FlowMatchScheduler
from ..nodes import optimized_scale
from einops import rearrange
from ..enhance_a_video.globals import disable_enhance
import comfy.model_management as mm
from comfy.utils import load_torch_file, ProgressBar, common_upscale
from comfy.clip_vision import clip_preprocess, ClipVisionModel
from comfy.cli_args import args, LatentPreviewMethod
script_directory = os.path.dirname(os.path.abspath(__file__))
#region Sampler
class WanVideoCausVidSampler:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("WANVIDEOMODEL",),
"text_embeds": ("WANVIDEOTEXTEMBEDS", ),
"image_embeds": ("WANVIDIMAGE_EMBEDS", ),
"steps": ("INT", {"default": 30, "min": 1}),
"shift": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 1000.0, "step": 0.01}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"force_offload": ("BOOLEAN", {"default": True, "tooltip": "Moves the model to the offload device after sampling"}),
"scheduler": ([
"flowmatch_causvid", "flowmatch_causvid_14b", "flowmatch_causvid_self_forcing",
#"unipc", "unipc/beta", "euler", "euler/beta", "lcm", "lcm/beta"
],
{
"default": 'flowmatch_causvid'
}),
"kv_cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}),
},
"optional": {
"samples": ("LATENT", {"tooltip": "init Latents to use for video2video process"} ),
"prefix_samples": ("LATENT", {"tooltip": "prefix latents"} ),
"denoise_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"rope_function": (["default", "comfy"], {"default": "default", "tooltip": "Comfy's RoPE implementation doesn't use complex numbers and can thus be compiled, that should be a lot faster when using torch.compile"}),
"experimental_args": ("EXPERIMENTALARGS", ),
}
}
RETURN_TYPES = ("LATENT", )
RETURN_NAMES = ("samples",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def _initialize_kv_cache(self, batch_size, dtype, device, num_blocks=30, num_heads=12):
"""
Initialize a Per-GPU KV cache for the Wan model.
"""
kv_cache1 = []
for _ in range(num_blocks):
kv_cache1.append({
"k": torch.zeros([batch_size, self.cache_window_size, num_heads, 128], dtype=dtype, device=device),
"v": torch.zeros([batch_size, self.cache_window_size, num_heads, 128], dtype=dtype, device=device),
"global_end_index": torch.tensor([0], dtype=torch.long, device=device),
"local_end_index": torch.tensor([0], dtype=torch.long, device=device)
})
self.kv_cache1 = kv_cache1 # always store the clean cache
def _initialize_crossattn_cache(self, batch_size, dtype, device, num_blocks=30, num_heads=12):
"""
Initialize a Per-GPU cross-attention cache for the Wan model.
"""
crossattn_cache = []
for _ in range(num_blocks):
crossattn_cache.append({
"k": torch.zeros([batch_size, 512, num_heads, 128], dtype=dtype, device=device),
"v": torch.zeros([batch_size, 512, num_heads, 128], dtype=dtype, device=device),
"is_init": False
})
self.crossattn_cache = crossattn_cache
def _shift_kv_cache(self):
"""
Shift the KV cache left by shift_blocks * num_frame_per_block * frame_seq_length.
This is called when kv_start exceeds window_size.
The first block is preserved, and shifting starts from the second block.
"""
shift_length = self.shift_blocks * self.num_frame_per_block * self.frame_seq_length
for block in self.kv_cache1:
block["k"] = torch.roll(block["k"], shifts=-shift_length, dims=1)
block["v"] = torch.roll(block["v"], shifts=-shift_length, dims=1)
# Clear the shifted-out part (except the first block)
block["k"][:, -shift_length:] = 0
block["v"][:, -shift_length:] = 0
# Update kv_start
self.kv_start -= shift_length
return shift_length
def process(self, model, text_embeds, image_embeds, shift, steps, seed, scheduler, kv_cache_device,
force_offload=True, samples=None, prefix_samples=None, denoise_strength=1.0, rope_function="default",
experimental_args=None):
#assert not (context_options and teacache_args), "Context options cannot currently be used together with teacache."
patcher = model
model = model.model
transformer = model.diffusion_model
dtype = model["dtype"]
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
if kv_cache_device == "main_device":
cache_device = mm.get_torch_device()
else:
cache_device = mm.unet_offload_device()
steps = int(steps/denoise_strength)
timesteps = None
if 'unipc' in scheduler:
sample_scheduler = FlowUniPCMultistepScheduler(shift=shift)
sample_scheduler.set_timesteps(steps, device=device, shift=shift, use_beta_sigmas=('beta' in scheduler))
elif 'euler' in scheduler:
sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta'))
sample_scheduler.set_timesteps(steps, device=device)
elif 'lcm' in scheduler:
sample_scheduler = FlowMatchLCMScheduler(shift=shift, use_beta_sigmas=(scheduler == 'lcm/beta'))
sample_scheduler.set_timesteps(steps, device=device)
elif 'flowmatch_causvid' in scheduler:
sample_scheduler = FlowMatchScheduler(
shift=shift, sigma_min=0.0, extra_one_step=True
)
sample_scheduler.set_timesteps(1000, training=True)
denoising_step_list = torch.tensor([1000, 757, 522], dtype=torch.long)
if "warp" in scheduler or "self_forcing" in scheduler:
denoising_step_list = torch.tensor([1000, 750, 500, 250] , dtype=torch.long)
timesteps = torch.cat((sample_scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32)))
denoising_step_list = timesteps[1000 - denoising_step_list]
elif "14b" in scheduler:
denoising_step_list = torch.tensor([1000, 934, 862, 756, 603, 410, 250, 140, 74], dtype=torch.long)
# sample_scheduler = FlowMatchScheduler(num_inference_steps=steps, shift=shift, sigma_min=0, extra_one_step=True)
# sample_scheduler.timesteps = torch.tensor(denoising_step_list).to(device)
# sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)])
#print(sample_scheduler.sigmas)
timesteps = denoising_step_list
#timesteps = torch.tensor(denoising_list).to(device)
print("timesteps", timesteps)
if denoise_strength < 1.0:
steps = int(steps * denoise_strength)
timesteps = timesteps[-(steps + 1):]
seed_g = torch.Generator(device=torch.device("cpu"))
seed_g.manual_seed(seed)
clip_fea, clip_fea_neg = None, None
vace_data, vace_context, vace_scale = None, None, None
image_cond = image_embeds.get("image_embeds", None)
target_shape = image_embeds.get("target_shape", None)
if target_shape is None:
raise ValueError("Empty image embeds must be provided for T2V (Text to Video")
has_ref = image_embeds.get("has_ref", False)
vace_context = image_embeds.get("vace_context", None)
vace_scale = image_embeds.get("vace_scale", None)
vace_start_percent = image_embeds.get("vace_start_percent", 0.0)
vace_end_percent = image_embeds.get("vace_end_percent", 1.0)
vace_seqlen = image_embeds.get("vace_seq_len", None)
vace_additional_embeds = image_embeds.get("additional_vace_inputs", [])
if vace_context is not None:
vace_data = [
{"context": vace_context,
"scale": vace_scale,
"start": vace_start_percent,
"end": vace_end_percent,
"seq_len": vace_seqlen
}
]
if len(vace_additional_embeds) > 0:
for i in range(len(vace_additional_embeds)):
if vace_additional_embeds[i].get("has_ref", False):
has_ref = True
vace_data.append({
"context": vace_additional_embeds[i]["vace_context"],
"scale": vace_additional_embeds[i]["vace_scale"],
"start": vace_additional_embeds[i]["vace_start_percent"],
"end": vace_additional_embeds[i]["vace_end_percent"],
"seq_len": vace_additional_embeds[i]["vace_seq_len"]
})
noise = torch.randn(
target_shape[0],
target_shape[1] + 1 if has_ref else target_shape[1],
target_shape[2],
target_shape[3],
dtype=torch.float32,
device=torch.device("cpu"),
generator=seed_g)
noise = noise.to(device, dtype)
latent_video_length = noise.shape[1]
if samples is not None:
input_samples = samples["samples"].squeeze(0).to(noise)
if input_samples.shape[1] != noise.shape[1]:
input_samples = torch.cat([input_samples[:, :1].repeat(1, noise.shape[1] - input_samples.shape[1], 1, 1), input_samples], dim=1)
original_image = input_samples.to(device)
if denoise_strength < 1.0:
latent_timestep = timesteps[:1].to(noise)
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
mask = samples.get("mask", None)
if mask is not None:
if mask.shape[2] != noise.shape[1]:
mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - mask.shape[2], noise.shape[2], noise.shape[3]), mask], dim=2)
init_latents = noise.to(device)
fps_embeds = None
if hasattr(transformer, "fps_embedding"):
fps = round(fps, 2)
log.info(f"Model has fps embedding, using {fps} fps")
fps_embeds = [fps]
fps_embeds = [0 if i == 16 else 1 for i in fps_embeds]
prefix_video = prefix_samples["samples"].to(noise) if prefix_samples is not None else None
prefix_video_latent_length = prefix_video.shape[2] if prefix_video is not None else 0
if prefix_video is not None:
log.info(f"Prefix video of length: {prefix_video_latent_length}")
init_latents[:, :prefix_video_latent_length] = prefix_video[0]
disable_enhance() #not sure if this can work, disabling for now to avoid errors if it's enabled by another sampler
freqs = None
transformer.rope_embedder.k = None
transformer.rope_embedder.num_frames = None
if rope_function=="comfy":
transformer.rope_embedder.k = 0
transformer.rope_embedder.num_frames = latent_video_length
else:
d = transformer.dim // transformer.num_heads
freqs = torch.cat([
rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=0),
rope_params(1024, 2 * (d // 6)),
rope_params(1024, 2 * (d // 6))
],
dim=1)
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1])
log.info(f"Seq len: {seq_len}")
seq_len = latent_video_length * 1560
log.info(f"Seq len: {seq_len}")
if args.preview_method in [LatentPreviewMethod.Auto, LatentPreviewMethod.Latent2RGB]: #default for latent2rgb
from latent_preview import prepare_callback
else:
from ..latent_preview import prepare_callback #custom for tiny VAE previews
#blockswap init
transformer_options = patcher.model_options.get("transformer_options", None)
if transformer_options is not None:
block_swap_args = transformer_options.get("block_swap_args", None)
if block_swap_args is not None:
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", True)
for name, param in transformer.named_parameters():
if "block" not 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, non_blocking=transformer.use_non_blocking)
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking)
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()
elif model["manual_offloading"]:
transformer.to(device)
use_fresca = False
if experimental_args is not None:
video_attention_split_steps = experimental_args.get("video_attention_split_steps", [])
if video_attention_split_steps:
transformer.video_attention_split_steps = [int(x.strip()) for x in video_attention_split_steps.split(",")]
else:
transformer.video_attention_split_steps = []
use_zero_init = experimental_args.get("use_zero_init", True)
use_cfg_zero_star = experimental_args.get("cfg_zero_star", False)
zero_star_steps = experimental_args.get("zero_star_steps", 0)
use_fresca = experimental_args.get("use_fresca", False)
if use_fresca:
fresca_scale_low = experimental_args.get("fresca_scale_low", 1.0)
fresca_scale_high = experimental_args.get("fresca_scale_high", 1.25)
fresca_freq_cutoff = experimental_args.get("fresca_freq_cutoff", 20)
#region model pred
def model_pred(z, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
vace_data=None, unianim_data=None, teacache_state=None, kv_cache=None, crossattn_cache=None, current_kv_cache_start=0, kv_start=0, kv_end=0):
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])):
nonlocal patcher
current_step_percentage = idx / len(timesteps)
control_lora_enabled = False
image_cond_input = image_cond
base_params = {
'seq_len': seq_len,
'device': device,
'freqs': freqs,
't': timestep,
'current_step': idx,
'control_lora_enabled': control_lora_enabled,
'vace_data': vace_data,
'unianim_data': unianim_data,
'kv_cache': kv_cache,
'crossattn_cache': crossattn_cache,
'current_kv_cache_start': current_kv_cache_start,
"kv_start": kv_start,
"kv_end": kv_end
}
#cond
noise_pred_cond, teacache_state_cond = transformer(
[z], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None,
clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage,
pred_id=teacache_state[0] if teacache_state else None,
**base_params
)
noise_pred_cond = noise_pred_cond[0].to(intermediate_device)
if use_fresca:
noise_pred_cond = fourier_filter(
noise_pred_cond,
scale_low=fresca_scale_low,
scale_high=fresca_scale_high,
freq_cutoff=fresca_freq_cutoff,
)
return noise_pred_cond, [teacache_state_cond]
def convert_flow_pred_to_x0(flow_pred: torch.Tensor, xt: torch.Tensor, timestep: torch.Tensor) -> torch.Tensor:
"""
Convert flow matching's prediction to x0 prediction.
flow_pred: the prediction with shape [B, C, H, W]
xt: the input noisy data with shape [B, C, H, W]
timestep: the timestep with shape [B]
pred = noise - x0
x_t = (1-sigma_t) * x0 + sigma_t * noise
we have x0 = x_t - sigma_t * pred
see derivations https://chatgpt.com/share/67bf8589-3d04-8008-bc6e-4cf1a24e2d0e
"""
# use higher precision for calculations
original_dtype = flow_pred.dtype
flow_pred, xt, sigmas, timesteps = map(
lambda x: x.double().to(flow_pred.device), [flow_pred, xt,
sample_scheduler.sigmas,
sample_scheduler.timesteps]
)
timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
x0_pred = xt - sigma_t * flow_pred
return x0_pred.to(original_dtype)
log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {init_latents.shape[3]*8}x{init_latents.shape[2]*8} with {steps} steps")
intermediate_device = device
#clear memory before sampling
mm.unload_all_models()
mm.soft_empty_cache()
gc.collect()
try:
torch.cuda.reset_peak_memory_stats(device)
except:
pass
#main loop
self.num_frame_per_block = 3
num_frames = noise.shape[1]
assert num_frames % self.num_frame_per_block == 0
num_blocks = num_frames // self.num_frame_per_block
print("num_blocks: ", num_blocks)
context_noise = 0
self.frame_seq_length = 1560
print("frame_seq_length: ", self.frame_seq_length)
self.cache_window_size = self.frame_seq_length * num_frames
self.shift_blocks = 1
self.kv_start = 0
self.kv_end = 0
output_latents = torch.zeros(
(target_shape[0],
target_shape[1],
target_shape[2],
target_shape[3]), device=device, dtype=dtype)
print("output_latents shape: ", output_latents.shape)
# Step 1: Initialize KV cache to all zeros
self._initialize_kv_cache(
batch_size=1,
dtype=noise.dtype,
device=cache_device,
num_blocks=transformer.num_layers,
num_heads=transformer.num_heads,
)
self._initialize_crossattn_cache(
batch_size=1,
dtype=noise.dtype,
device=cache_device,
num_blocks=transformer.num_layers,
num_heads=transformer.num_heads,
)
# Step 2: Cache context feature
current_kv_cache_start_frame = 0
num_input_frames = 0
# Step 3: Temporal denoising loop
all_num_frames = [self.num_frame_per_block] * num_blocks
print("all_num_frames", all_num_frames)
pbar = ProgressBar(num_blocks)
callback = prepare_callback(patcher, num_blocks)
for i,current_num_frames in enumerate(all_num_frames):
print("current_kv_cache_start_frame: ", current_kv_cache_start_frame)
#noisy_input = noise[:, current_kv_cache_start_frame - num_input_frames:current_kv_cache_start_frame + current_num_frames - num_input_frames]
noisy_input = noise[:, i * self.num_frame_per_block:(i + 1) * self.num_frame_per_block]
print("noisy_input shape: ", noisy_input.shape)
kv_end = self.kv_start + self.num_frame_per_block * self.frame_seq_length
print("kv_end: ", kv_end)
# Spatial denoising loop
for step_index, current_timestep in enumerate(timesteps):
print(f"current_timestep: {current_timestep}")
# set current timestep
timestep = torch.ones(
[1, current_num_frames],
device=noise.device,
dtype=torch.int64) * current_timestep
if step_index < len(timesteps) - 1:
flow_pred, self.teacache_state = model_pred(
noisy_input.to(dtype),
text_embeds["prompt_embeds"],
text_embeds["negative_prompt_embeds"],
timestep, step_index, image_cond, clip_fea,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_kv_cache_start=current_kv_cache_start_frame * self.frame_seq_length,
kv_start = self.kv_start,
kv_end = kv_end
)
#print("noise_pred shape: ", noise_pred.shape)
denoised_pred = convert_flow_pred_to_x0(
flow_pred=flow_pred.transpose(0, 1),
xt=noisy_input.transpose(0, 1),
timestep=timestep.flatten(0, 1)
)
next_timestep = timesteps[step_index + 1]
print("step_index: ", step_index, "next_timestep: ", next_timestep)
noisy_input = sample_scheduler.add_noise(
denoised_pred,
torch.randn_like(denoised_pred),
next_timestep * torch.ones(
[current_num_frames], device=noise.device, dtype=torch.long)
)
noisy_input = noisy_input.transpose(0, 1)
else:
# for getting real output
flow_pred, self.teacache_state = model_pred(
noisy_input.to(dtype),
text_embeds["prompt_embeds"],
text_embeds["negative_prompt_embeds"],
timestep, step_index, image_cond, clip_fea,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_kv_cache_start=current_kv_cache_start_frame * self.frame_seq_length,
kv_start = self.kv_start,
kv_end = kv_end
)
denoised_pred = convert_flow_pred_to_x0(
flow_pred=flow_pred.transpose(0, 1),
xt=noisy_input.transpose(0, 1),
timestep=timestep.flatten(0, 1)
)
denoised_pred = denoised_pred.transpose(0, 1)
# Step 3.2: record the model's output
#print("denoised_pred shape before output: ", denoised_pred.shape)
#output_latents[:, current_kv_cache_start_frame:current_kv_cache_start_frame + current_num_frames] = denoised_pred
output_latents[:, i * self.num_frame_per_block:(i + 1) * self.num_frame_per_block] = denoised_pred
# Step 3.3: rerun with timestep zero to update KV cache using clean context
print("cleaning KV cache")
context_timestep = torch.ones_like(timestep) * context_noise
model_pred(
denoised_pred.to(dtype),
text_embeds["prompt_embeds"],
text_embeds["negative_prompt_embeds"],
context_timestep, step_index, image_cond, clip_fea,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_kv_cache_start=current_kv_cache_start_frame * self.frame_seq_length,
kv_start = self.kv_start,
kv_end = kv_end
)
# Update positions for next block
r_shift_length = self.num_frame_per_block * self.frame_seq_length
self.kv_start += r_shift_length
kv_end += r_shift_length
#self.rope_start += r_shift_length
# Check if we need to shift the cache
if kv_end > self.cache_window_size:
print("Shifting KV cache")
kv_end -= self._shift_kv_cache()
# Step 3.4: update the start and end frame indices
current_kv_cache_start_frame += current_num_frames
if callback is not None:
#callback_latent = output_latents[:, :current_kv_cache_start_frame].float().detach().permute(1,0,2,3)
callback_latent = denoised_pred.float().detach().permute(1,0,2,3)
callback(i, callback_latent, None, num_blocks)
else:
pbar.update(1)
# reset cross attn cache
for block_index in range(transformer.num_layers):
self.crossattn_cache[block_index]["is_init"] = False
# reset kv cache
for block_index in range(len(self.kv_cache1)):
self.kv_cache1[block_index]["global_end_index"] = torch.tensor(
[0], dtype=torch.long, device=noise.device)
self.kv_cache1[block_index]["local_end_index"] = torch.tensor(
[0], dtype=torch.long, device=noise.device)
self.kv_cache1 = None
self.crossattn_cache = None
if force_offload:
if model["manual_offloading"]:
transformer.to(offload_device)
mm.soft_empty_cache()
gc.collect()
try:
print_memory(device)
torch.cuda.reset_peak_memory_stats(device)
except:
pass
return ({
"samples": output_latents.unsqueeze(0).cpu(),
}, )
NODE_CLASS_MAPPINGS = {
"WanVideoCausVidSampler": WanVideoCausVidSampler,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoCausVidSampler": "WanVideo CausVid Sampler",
}
+75 -1
View File
@@ -1,6 +1,7 @@
import numpy as np
from typing import Callable, Optional, List
import torch
from ..utils import log
def ordered_halving(val):
bin_str = f"{val:064b}"
@@ -182,3 +183,76 @@ def get_total_steps(
)
for i in range(len(timesteps))
)
def create_window_mask(noise_pred_context, c, latent_video_length, context_overlap, looped=False, window_type="linear"):
window_mask = torch.ones_like(noise_pred_context)
if window_type == "pyramid":
# Create pyramid weights that peak in the middle
length = noise_pred_context.shape[1]
if length % 2 == 0:
max_weight = length // 2
weight_sequence = list(range(1, max_weight + 1, 1)) + list(range(max_weight, 0, -1))
else:
max_weight = (length + 1) // 2
weight_sequence = list(range(1, max_weight, 1)) + [max_weight] + list(range(max_weight - 1, 0, -1))
# Normalize weights to range from 0 to 1
max_val = max(weight_sequence)
weight_sequence = [w / max_val for w in weight_sequence]
# Apply the weights to create the mask
weights_tensor = torch.tensor(weight_sequence, device=noise_pred_context.device)
weights_tensor = weights_tensor.view(1, -1, 1, 1)
window_mask = weights_tensor.expand_as(window_mask).clone()
# Adjust for position in sequence if needed
if not looped:
if min(c) == 0: # First chunk
left_ramp = torch.linspace(0, 1, context_overlap, device=noise_pred_context.device).view(1, -1, 1, 1)
# Clone to avoid in-place memory conflict
left_section = window_mask[:, :context_overlap].clone()
window_mask[:, :context_overlap] = torch.maximum(left_section, left_ramp)
if max(c) == latent_video_length - 1: # Last chunk
right_ramp = torch.linspace(1, 0, context_overlap, device=noise_pred_context.device).view(1, -1, 1, 1)
# Clone to avoid in-place memory conflict
right_section = window_mask[:, -context_overlap:].clone()
window_mask[:, -context_overlap:] = torch.maximum(right_section, right_ramp)
else: # Original "linear" window masking
# Apply left-side blending for all except first chunk (or always in loop mode)
if min(c) > 0 or (looped and max(c) == latent_video_length - 1):
ramp_up = torch.linspace(0, 1, context_overlap, device=noise_pred_context.device)
ramp_up = ramp_up.view(1, -1, 1, 1)
window_mask[:, :context_overlap] = ramp_up
# Apply right-side blending for all except last chunk (or always in loop mode)
if max(c) < latent_video_length - 1 or (looped and min(c) == 0):
ramp_down = torch.linspace(1, 0, context_overlap, device=noise_pred_context.device)
ramp_down = ramp_down.view(1, -1, 1, 1)
window_mask[:, -context_overlap:] = ramp_down
return window_mask
class WindowTracker:
def __init__(self, verbose=False):
self.window_map = {} # Maps frame sequence to persistent ID
self.next_id = 0
self.cache_states = {} # Maps persistent ID to teacache state
self.verbose = verbose
def get_window_id(self, frames):
key = tuple(sorted(frames)) # Order-independent frame sequence
if key not in self.window_map:
self.window_map[key] = self.next_id
if self.verbose:
log.info(f"New window pattern {key} -> ID {self.next_id}")
self.next_id += 1
return self.window_map[key]
def get_teacache(self, window_id, base_state):
if window_id not in self.cache_states:
if self.verbose:
log.info(f"Initializing persistent teacache for window {window_id}")
self.cache_states[window_id] = base_state.copy()
return self.cache_states[window_id]
+7 -4
View File
@@ -41,9 +41,11 @@ class WanVideoControlnetLoader:
model_path = folder_paths.get_full_path_or_raise("controlnet", model)
sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True)
num_layers = 8 if "blocks.7.scale_shift_table" in sd else 6
out_proj_dim = 5120 if num_layers == 6 else 1536
out_proj_dim = sd["controlnet_blocks.0.bias"].shape[0]
downscale_coef = 16 if out_proj_dim == 3072 else 8
vae_channels = 48 if out_proj_dim == 3072 else 16
if not "control_encoder.0.0.weight" in sd:
raise ValueError("Invalid ControlNet model")
@@ -52,7 +54,7 @@ class WanVideoControlnetLoader:
"added_kv_proj_dim": None,
"attention_head_dim": 128,
"cross_attn_norm": None,
"downscale_coef": 8,
"downscale_coef": downscale_coef,
"eps": 1e-06,
"ffn_dim": 8960,
"freq_dim": 256,
@@ -69,8 +71,9 @@ class WanVideoControlnetLoader:
"qk_norm": "rms_norm_across_heads",
"rope_max_seq_len": 1024,
"text_dim": 4096,
"vae_channels": 16
"vae_channels": vae_channels
}
print(f"Loading WanControlnet with config: {controlnet_cfg}")
from .wan_controlnet import WanControlnet
+32 -63
View File
@@ -10,20 +10,14 @@ from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscal
from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.transformers.transformer_wan import (
WanTimeTextImageEmbedding,
WanRotaryPosEmbed,
WanTimeTextImageEmbedding,
WanRotaryPosEmbed,
WanTransformerBlock
)
def zero_module(module):
for p in module.parameters():
nn.init.zeros_(p)
return module
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin):
r"""
A Controlnet Transformer model for video-like data used in the Wan model.
@@ -70,7 +64,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
_no_split_modules = ["WanTransformerBlock"]
_keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"]
_keys_to_ignore_on_load_unexpected = ["norm_added_q"]
@register_to_config
def __init__(
self,
@@ -101,16 +95,16 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
## Spatial compression with time awareness
nn.Sequential(
nn.Conv3d(
in_channels,
input_channels[0],
in_channels,
input_channels[0],
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)
),
nn.GELU(approximate="tanh"),
nn.GroupNorm(2, input_channels[0]),
),
## Temporal compression with spatial awareness
## Spatio-Temporal compression with spatial awareness
nn.Sequential(
nn.Conv3d(input_channels[0], input_channels[1], kernel_size=3, stride=(2, 1, 1), padding=1),
nn.GELU(approximate="tanh"),
@@ -123,9 +117,9 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
nn.GroupNorm(2, input_channels[2]),
)
])
inner_dim = num_attention_heads * attention_head_dim
# 1. Patch & position embedding
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)
@@ -154,11 +148,10 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
for _ in range(len(self.blocks)):
controlnet_block = nn.Linear(inner_dim, out_proj_dim)
controlnet_block = zero_module(controlnet_block)
self.controlnet_blocks.append(controlnet_block)
self.gradient_checkpointing = False
def forward(
self,
hidden_states: torch.Tensor,
@@ -183,28 +176,42 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
logger.warning(
"Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective."
)
rotary_emb = self.rope(hidden_states)
# 0. Controlnet encoder
for control_encoder_block in self.control_encoder:
controlnet_states = control_encoder_block(controlnet_states)
hidden_states = torch.cat([hidden_states, controlnet_states], dim=1)
## 1. Patch embedding and stack
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
if timestep.ndim == 2:
## for ComfyUI workflow
if hidden_states.shape[1] != timestep.shape[1]:
timestep = timestep.repeat_interleave(hidden_states.shape[1] // timestep.shape[1], dim=1)
ts_seq_len = timestep.shape[1]
timestep = timestep.flatten() # batch_size * seq_len
else:
ts_seq_len = None
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image
timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len
)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if ts_seq_len is not None:
# batch_size, seq_len, 6, inner_dim
timestep_proj = timestep_proj.unflatten(2, (6, -1))
else:
# batch_size, 6, inner_dim
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
# 2. Transformer blocks
# 4. Transformer blocks
controlnet_hidden_states = ()
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block, controlnet_block in zip(self.blocks, self.controlnet_blocks):
@@ -226,42 +233,4 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
return (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,
}
controlnet = WanControlnet(**parameters)
hidden_states = torch.rand(1, 16, 21, 60, 90)
timestep = 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, 81, 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)
+281
View File
@@ -0,0 +1,281 @@
import torch
import torch.nn as nn
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
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())
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 + "."
module_prefix = module_prefix.replace("_orig_mod.", "")
_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 and "dual_controller" not in module_prefix and name not in modules_to_not_convert:
weight_key = module_prefix + "weight"
if weight_key not in state_dict:
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_weight = scale_weights.get(scale_key)
with init_empty_weights():
model._modules[name] = CustomLinear(
in_features,
out_features,
module.bias is not None,
compute_dtype=compute_dtype,
scale_weight=scale_weight,
allow_compile=allow_compile,
is_gguf=is_gguf
)
model._modules[name].source_cls = type(module)
model._modules[name].requires_grad_(False)
return model
def set_lora_params(module, patches, module_prefix="", device=torch.device("cpu")):
remove_lora_from_module(module)
# Recursively set lora_diffs and lora_strengths for all CustomLinear layers
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(child, patches, child_prefix, device)
if isinstance(module, CustomLinear):
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, [])
#print(f"Processing LoRA patches for {key}: {len(patch)} patches found")
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
lora_strengths = [p[0] for p in patch]
module.set_lora_diffs(lora_diffs, device=device)
module.set_lora_strengths(lora_strengths, device=device)
module._step.fill_(0) # Initialize step for LoRA scheduling
class CustomLinear(nn.Linear):
def __init__(
self,
in_features,
out_features,
bias=False,
compute_dtype=None,
device=None,
scale_weight=None,
allow_compile=False,
is_gguf=False
) -> None:
super().__init__(in_features, out_features, bias, device)
self.compute_dtype = compute_dtype
self.lora_diffs = []
self.register_buffer("_step", torch.zeros((), dtype=torch.long))
self.scale_weight = scale_weight
self.lora_strengths = []
self.allow_compile = allow_compile
self.is_gguf = is_gguf
if not allow_compile:
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")):
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 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):
"""Apply LoRA using custom ops to avoid graph breaks"""
if not hasattr(self, "lora_diff_0_0"):
return weight
for idx, lora_diff_names in enumerate(self.lora_diffs):
lora_strength = self._get_lora_strength(idx)
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])
weight = self._apply_lora_impl(
weight, lora_diff_0, lora_diff_1,
float(lora_diff_2) if lora_diff_2 is not None else 0.0, lora_strength
)
else:
lora_diff = getattr(self, lora_diff_names)
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
def forward(self, input):
weight = self._prepare_weight(input)
if self.bias is not None:
bias = self.bias.to(input if not self.is_gguf else self.compute_dtype)
else:
bias = 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():
weight = weight * self.scale_weight
else:
input = input * self.scale_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)
def remove_lora_from_module(module):
for name, submodule in module.named_modules():
if hasattr(submodule, "lora_diffs"):
for i in range(len(submodule.lora_diffs)):
if hasattr(submodule, f"lora_diff_{i}_0"):
delattr(submodule, f"lora_diff_{i}_0")
if hasattr(submodule, f"lora_diff_{i}_1"):
delattr(submodule, f"lora_diff_{i}_1")
if hasattr(submodule, f"lora_diff_{i}_2"):
delattr(submodule, f"lora_diff_{i}_2")
+104
View File
@@ -0,0 +1,104 @@
import torch
from comfy.model_management import get_autocast_device, get_torch_device
@torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False)
@torch.compiler.disable()
def rope_apply_z(x, grid_sizes, freqs, inner_t, shift=6):
n, c = x.size(2), x.size(3) // 2
# loop over samples
output = []
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
seq_len = f * h * w
# precompute multipliers
x_i = torch.view_as_complex(
x[i, :seq_len].to(torch.float64).reshape(seq_len, n, -1, 2)
)
start_ind = [sum(inner_t[i][:_]) for _ in range(len(inner_t[i]))]
end_ind = [sum(inner_t[i][:_+1]) for _ in range(len(inner_t[i]))]
freq_select = []
for shot_ind, (s, e) in enumerate(zip(start_ind, end_ind)):
freq_select += [shot_ind * shift] * (e - s)
shot_freqs = freqs[freq_select]
freqs_i = shot_freqs.view(f, 1, 1, -1).expand(f, h, w, -1).reshape(seq_len, 1, -1)
# apply rotary embedding
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
x_i = torch.cat([x_i, x[i, seq_len:]])
# append to collection
output.append(x_i)
return torch.stack(output).float()
@torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False)
@torch.compiler.disable()
def rope_apply_c(x, freqs, inner_c, shift=6):
b, s, n, c = x.size(0), x.size(1), x.size(2), x.size(3) // 2
# loop over samples
output = []
for i in range(b):
# precompute multipliers
x_i = torch.view_as_complex(
x[i].to(torch.float64).reshape(s, n, -1, 2)
)
freq_select = []
for shot_ind, c_len in enumerate(inner_c[i]):
freq_select += [shot_ind * shift] * c_len
freq_select += [shot_ind+10] * (s-len(freq_select)) # extra suppression for the empty token
shot_freqs = freqs[freq_select]
freqs_i = shot_freqs.view(s, 1, -1)
# apply rotary embedding
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
# append to collection
output.append(x_i)
return torch.stack(output).float()
@torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False)
@torch.compiler.disable()
def rope_apply_echoshot(x, grid_sizes, freqs, inner_t, shift=4):
n, c = x.size(2), x.size(3) // 2
# split freqs
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
# loop over samples
output = []
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
seq_len = f * h * w
# precompute multipliers
x_i = torch.view_as_complex(
x[i, :seq_len].to(torch.float64).reshape(seq_len, n, -1, 2)
)
start_ind = [sum(inner_t[i][:_]) for _ in range(len(inner_t[i]))]
end_ind = [sum(inner_t[i][:_+1]) for _ in range(len(inner_t[i]))]
freq_select = []
for shot_ind, (s, e) in enumerate(zip(start_ind, end_ind)):
freq_select += list(range(shot_ind * shift + s, shot_ind * shift + e))
t_freqs = freqs[0][freq_select]
freqs_i = torch.cat([
# freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
t_freqs.view(f, 1, 1, -1).expand(f, h, w, -1), ###
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
], dim=-1).reshape(seq_len, 1, -1)
# apply rotary embedding
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
x_i = torch.cat([x_i, x[i, seq_len:]])
# append to collection
output.append(x_i)
return torch.stack(output).float()
File diff suppressed because it is too large Load Diff
File diff suppressed because one or more lines are too long
Binary file not shown.
File diff suppressed because it is too large Load Diff
@@ -1587,9 +1587,9 @@
"link": null
},
{
"name": "teacache_args",
"name": "cache_args",
"shape": 7,
"type": "TEACACHEARGS",
"type": "CACHEARGS",
"link": 335
},
{
@@ -1668,8 +1668,8 @@
"inputs": [],
"outputs": [
{
"name": "teacache_args",
"type": "TEACACHEARGS",
"name": "cache_args",
"type": "CACHEARGS",
"links": [
335
]
File diff suppressed because it is too large Load Diff
@@ -1751,8 +1751,8 @@
"inputs": [],
"outputs": [
{
"name": "teacache_args",
"type": "TEACACHEARGS",
"name": "cache_args",
"type": "CACHEARGS",
"links": [
334
]
@@ -1789,8 +1789,8 @@
"inputs": [],
"outputs": [
{
"name": "teacache_args",
"type": "TEACACHEARGS",
"name": "cache_args",
"type": "CACHEARGS",
"links": [
335
]
@@ -2252,8 +2252,8 @@
"inputs": [],
"outputs": [
{
"name": "teacache_args",
"type": "TEACACHEARGS",
"name": "cache_args",
"type": "CACHEARGS",
"links": [
350
]
@@ -2840,9 +2840,9 @@
"link": null
},
{
"name": "teacache_args",
"name": "cache_args",
"shape": 7,
"type": "TEACACHEARGS",
"type": "CACHEARGS",
"link": 350
},
{
@@ -2966,9 +2966,9 @@
"link": null
},
{
"name": "teacache_args",
"name": "cache_args",
"shape": 7,
"type": "TEACACHEARGS",
"type": "CACHEARGS",
"link": 335
},
{
@@ -5220,9 +5220,9 @@
"link": null
},
{
"name": "teacache_args",
"name": "cache_args",
"shape": 7,
"type": "TEACACHEARGS",
"type": "CACHEARGS",
"link": 334
},
{
@@ -176,8 +176,8 @@
"inputs": [],
"outputs": [
{
"name": "teacache_args",
"type": "TEACACHEARGS",
"name": "cache_args",
"type": "CACHEARGS",
"links": [
103
]
@@ -826,9 +826,9 @@
"link": null
},
{
"name": "teacache_args",
"name": "cache_args",
"shape": 7,
"type": "TEACACHEARGS",
"type": "CACHEARGS",
"link": 103
},
{
@@ -747,8 +747,8 @@
"inputs": [],
"outputs": [
{
"name": "teacache_args",
"type": "TEACACHEARGS",
"name": "cache_args",
"type": "CACHEARGS",
"links": [
56
]
@@ -1263,9 +1263,9 @@
"link": null
},
{
"name": "teacache_args",
"name": "cache_args",
"shape": 7,
"type": "TEACACHEARGS",
"type": "CACHEARGS",
"link": 56
},
{
@@ -794,8 +794,8 @@
"inputs": [],
"outputs": [
{
"name": "teacache_args",
"type": "TEACACHEARGS",
"name": "cache_args",
"type": "CACHEARGS",
"links": [
56
]
@@ -1719,9 +1719,9 @@
"link": null
},
{
"name": "teacache_args",
"name": "cache_args",
"shape": 7,
"type": "TEACACHEARGS",
"type": "CACHEARGS",
"link": 56
},
{
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
@@ -485,9 +485,9 @@
"link": null
},
{
"name": "teacache_args",
"name": "cache_args",
"shape": 7,
"type": "TEACACHEARGS",
"type": "CACHEARGS",
"link": 89
},
{
@@ -1699,8 +1699,8 @@
"inputs": [],
"outputs": [
{
"name": "teacache_args",
"type": "TEACACHEARGS",
"name": "cache_args",
"type": "CACHEARGS",
"links": [
89
]
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 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 it is too large Load Diff
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 it is too large Load Diff
@@ -587,8 +587,8 @@
"inputs": [],
"outputs": [
{
"name": "teacache_args",
"type": "TEACACHEARGS",
"name": "cache_args",
"type": "CACHEARGS",
"links": [
56
]
@@ -3788,9 +3788,9 @@
"link": null
},
{
"name": "teacache_args",
"name": "cache_args",
"shape": 7,
"type": "TEACACHEARGS",
"type": "CACHEARGS",
"link": 56
},
{
@@ -2,7 +2,7 @@
"id": "206247b6-9fec-4ed2-8927-e4f388c674d4",
"revision": 0,
"last_node_id": 196,
"last_link_id": 303,
"last_link_id": 306,
"nodes": [
{
"id": 42,
@@ -1012,8 +1012,8 @@
"inputs": [],
"outputs": [
{
"name": "teacache_args",
"type": "TEACACHEARGS",
"name": "cache_args",
"type": "CACHEARGS",
"links": [
195
]
@@ -1051,8 +1051,8 @@
"mode": 0,
"inputs": [
{
"name": "TEACACHEARGS",
"type": "TEACACHEARGS",
"name": "CACHEARGS",
"type": "CACHEARGS",
"link": 195
}
],
@@ -1109,38 +1109,6 @@
"ExpArgs"
]
},
{
"id": 140,
"type": "GetNode",
"pos": [
724.8881225585938,
-579.0637817382812
],
"size": [
210,
34
],
"flags": {
"collapsed": true
},
"order": 18,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "TEACACHEARGS",
"type": "TEACACHEARGS",
"links": [
235
]
}
],
"title": "Get_TeaCache",
"properties": {},
"widgets_values": [
"TeaCache"
]
},
{
"id": 141,
"type": "GetNode",
@@ -1155,7 +1123,7 @@
"flags": {
"collapsed": true
},
"order": 19,
"order": 18,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1185,7 +1153,7 @@
226
],
"flags": {},
"order": 20,
"order": 19,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1225,7 +1193,7 @@
106
],
"flags": {},
"order": 21,
"order": 20,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1300,7 +1268,7 @@
"flags": {
"collapsed": true
},
"order": 22,
"order": 21,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1332,15 +1300,16 @@
"flags": {
"collapsed": true
},
"order": 23,
"order": 22,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "TEACACHEARGS",
"type": "TEACACHEARGS",
"name": "CACHEARGS",
"type": "CACHEARGS",
"links": [
196
196,
305
]
}
],
@@ -1364,7 +1333,7 @@
"flags": {
"collapsed": true
},
"order": 24,
"order": 23,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1396,7 +1365,7 @@
"flags": {
"collapsed": true
},
"order": 25,
"order": 24,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1423,7 +1392,7 @@
],
"size": [
428.4000244140625,
860.4000244140625
880.4000244140625
],
"flags": {},
"order": 88,
@@ -1457,10 +1426,10 @@
"link": 183
},
{
"name": "teacache_args",
"name": "cache_args",
"shape": 7,
"type": "TEACACHEARGS",
"link": 235
"type": "CACHEARGS",
"link": 304
},
{
"name": "slg_args",
@@ -1506,7 +1475,8 @@
true,
"unipc",
1,
"comfy"
"comfy",
""
]
},
{
@@ -1521,7 +1491,7 @@
154
],
"flags": {},
"order": 26,
"order": 25,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1559,7 +1529,7 @@
58
],
"flags": {},
"order": 27,
"order": 26,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1594,7 +1564,7 @@
"flags": {
"collapsed": true
},
"order": 28,
"order": 27,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1686,7 +1656,7 @@
"flags": {
"collapsed": true
},
"order": 29,
"order": 28,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1771,7 +1741,7 @@
"flags": {
"collapsed": true
},
"order": 30,
"order": 29,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1805,7 +1775,7 @@
"flags": {
"collapsed": true
},
"order": 31,
"order": 30,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1839,7 +1809,7 @@
"flags": {
"collapsed": true
},
"order": 32,
"order": 31,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1873,15 +1843,16 @@
"flags": {
"collapsed": true
},
"order": 33,
"order": 32,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "TEACACHEARGS",
"type": "TEACACHEARGS",
"name": "CACHEARGS",
"type": "CACHEARGS",
"links": [
260
260,
306
]
}
],
@@ -1905,7 +1876,7 @@
"flags": {
"collapsed": true
},
"order": 34,
"order": 33,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1937,7 +1908,7 @@
"flags": {
"collapsed": true
},
"order": 35,
"order": 34,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1955,101 +1926,6 @@
"ExpArgs"
]
},
{
"id": 165,
"type": "WanVideoDiffusionForcingSampler",
"pos": [
5483.89599609375,
-510.88037109375
],
"size": [
428.4000244140625,
860.4000244140625
],
"flags": {},
"order": 103,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "WANVIDEOMODEL",
"link": 256
},
{
"name": "text_embeds",
"type": "WANVIDEOTEXTEMBEDS",
"link": 296
},
{
"name": "image_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"link": 258
},
{
"name": "samples",
"shape": 7,
"type": "LATENT",
"link": null
},
{
"name": "prefix_samples",
"shape": 7,
"type": "LATENT",
"link": 259
},
{
"name": "teacache_args",
"shape": 7,
"type": "TEACACHEARGS",
"link": 260
},
{
"name": "slg_args",
"shape": 7,
"type": "SLGARGS",
"link": 261
},
{
"name": "experimental_args",
"shape": 7,
"type": "EXPERIMENTALARGS",
"link": 262
},
{
"name": "unianimate_poses",
"shape": 7,
"type": "UNIANIMATE_POSE",
"link": null
}
],
"outputs": [
{
"name": "samples",
"type": "LATENT",
"links": [
252
]
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "e5a326c9811514f2c08c89bccea9a7c731d9a503",
"Node name for S&R": "WanVideoDiffusionForcingSampler"
},
"widgets_values": [
10,
24.000000000000004,
30,
4.000000000000001,
5.000000000000001,
0,
"fixed",
true,
"unipc",
1,
"comfy"
]
},
{
"id": 153,
"type": "WanVideoEncode",
@@ -2130,6 +2006,7 @@
},
{
"name": "negative",
"shape": 7,
"type": "CONDITIONING",
"link": 55
}
@@ -2207,7 +2084,7 @@
60
],
"flags": {},
"order": 36,
"order": 35,
"mode": 0,
"inputs": [],
"outputs": [
@@ -2343,7 +2220,7 @@
"flags": {
"collapsed": true
},
"order": 37,
"order": 36,
"mode": 0,
"inputs": [],
"outputs": [
@@ -2654,101 +2531,6 @@
80
]
},
{
"id": 104,
"type": "WanVideoDiffusionForcingSampler",
"pos": [
3010,
-400
],
"size": [
428.4000244140625,
860.4000244140625
],
"flags": {},
"order": 97,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "WANVIDEOMODEL",
"link": 190
},
{
"name": "text_embeds",
"type": "WANVIDEOTEXTEMBEDS",
"link": 192
},
{
"name": "image_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"link": 207
},
{
"name": "samples",
"shape": 7,
"type": "LATENT",
"link": null
},
{
"name": "prefix_samples",
"shape": 7,
"type": "LATENT",
"link": 180
},
{
"name": "teacache_args",
"shape": 7,
"type": "TEACACHEARGS",
"link": 196
},
{
"name": "slg_args",
"shape": 7,
"type": "SLGARGS",
"link": 240
},
{
"name": "experimental_args",
"shape": 7,
"type": "EXPERIMENTALARGS",
"link": 198
},
{
"name": "unianimate_poses",
"shape": 7,
"type": "UNIANIMATE_POSE",
"link": null
}
],
"outputs": [
{
"name": "samples",
"type": "LATENT",
"links": [
178
]
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "e5a326c9811514f2c08c89bccea9a7c731d9a503",
"Node name for S&R": "WanVideoDiffusionForcingSampler"
},
"widgets_values": [
24,
24.000000000000004,
30,
4.000000000000001,
5.000000000000001,
0,
"fixed",
true,
"unipc",
1,
"comfy"
]
},
{
"id": 90,
"type": "VHS_VideoCombine",
@@ -3137,7 +2919,7 @@
"flags": {
"collapsed": true
},
"order": 38,
"order": 37,
"mode": 0,
"inputs": [],
"outputs": [
@@ -3434,7 +3216,7 @@
203.9819793701172
],
"flags": {},
"order": 39,
"order": 38,
"mode": 0,
"inputs": [],
"outputs": [
@@ -3474,7 +3256,7 @@
139.27662658691406
],
"flags": {},
"order": 40,
"order": 39,
"mode": 0,
"inputs": [],
"outputs": [],
@@ -3621,7 +3403,7 @@
"flags": {
"collapsed": true
},
"order": 41,
"order": 40,
"mode": 0,
"inputs": [],
"outputs": [
@@ -3663,6 +3445,7 @@
},
{
"name": "negative",
"shape": 7,
"type": "CONDITIONING",
"link": 265
}
@@ -3695,7 +3478,7 @@
95.97175598144531
],
"flags": {},
"order": 42,
"order": 41,
"mode": 0,
"inputs": [],
"outputs": [],
@@ -3718,7 +3501,7 @@
106
],
"flags": {},
"order": 43,
"order": 42,
"mode": 0,
"inputs": [],
"outputs": [
@@ -3753,10 +3536,10 @@
],
"size": [
528.6734619140625,
234
254
],
"flags": {},
"order": 44,
"order": 43,
"mode": 0,
"inputs": [
{
@@ -3788,6 +3571,12 @@
"shape": 7,
"type": "VACEPATH",
"link": null
},
{
"name": "fantasytalking_model",
"shape": 7,
"type": "FANTASYTALKINGMODEL",
"link": null
}
],
"outputs": [
@@ -3829,7 +3618,7 @@
"flags": {
"collapsed": true
},
"order": 45,
"order": 44,
"mode": 0,
"inputs": [],
"outputs": [
@@ -3977,7 +3766,7 @@
88
],
"flags": {},
"order": 46,
"order": 45,
"mode": 0,
"inputs": [],
"outputs": [],
@@ -4000,7 +3789,7 @@
451.9747314453125
],
"flags": {},
"order": 47,
"order": 46,
"mode": 0,
"inputs": [
{
@@ -4090,6 +3879,12 @@
"type": "IMAGE",
"link": 298
},
{
"name": "get_image_size",
"shape": 7,
"type": "IMAGE",
"link": null
},
{
"name": "width_input",
"shape": 7,
@@ -4101,12 +3896,6 @@
"shape": 7,
"type": "INT",
"link": null
},
{
"name": "get_image_size",
"shape": 7,
"type": "IMAGE",
"link": null
}
],
"outputs": [
@@ -4158,7 +3947,7 @@
85.25131225585938
],
"flags": {},
"order": 48,
"order": 47,
"mode": 0,
"inputs": [],
"outputs": [
@@ -4233,7 +4022,7 @@
88
],
"flags": {},
"order": 49,
"order": 48,
"mode": 0,
"inputs": [],
"outputs": [],
@@ -4371,7 +4160,7 @@
314
],
"flags": {},
"order": 50,
"order": 49,
"mode": 0,
"inputs": [],
"outputs": [
@@ -4452,7 +4241,7 @@
209.54696655273438
],
"flags": {},
"order": 51,
"order": 50,
"mode": 0,
"inputs": [],
"outputs": [],
@@ -4524,7 +4313,7 @@
130
],
"flags": {},
"order": 52,
"order": 51,
"mode": 4,
"inputs": [],
"outputs": [
@@ -4564,7 +4353,7 @@
"flags": {
"collapsed": true
},
"order": 53,
"order": 52,
"mode": 0,
"inputs": [],
"outputs": [
@@ -4658,7 +4447,7 @@
"flags": {
"collapsed": true
},
"order": 54,
"order": 53,
"mode": 0,
"inputs": [],
"outputs": [
@@ -4692,7 +4481,7 @@
"flags": {
"collapsed": true
},
"order": 55,
"order": 54,
"mode": 0,
"inputs": [],
"outputs": [
@@ -4724,7 +4513,7 @@
166.29330444335938
],
"flags": {},
"order": 56,
"order": 55,
"mode": 0,
"inputs": [],
"outputs": [],
@@ -4734,6 +4523,229 @@
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 140,
"type": "GetNode",
"pos": [
724.8883666992188,
-591.3699340820312
],
"size": [
210,
50
],
"flags": {
"collapsed": true
},
"order": 56,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "CACHEARGS",
"type": "CACHEARGS",
"links": [
235,
304
]
}
],
"title": "Get_TeaCache",
"properties": {},
"widgets_values": [
"TeaCache"
]
},
{
"id": 104,
"type": "WanVideoDiffusionForcingSampler",
"pos": [
3010,
-400
],
"size": [
428.4000244140625,
860.4000244140625
],
"flags": {},
"order": 97,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "WANVIDEOMODEL",
"link": 190
},
{
"name": "text_embeds",
"type": "WANVIDEOTEXTEMBEDS",
"link": 192
},
{
"name": "image_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"link": 207
},
{
"name": "samples",
"shape": 7,
"type": "LATENT",
"link": null
},
{
"name": "prefix_samples",
"shape": 7,
"type": "LATENT",
"link": 180
},
{
"name": "cache_args",
"shape": 7,
"type": "CACHEARGS",
"link": 305
},
{
"name": "slg_args",
"shape": 7,
"type": "SLGARGS",
"link": 240
},
{
"name": "experimental_args",
"shape": 7,
"type": "EXPERIMENTALARGS",
"link": 198
},
{
"name": "unianimate_poses",
"shape": 7,
"type": "UNIANIMATE_POSE",
"link": null
}
],
"outputs": [
{
"name": "samples",
"type": "LATENT",
"links": [
178
]
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "e5a326c9811514f2c08c89bccea9a7c731d9a503",
"Node name for S&R": "WanVideoDiffusionForcingSampler"
},
"widgets_values": [
24,
24.000000000000004,
30,
4.000000000000001,
5.000000000000001,
0,
"fixed",
true,
"unipc",
1,
"comfy"
]
},
{
"id": 165,
"type": "WanVideoDiffusionForcingSampler",
"pos": [
5483.89599609375,
-510.88037109375
],
"size": [
428.4000244140625,
860.4000244140625
],
"flags": {},
"order": 103,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "WANVIDEOMODEL",
"link": 256
},
{
"name": "text_embeds",
"type": "WANVIDEOTEXTEMBEDS",
"link": 296
},
{
"name": "image_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"link": 258
},
{
"name": "samples",
"shape": 7,
"type": "LATENT",
"link": null
},
{
"name": "prefix_samples",
"shape": 7,
"type": "LATENT",
"link": 259
},
{
"name": "cache_args",
"shape": 7,
"type": "CACHEARGS",
"link": 306
},
{
"name": "slg_args",
"shape": 7,
"type": "SLGARGS",
"link": 261
},
{
"name": "experimental_args",
"shape": 7,
"type": "EXPERIMENTALARGS",
"link": 262
},
{
"name": "unianimate_poses",
"shape": 7,
"type": "UNIANIMATE_POSE",
"link": null
}
],
"outputs": [
{
"name": "samples",
"type": "LATENT",
"links": [
252
]
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "e5a326c9811514f2c08c89bccea9a7c731d9a503",
"Node name for S&R": "WanVideoDiffusionForcingSampler"
},
"widgets_values": [
10,
24.000000000000004,
30,
4.000000000000001,
5.000000000000001,
0,
"fixed",
true,
"unipc",
1,
"comfy"
]
}
],
"links": [
@@ -4889,14 +4901,6 @@
0,
"*"
],
[
196,
115,
0,
104,
5,
"TEACACHEARGS"
],
[
197,
87,
@@ -5089,14 +5093,6 @@
0,
"IMAGE"
],
[
235,
140,
0,
103,
5,
"TEACACHEARGS"
],
[
236,
141,
@@ -5217,14 +5213,6 @@
4,
"LATENT"
],
[
260,
162,
0,
165,
5,
"TEACACHEARGS"
],
[
261,
163,
@@ -5456,6 +5444,30 @@
156,
2,
"INT"
],
[
304,
140,
0,
103,
5,
"CACHEARGS"
],
[
305,
115,
0,
104,
5,
"CACHEARGS"
],
[
306,
162,
0,
165,
5,
"CACHEARGS"
]
],
"groups": [
@@ -5528,13 +5540,13 @@
"config": {},
"extra": {
"ds": {
"scale": 1.191817653772724,
"scale": 0.611590904484147,
"offset": [
1695.7620823297345,
1138.5391291690546
426.87167769967925,
1142.3743330459465
]
},
"frontendVersion": "1.17.3",
"frontendVersion": "1.22.0",
"node_versions": {
"ComfyUI-WanVideoWrapper": "5a2383621a05825d0d0437781afcb8552d9590fd",
"comfy-core": "0.3.26",
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
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 one or more lines are too long
File diff suppressed because it is too large Load Diff

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