162 Commits
Author SHA1 Message Date
kijai fdb23dec7d Update model.py 2026-01-05 22:11:04 +02:00
kijai 07d7d8ca8e remove prints 2026-01-05 22:10:02 +02:00
kijai 01869d4bf5 Merge branch 'main' into longvie2 2026-01-05 18:47:48 +02:00
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 55c672028b Merge branch 'main' into longvie2 2025-12-29 15:39:43 +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 b551ec9e31 Merge branch 'main' into longvie2 2025-12-29 15:03:53 +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
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
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 9f019d7dfb Merge branch 'main' into longvie2 2025-12-23 23:40:25 +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 fc5322fae4 Merge branch 'main' into longvie2 2025-12-23 22:04:15 +02:00
kijai e75f814312 LongCat-Avatar example 2025-12-23 20:43:10 +02:00
kijai 222fc70eb7 Update nodes.py 2025-12-23 17:18:55 +02:00
kijai 8509236da1 init 2025-12-23 14:20:18 +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
wzxysf 6a37c0b2d6 Enable vae tiling with end frame 2025-10-24 22:23:31 +08:00
cmeka f653544f7b Merge branch 'kijai:main' into custom-sigmas 2025-10-21 12:48:02 -04:00
cmeka 425035d810 Fix custom sigmas for supported schedulers
Previously, only unipc, dpm++, and dpm++_sde schedulers preserved custom input sigmas exactly. Other schedulers such as (euler, lcm, deis,
  etc.) would transform or modify the sigmas through their set_timesteps() methods, causing inconsistent behavior.
2025-10-21 12:46:58 -04:00
86 changed files with 36324 additions and 13558 deletions
+149 -2
View File
@@ -2,6 +2,10 @@ 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__(
@@ -22,7 +26,7 @@ class FeedForwardSwiGLU(nn.Module):
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.
@@ -62,4 +66,147 @@ class TimestepEmbedder(nn.Module):
if t_freq.dtype != dtype:
t_freq = t_freq.to(dtype)
t_emb = self.mlp(t_freq)
return t_emb
return t_emb
class SingleStreamAttention(nn.Module):
def __init__(
self,
dim: int,
encoder_hidden_states_dim: int,
num_heads: int,
qkv_bias: bool,
qk_norm: bool,
attn_drop: float = 0.0,
proj_drop: float = 0.0,
eps: float = 1e-6,
class_range: int = 24,
class_interval: int = 4,
attention_mode: str = "sdpa",
) -> None:
super().__init__()
assert dim % num_heads == 0, "dim should be divisible by num_heads"
self.dim = dim
self.encoder_hidden_states_dim = encoder_hidden_states_dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.scale = self.head_dim**-0.5
self.q_linear = nn.Linear(dim, dim, bias=qkv_bias)
self.q_norm = WanRMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
self.attn_drop = nn.Dropout(attn_drop)
self.proj = nn.Linear(dim, dim)
self.proj_drop = nn.Dropout(proj_drop)
self.kv_linear = nn.Linear(encoder_hidden_states_dim, dim * 2, bias=qkv_bias)
self.k_norm = WanRMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
self.attention_mode = attention_mode
# multitalk related params
self.class_interval = class_interval
self.class_range = class_range
self.rope_h1 = (0, self.class_interval)
self.rope_h2 = (self.class_range - self.class_interval, self.class_range)
self.rope_bak = int(self.class_range // 2)
self.rope_1d = RotaryPositionalEmbedding1D(self.head_dim)
def _process_cross_attn(self, x, cond, frames_num=None, x_ref_attn_map=None):
N_t = frames_num
out_dtype = x.dtype
x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t)
# get q for hidden_state
B, N, C = x.shape
q = self.q_linear(x)
q_shape = (B, N, self.num_heads, self.head_dim)
q = q.view(q_shape).permute((0, 2, 1, 3)) # [B, H, N, D]
q = self.q_norm(q.to(self.q_norm.weight.dtype)).to(q.dtype)
# multitalk with rope1d pe
if x_ref_attn_map is not None:
max_values = x_ref_attn_map.max(1).values[:, None, None]
min_values = x_ref_attn_map.min(1).values[:, None, None]
max_min_values = torch.cat([max_values, min_values], dim=2)
human1_max_value, human1_min_value = max_min_values[0, :, 0].max(), max_min_values[0, :, 1].min()
human2_max_value, human2_min_value = max_min_values[1, :, 0].max(), max_min_values[1, :, 1].min()
human1 = normalize_and_scale(x_ref_attn_map[0], (human1_min_value, human1_max_value), (self.rope_h1[0], self.rope_h1[1]))
human2 = normalize_and_scale(x_ref_attn_map[1], (human2_min_value, human2_max_value), (self.rope_h2[0], self.rope_h2[1]))
back = torch.full((x_ref_attn_map.size(1),), self.rope_bak, dtype=human1.dtype).to(human1.device)
max_indices = x_ref_attn_map.argmax(dim=0)
normalized_map = torch.stack([human1, human2, back], dim=1)
normalized_pos = normalized_map[range(x_ref_attn_map.size(1)), max_indices]
q = rearrange(q, "(B N_t) H S C -> B H (N_t S) C", N_t=N_t)
q = self.rope_1d(q, normalized_pos)
q = rearrange(q, "B H (N_t S) C -> (B N_t) H S C", N_t=N_t)
# get kv from encoder_hidden_states
_, N_a, _ = cond.shape
encoder_kv = self.kv_linear(cond)
encoder_kv_shape = (B, N_a, 2, self.num_heads, self.head_dim)
encoder_kv = encoder_kv.view(encoder_kv_shape).permute((2, 0, 3, 1, 4))
encoder_k, encoder_v = encoder_kv.unbind(0)
encoder_k = self.k_norm(encoder_k.to(self.k_norm.weight.dtype)).to(encoder_k.dtype)
# multitalk with rope1d pe
if x_ref_attn_map is not None:
per_frame = torch.zeros(N_a, dtype=encoder_k.dtype).to(encoder_k.device)
per_frame[:per_frame.size(0)//2] = (self.rope_h1[0] + self.rope_h1[1]) / 2
per_frame[per_frame.size(0)//2:] = (self.rope_h2[0] + self.rope_h2[1]) / 2
encoder_pos = torch.concat([per_frame]*N_t, dim=0)
encoder_k = rearrange(encoder_k, "(B N_t) H S C -> B H (N_t S) C", N_t=N_t)
encoder_k = self.rope_1d(encoder_k, encoder_pos)
encoder_k = rearrange(encoder_k, "B H (N_t S) C -> (B N_t) H S C", N_t=N_t)
# Input tensors must be in format ``[B, M, H, K]``, where B is the batch size, M \
# the sequence length, H the number of heads, and K the embeding size per head
q = rearrange(q, "B H M K -> B M H K")
encoder_k = rearrange(encoder_k, "B H M K -> B M H K")
encoder_v = rearrange(encoder_v, "B H M K -> B M H K")
x = attention(q, encoder_k, encoder_v, attention_mode=self.attention_mode)
x = rearrange(x, "B M H K -> B H M K")
# linear transform
x_output_shape = (B, N, C)
x = x.transpose(1, 2)
x = x.reshape(x_output_shape)
x = self.proj(x)
x = self.proj_drop(x)
# reshape x to origin shape
x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t)
return x.type(out_dtype)
def forward(self, x, cond, num_latent_frames=None, num_cond_latents=None, x_ref_attn_map=None, human_num=None):
B, N, C = x.shape
if (num_cond_latents is None or num_cond_latents == 0):
# text to video
output = self._process_cross_attn(x, cond, num_latent_frames, x_ref_attn_map)
return None, output
elif num_cond_latents is not None and num_cond_latents > 0:
# image to video or video continuation
num_cond_latents_thw = num_cond_latents * (N // num_latent_frames)
x_noise = x[:, num_cond_latents_thw:]
cond = rearrange(cond, "(B N_t) M C -> B N_t M C", B=B)
cond = cond[:, num_cond_latents:]
cond = rearrange(cond, "B N_t M C -> (B N_t) M C")
frames_num = num_latent_frames - num_cond_latents
if human_num is not None and human_num == 2:
# multitalk mode
output_noise = self._process_cross_attn(x_noise, cond, frames_num, x_ref_attn_map)
else:
# singletalk mode
output_noise = self._process_cross_attn(x_noise, cond, frames_num)
output_cond = torch.zeros((B, num_cond_latents_thw, C), dtype=output_noise.dtype, device=output_noise.device)
return output_cond, output_noise
else:
raise NotImplementedError
+120
View File
@@ -0,0 +1,120 @@
import torch
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"),
],
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) -> 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
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 = 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)
NODE_CLASS_MAPPINGS = {
"WanVideoLongCatAvatarExtendEmbeds": WanVideoLongCatAvatarExtendEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoLongCatAvatarExtendEmbeds": "WanVideo LongCat Avatar Extend Embeds",
}
+199
View File
@@ -0,0 +1,199 @@
import torch
import torch.nn as nn
from einops import rearrange
from ..wanvideo.modules.attention import attention
def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor):
return (x * (1 + scale) + shift)
def sinusoidal_embedding_1d(dim, position):
sinusoid = torch.outer(position.type(torch.float64), torch.pow(
10000, -torch.arange(dim//2, dtype=torch.float64, device=position.device).div(dim//2)))
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
return x.to(position.dtype)
def precompute_freqs_cis_3d(dim: int, end: int = 1024, theta: float = 10000.0):
# 3d rope precompute
f_freqs_cis = precompute_freqs_cis(dim - 2 * (dim // 3), end, theta)
h_freqs_cis = precompute_freqs_cis(dim // 3, end, theta)
w_freqs_cis = precompute_freqs_cis(dim // 3, end, theta)
return f_freqs_cis, h_freqs_cis, w_freqs_cis
def precompute_freqs_cis(dim: int, end: int = 1024, theta: float = 10000.0):
# 1d rope precompute
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)
[: (dim // 2)].double() / dim))
freqs = torch.outer(torch.arange(end, device=freqs.device), freqs)
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64
return freqs_cis
def rope_apply(x, freqs, num_heads):
x = rearrange(x, "b s (n d) -> b s n d", n=num_heads)
x_out = torch.view_as_complex(x.to(torch.float64).reshape(
x.shape[0], x.shape[1], x.shape[2], -1, 2))
x_out = torch.view_as_real(x_out * freqs).flatten(2)
return x_out.to(x.dtype)
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-5):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def norm(self, x):
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
def forward(self, x):
dtype = x.dtype
return self.norm(x.float()).to(dtype) * self.weight
class AttentionModule(nn.Module):
def __init__(self, num_heads, head_dim):
super().__init__()
self.num_heads = num_heads
self.head_dim = head_dim
def forward(self, q, k, v):
b, n, d = q.size(0), self.num_heads, self.head_dim
x = attention(
q.view(b, -1, n, d),
k.view(b, -1, n, d),
v.view(b, -1, n, d)
)
return x.flatten(2)
class SelfAttention(nn.Module):
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6):
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.q = nn.Linear(dim, dim)
self.k = nn.Linear(dim, dim)
self.v = nn.Linear(dim, dim)
self.o = nn.Linear(dim, dim)
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
self.attn = AttentionModule(self.num_heads, self.head_dim)
def forward(self, x, freqs):
q = self.norm_q(self.q(x))
k = self.norm_k(self.k(x))
v = self.v(x)
q = rope_apply(q, freqs, self.num_heads)
k = rope_apply(k, freqs, self.num_heads)
x = self.attn(q, k, v)
return self.o(x)
class CrossAttention(nn.Module):
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6, clip_fea: torch.Tensor = None):
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.q = nn.Linear(dim, dim)
self.k = nn.Linear(dim, dim)
self.v = nn.Linear(dim, dim)
self.o = nn.Linear(dim, dim)
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
self.k_img = nn.Linear(dim, dim)
self.v_img = nn.Linear(dim, dim)
self.norm_k_img = RMSNorm(dim, eps=eps)
self.attn = AttentionModule(self.num_heads, self.head_dim)
def forward(self, x: torch.Tensor, y: torch.Tensor, clip_fea: torch.Tensor = None):
ctx = y
q = self.norm_q(self.q(x))
k = self.norm_k(self.k(ctx))
v = self.v(ctx)
x = self.attn(q, k, v)
if clip_fea is not None:
k_img = self.norm_k_img(self.k_img(clip_fea))
v_img = self.v_img(clip_fea)
y = self.attn(q, k_img, v_img)
x = x + y
return self.o(x)
class GateModule(nn.Module):
def __init__(self,):
super().__init__()
def forward(self, x, gate, residual):
return x + gate * residual
class DiTBlock(nn.Module):
def __init__(self, dim: int, num_heads: int, ffn_dim: int, eps: float = 1e-6):
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.ffn_dim = ffn_dim
self.self_attn = SelfAttention(dim, num_heads, eps)
self.cross_attn = CrossAttention(dim, num_heads, eps)
self.norm1 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
self.norm2 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
self.norm3 = nn.LayerNorm(dim, eps=eps)
self.ffn = nn.Sequential(nn.Linear(dim, ffn_dim), nn.GELU(
approximate='tanh'), nn.Linear(ffn_dim, dim))
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
self.gate = GateModule()
def forward(self, x, context, t_mod, freqs, clip_fea=None):
has_seq = len(t_mod.shape) == 4
chunk_dim = 2 if has_seq else 1
# msa: multi-head self-attention mlp: multi-layer perceptron
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
self.modulation.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod).chunk(6, dim=chunk_dim)
if has_seq:
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
shift_msa.squeeze(2), scale_msa.squeeze(2), gate_msa.squeeze(2),
shift_mlp.squeeze(2), scale_mlp.squeeze(2), gate_mlp.squeeze(2),
)
input_x = modulate(self.norm1(x), shift_msa, scale_msa)
x = self.gate(x, gate_msa, self.self_attn(input_x, freqs))
x = x + self.cross_attn(self.norm3(x), context, clip_fea=clip_fea)
input_x = modulate(self.norm2(x), shift_mlp, scale_mlp)
x = self.gate(x, gate_mlp, self.ffn(input_x))
return x
class WanModelDualControl(torch.nn.Module):
def __init__(self, dim: int, ffn_dim: int, eps: float, num_heads: int, control_layers = 12):
super().__init__()
self.control_layers = control_layers
self.control_blocks_dense = nn.ModuleList([
DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps)
for _ in range(self.control_layers)
])
self.control_blocks_sparse = nn.ModuleList([
DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps)
for _ in range(self.control_layers)
])
self.control_initial_combine_linear_dense = torch.nn.Linear(dim, dim//2)
self.control_initial_combine_linear_sparse = torch.nn.Linear(dim, dim//2)
self.control_text_linear = torch.nn.Linear(dim, dim//2)
self.control_t_mod = torch.nn.Linear(dim, dim//2)
self.control_combine_linears = torch.nn.ModuleList([torch.nn.Linear(dim//2, dim) for _ in range(self.control_layers)])
head_dim = dim // num_heads
self.freqs = precompute_freqs_cis_3d(head_dim)
+88
View File
@@ -0,0 +1,88 @@
import torch
from ..utils import log
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoAddDualControlEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"vae": ("WANVAE", {"tooltip": "VAE model"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
"first_frame_noise_level": ("FLOAT", {"default": 0.925926, "min": 0.0, "max": 1.0, "step": 0.000001, "tooltip": "Noise level for the first frame when using previous frames"}),
},
"optional": {
"dense": ("IMAGE", {"tooltip": "Dense control signal (depth) video input"}),
"sparse": ("IMAGE", {"tooltip": "Sparse control signal (tracks) video input"}),
"prev_images": ("IMAGE", {"tooltip": "Previous frames for temporal consistency, default is 8 frames"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, vae, strength, start_percent, end_percent, first_frame_noise_level, dense=None, sparse=None, prev_images=None):
updated = dict(embeds)
updated.setdefault("dual_control", {})
if dense is None and sparse is None:
raise ValueError("At least one of dense or sparse inputs must be provided.")
num_frames = dense.shape[0] if dense is not None else sparse.shape[0]
height = dense.shape[1] if dense is not None else sparse.shape[1]
width = dense.shape[2] if dense is not None else sparse.shape[2]
msk = torch.ones(1, num_frames, height//8, width//8, device=device)
msk[:, 1:] = 0
msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
msk = msk.view(1, msk.shape[1] // 4, 4, height//8, width//8)
msk = msk.transpose(1, 2)
dense_input_latent = sparse_input_latent = None
vae.to(device)
if dense is not None:
dense_images = 1 - dense[..., :3] # Invert colors for depth to match the usual range in comfy
dense_images = dense_images.permute(3, 0, 1, 2) * 2 - 1
dense_video_latent = vae.encode([dense_images.to(device, vae.dtype)], device, tiled=False)
dense_first = (dense_images[:, :1]).to(device, vae.dtype)
vae_input_dense = torch.cat([dense_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1)
dense_concat_latent = vae.encode([vae_input_dense], device, tiled=False)
dense_concat_latent = torch.cat([msk, dense_concat_latent], dim=1)
dense_input_latent = torch.cat([dense_video_latent, dense_concat_latent],dim=1)
if sparse is not None:
sparse_images = sparse[..., :3].permute(3, 0, 1, 2) * 2 - 1
sparse_video_latent = vae.encode([sparse_images.to(device, vae.dtype)], device, tiled=False)
sparse_first = (sparse_images[:, :1]).to(device, vae.dtype)
vae_input_sparse = torch.cat([sparse_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1)
sparse_concat_latent = vae.encode([vae_input_sparse], device, tiled=False)
sparse_concat_latent = torch.cat([msk, sparse_concat_latent], dim=1)
sparse_input_latent = torch.cat([sparse_video_latent, sparse_concat_latent],dim=1)
if prev_images is not None:
prev_images = prev_images[..., :3].permute(3, 0, 1, 2) * 2 - 1
prev_video_latent = vae.encode([prev_images.to(device, vae.dtype)], device, tiled=False)
updated["dual_control"]["prev_latent"] = prev_video_latent[0]
vae.to(offload_device)
updated["dual_control"]["dense_input_latent"] = dense_input_latent
updated["dual_control"]["sparse_input_latent"] = sparse_input_latent
updated["dual_control"]["strength"] = strength
updated["dual_control"]["start_percent"] = start_percent
updated["dual_control"]["end_percent"] = end_percent
updated["dual_control"]["first_frame_noise_level"] = first_frame_noise_level
return (updated,)
NODE_CLASS_MAPPINGS = {
"WanVideoAddDualControlEmbeds": WanVideoAddDualControlEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddDualControlEmbeds": "WanVideo Add Dual Control Embeds",
}
+83 -12
View File
@@ -31,7 +31,7 @@ def p3d_to_p2d(point_3d, height, width): # point3d n*1024*3
def get_pose_images(smpl_data, offset):
pose_images = []
for data in smpl_data:
for data in smpl_data:
if isinstance(data, np.ndarray):
joints3d = data
else:
@@ -43,28 +43,33 @@ def get_pose_images(smpl_data, offset):
return pose_images
def get_control_conditions(poses, h, w):
video_transforms = transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True)
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)
canvas = draw_3d_points(
canvas,
joints3d[0],
stickwidth=int(h / 350),
)
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 as e:
print("wrong:", e)
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.
print("control_pixel_values.shape", control_pixel_values.shape)
#control_pixel_values = video_transforms(control_pixel_values)
return control_pixel_values
@@ -140,3 +145,69 @@ def draw_3d_points(canvas, points, stickwidth=2, r=2, draw_line=True):
cv2.fillConvexPoly(canvas, polygon, connection_colors[i%17])
return canvas
def draw_3d_points_scail(canvas, points, stickwidth=2, r=2, draw_line=True):
connetions = [
[15,12],[12, 16],[16, 18],[18, 20],[20, 22], # 0-4: Left arm chain
[12,17],[17,19],[19,21], # 5-7: Right arm chain
[21,23], # 8: Right hand
[12,1],[1,4],[4,7], # 9-11: Neck to left leg (hip, thigh, shin)
[12,2],[2,5],[5,8], # 12-14: Neck to right leg (hip, thigh, shin)
]
# Warm colors for right side, cool colors for left side
connection_colors = [
[180, 180, 180], # 0: [15,12] - L. clavicle (Bright Cyan)
[0, 200, 255], # 1: [12,16] - L. shoulder (Bright Cyan)
[0, 120, 255], # 2: [16,18] - L. upper arm (Bright Blue)
[0, 60, 255], # 3: [18,20] - L. forearm (Deep Blue)
[60, 0, 255], # 4: [20,22] - L. hand (Blue-Purple)
[255, 0, 0], # 5: [12,17] - R. clavicle (Bright Red)
[255, 100, 0], # 6: [17,19] - R. upper arm (Bright Orange)
[255, 180, 0], # 7: [19,21] - R. forearm (Golden Orange)
[255, 255, 0], # 8: [21,23] - R. hand (Bright Yellow)
[30, 27, 160], # 9: [12,1] - Neck to L. hip (purple-blue)
[73, 27, 177], # 10: [1,4] - L. thigh (purple)
[145, 27, 194], # 11: [4,7] - L. shin (magenta)
[200, 255, 100], # 12: [12,2] - Neck to R. hip (yellow)
[54, 201, 52], # 13: [2,5] - R. thigh (green)
[30, 176, 85], # 14: [5,8] - R. shin (green)
]
# draw line
if draw_line:
# Collect all joints that are part of connections
joints_in_use = set()
for connection in connetions:
joints_in_use.add(connection[0])
joints_in_use.add(connection[1])
for i in range(len(connetions)):
point1_idx, point2_idx = connetions[i][0:2]
point1 = points[point1_idx]
point2 = points[point2_idx]
x1, y1 = int(point1[0]), int(point1[1])
x2, y2 = int(point2[0]), int(point2[1])
cv2.line(canvas, (x1, y1), (x2, y2), connection_colors[i], stickwidth)
# draw points for joints that have connections
joints_in_use = set()
for connection in connetions:
joints_in_use.add(connection[0])
joints_in_use.add(connection[1])
for joint_idx in joints_in_use:
if joint_idx >= len(points):
continue
x, y = points[joint_idx][0:2]
x, y = int(x), int(y)
# Use the color from the first connection involving this joint
joint_color = [180, 180, 180] # default grey
for i, connection in enumerate(connetions):
if connection[0] == joint_idx or connection[1] == joint_idx:
joint_color = connection_colors[i]
break
cv2.circle(canvas, (x, y), r, joint_color, thickness=-1)
return canvas
View File
+146 -48
View File
@@ -1,34 +1,56 @@
import os
import torch
import gc
from ..utils import log, dict_to_device
from ..utils import log
import numpy as np
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
import comfy.model_management as mm
from comfy.utils import load_torch_file
import folder_paths
script_directory = os.path.dirname(os.path.abspath(__file__))
script_directory = os.path.dirname(os.path.abspath(__file__))
device = mm.get_torch_device()
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
from .mtv import prepare_motion_embeddings
def check_jit_script_function():
if torch.jit.script.__name__ != "script":
# Get more details about what modified it
module = torch.jit.script.__module__
qualname = getattr(torch.jit.script, '__qualname__', 'unknown')
code_file = None
try:
code_file = torch.jit.script.__code__.co_filename
code_line = torch.jit.script.__code__.co_firstlineno
log.warning(f"torch.jit.script has been modified by another custom node.\n"
f" Function name: {torch.jit.script.__name__}\n"
f" Module: {module}\n"
f" Qualified name: {qualname}\n"
f" Defined in: {code_file}:{code_line}\n"
f"This may cause issues with the NLF model.")
except:
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": (
[
"https://github.com/isarandi/nlf/releases/download/v0.3.2/nlf_l_multi_0.3.2.torchscript"
],
)
"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"}),
},
}
@@ -37,8 +59,11 @@ class DownloadAndLoadNLFModel:
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, url):
def loadmodel(self, url, warmup=True):
if url not in model_list:
raise ValueError(f"URL {url} is not in the list of allowed models.")
check_jit_script_function()
if not os.path.exists(local_model_path):
log.info(f"Downloading NLF model to: {local_model_path}")
import requests
@@ -52,6 +77,20 @@ class DownloadAndLoadNLFModel:
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:
@@ -59,8 +98,12 @@ class LoadNLFModel:
def INPUT_TYPES(s):
return {
"required": {
"path": ("STRING", {"default": local_model_path}),
"nlf_model": (folder_paths.get_filename_list("nlf"), {"tooltip": "These models are loaded from the 'ComfyUI/models/nlf' -folder",}),
},
"optional": {
"warmup": ("BOOLEAN", {"default": True, "tooltip": "Whether to warmup the model after loading"}),
},
}
RETURN_TYPES = ("NLFMODEL",)
@@ -68,8 +111,22 @@ class LoadNLFModel:
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, path):
model = torch.jit.load(path).eval()
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,
@@ -108,7 +165,7 @@ class LoadVQVAE:
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)
@@ -131,15 +188,6 @@ class MTVCrafterEncodePoses:
def encode(self, vqvae, poses):
# import pickle
# with open(os.path.join(script_directory, "data", "sampled_data.pkl"), 'rb') as f:
# data_list = pickle.load(f)
# if not isinstance(data_list, list):
# data_list = [data_list]
# print(data_list)
# smpl_poses = data_list[1]['pose']
global_mean = np.load(os.path.join(script_directory, "data", "mean.npy")) #global_mean.shape: (24, 3)
global_std = np.load(os.path.join(script_directory, "data", "std.npy"))
@@ -153,7 +201,7 @@ class MTVCrafterEncodePoses:
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)
@@ -162,7 +210,7 @@ class MTVCrafterEncodePoses:
'global_mean': global_mean,
'global_std': global_std
}
return poses_dict, recon_motion
@@ -173,32 +221,74 @@ class NLFPredict:
"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", )
RETURN_NAMES = ("pose_results",)
RETURN_TYPES = ("NLFPRED", "BBOX",)
RETURN_NAMES = ("pose_results", "bboxes")
FUNCTION = "predict"
CATEGORY = "WanVideoWrapper"
def predict(self, model, images):
model.to(device)
pred = model.detect_smpl_batched(images.permute(0, 3, 1, 2).to(device))
model.to(offload_device)
def predict(self, model, images, per_batch=-1):
pred = dict_to_device(pred, offload_device)
check_jit_script_function()
model = model.to(device)
num_images = images.shape[0]
# Determine batch size
if per_batch == -1:
batch_size = num_images
else:
batch_size = per_batch
# Initialize result containers
all_boxes = []
all_joints3d_nonparam = []
# Process in batches
for i in range(0, num_images, batch_size):
end_idx = min(i + batch_size, num_images)
batch_images = images[i:end_idx]
jit_profiling_prev_state = torch._C._jit_set_profiling_executor(True)
try:
pred = model.detect_smpl_batched(batch_images.permute(0, 3, 1, 2).to(device))
finally:
torch._C._jit_set_profiling_executor(jit_profiling_prev_state)
# Collect boxes and joints from this batch
if 'boxes' in pred:
all_boxes.extend(pred['boxes'])
if 'joints3d_nonparam' in pred:
all_joints3d_nonparam.extend(pred['joints3d_nonparam'])
model = model.to(offload_device)
# Move collected results to offload device
all_boxes = [box.to(offload_device) for box in all_boxes]
all_joints3d_nonparam = [joints.to(offload_device) for joints in all_joints3d_nonparam]
# Maintain the original nested format: wrap in a list to match expected structure
pose_results = {
'joints3d_nonparam': [],
'joints3d_nonparam': [all_joints3d_nonparam],
}
# Collect pose data
for key in pose_results.keys():
if key in pred:
pose_results[key].append(pred[key])
# 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:
pose_results[key].append(None)
return (pose_results,)
# 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
@@ -208,25 +298,32 @@ class DrawNLFPoses:
"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):
def predict(self, poses, width, height, stick_width=1.0, point_radius=2, style="original"):
from .draw_pose import get_control_conditions
print(type(poses))
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)
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,
@@ -234,6 +331,7 @@ NODE_CLASS_MAPPINGS = {
"MTVCrafterEncodePoses": MTVCrafterEncodePoses
}
NODE_DISPLAY_NAME_MAPPINGS = {
"LoadNLFModel": "Load NLF Model",
"DownloadAndLoadNLFModel": "(Download)Load NLF Model",
"NLFPredict": "NLF Predict",
"DrawNLFPoses": "Draw NLF Poses",
+13 -3
View File
@@ -8,6 +8,8 @@ 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()
@@ -216,9 +218,11 @@ class WanVideoOviCFG:
def INPUT_TYPES(s):
return {"required": {
"original_text_embeds": ("WANVIDEOTEXTEMBEDS",),
"ovi_negative_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", )
@@ -227,10 +231,16 @@ class WanVideoOviCFG:
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_negative_text_embeds, ovi_audio_cfg):
negative_text_embeds = ovi_negative_text_embeds.get("negative_prompt_embeds", None)
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({
+13 -6
View File
@@ -75,14 +75,21 @@ class VAE(nn.Module):
super().__init__()
if data_dim == 80:
self.data_mean = nn.Buffer(torch.tensor(DATA_MEAN_80D, dtype=torch.float32))
self.data_std = nn.Buffer(torch.tensor(DATA_STD_80D, dtype=torch.float32))
data_mean = torch.tensor(DATA_MEAN_80D, dtype=torch.float32)
data_std = torch.tensor(DATA_STD_80D, dtype=torch.float32)
elif data_dim == 128:
self.data_mean = nn.Buffer(torch.tensor(DATA_MEAN_128D, dtype=torch.float32))
self.data_std = nn.Buffer(torch.tensor(DATA_STD_128D, dtype=torch.float32))
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")
self.data_mean = self.data_mean.view(1, -1, 1)
self.data_std = self.data_std.view(1, -1, 1)
# 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,
+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
+60 -113
View File
@@ -1,128 +1,75 @@
try:
from .utils import check_duplicate_nodes, log
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" - {dir_path}\n"
log.warning(warning_msg + "Please remove duplicates to avoid possible conflicts.")
warning_msg += f" - {color_text(dir_path, 'yellow')}\n"
log.warning(color_text(warning_msg + "Please remove duplicates to avoid possible conflicts.", "red"))
except:
pass
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 .skyreels.nodes import NODE_CLASS_MAPPINGS as SKYREELS_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as SKYREELS_NODE_DISPLAY_NAME_MAPPINGS
from .fantasytalking.nodes import NODE_CLASS_MAPPINGS as FANTASYTALKING_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FANTASYTALKING_NODE_DISPLAY_NAME_MAPPINGS
from .nodes_sampler import NODE_CLASS_MAPPINGS as SAMPLER_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as SAMPLER_NODE_DISPLAY_NAME_MAPPINGS
from .fun_camera.nodes import NODE_CLASS_MAPPINGS as FUN_CAMERA_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FUN_CAMERA_NODE_DISPLAY_NAME_MAPPINGS
from .uni3c.nodes import NODE_CLASS_MAPPINGS as UNI3C_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNI3C_NODE_DISPLAY_NAME_MAPPINGS
from .controlnet.nodes import NODE_CLASS_MAPPINGS as CONTROLNET_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as CONTROLNET_NODE_DISPLAY_NAME_MAPPINGS
from .ATI.nodes import NODE_CLASS_MAPPINGS as ATI_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as ATI_NODE_DISPLAY_NAME_MAPPINGS
from .multitalk.nodes import NODE_CLASS_MAPPINGS as MULTITALK_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MULTITALK_NODE_DISPLAY_NAME_MAPPINGS
from .nodes_model_loading import NODE_CLASS_MAPPINGS as MODEL_LOADING_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MODEL_LOADING_NODE_DISPLAY_NAME_MAPPINGS
from .nodes_utility import NODE_CLASS_MAPPINGS as UTILITY_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UTILITY_NODE_DISPLAY_NAME_MAPPINGS
from .cache_methods.nodes_cache import NODE_CLASS_MAPPINGS as NODE_CACHE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as NODE_CACHE_DISPLAY_NAME_MAPPINGS
from .nodes_deprecated import NODE_CLASS_MAPPINGS as DEPRECATED_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS
from .s2v.nodes import NODE_CLASS_MAPPINGS as S2V_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as S2V_NODE_DISPLAY_NAME_MAPPINGS
from .FlashVSR.flashvsr_nodes import NODE_CLASS_MAPPINGS as FLASHVSR_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FLASHVSR_NODE_DISPLAY_NAME_MAPPINGS
from .mocha.nodes import NODE_CLASS_MAPPINGS as MOCHA_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MOCHA_NODE_DISPLAY_NAME_MAPPINGS
from .utils import log
try:
from .qwen.qwen import NODE_CLASS_MAPPINGS as QWEN_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as QWEN_NODE_DISPLAY_NAME_MAPPINGS
except Exception as e:
log.warning(f"WanVideoWrapper WARNING: Qwen nodes not available due to error in importing them: {e}")
QWEN_NODE_CLASS_MAPPINGS = {}
QWEN_NODE_DISPLAY_NAME_MAPPINGS = {}
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
# Required modules (will raise on import failure)
REQUIRED_MODULES = [
(".nodes", "Main"),
(".nodes_sampler", "Sampler"),
(".nodes_model_loading", "ModelLoading"),
(".nodes_utility", "Utility"),
(".cache_methods.nodes_cache", "Cache"),
]
try:
from .fantasyportrait.nodes import NODE_CLASS_MAPPINGS as FANTASYPORTRAIT_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FANTASYPORTRAIT_NODE_DISPLAY_NAME_MAPPINGS
except Exception as e:
log.warning(f"WanVideoWrapper WARNING: FantasyPortrait nodes not available due to error in importing them: {e}")
FANTASYPORTRAIT_NODE_CLASS_MAPPINGS = {}
FANTASYPORTRAIT_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"),
]
try:
from .unianimate.nodes import NODE_CLASS_MAPPINGS as UNIANIMATE_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS
except Exception as e:
log.warning(f"WanVideoWrapper WARNING: UniAnimate nodes not available due to error in importing them: {e}")
UNIANIMATE_NODE_CLASS_MAPPINGS = {}
UNIANIMATE_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
try:
from .MTV.nodes import NODE_CLASS_MAPPINGS as MTV_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MTV_NODE_DISPLAY_NAME_MAPPINGS
except Exception as e:
log.warning(f"WanVideoWrapper WARNING: MTV nodes not available due to error in importing them: {e}")
MTV_NODE_CLASS_MAPPINGS = {}
MTV_NODE_DISPLAY_NAME_MAPPINGS = {}
# Register all node modules
for module_path, name in REQUIRED_MODULES:
register_nodes(module_path, name, optional=False)
try:
from .HuMo.nodes import NODE_CLASS_MAPPINGS as HUMO_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as HUMO_NODE_DISPLAY_NAME_MAPPINGS
except Exception as e:
log.warning(f"WanVideoWrapper WARNING: HuMo nodes not available due to error in importing them: {e}")
HUMO_NODE_CLASS_MAPPINGS = {}
HUMO_NODE_DISPLAY_NAME_MAPPINGS = {}
for module_path, name in OPTIONAL_MODULES:
register_nodes(module_path, name, optional=True)
try:
from .lynx.nodes import NODE_CLASS_MAPPINGS as LYNX_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as LYNX_NODE_DISPLAY_NAME_MAPPINGS
except Exception as e:
log.warning(f"WanVideoWrapper WARNING: Lynx nodes not available due to error in importing them: {e}")
LYNX_NODE_CLASS_MAPPINGS = {}
LYNX_NODE_DISPLAY_NAME_MAPPINGS = {}
try:
from .Ovi.nodes_ovi import NODE_CLASS_MAPPINGS as OVI_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as OVI_NODE_DISPLAY_NAME_MAPPINGS
except Exception as e:
log.warning(f"WanVideoWrapper WARNING: Ovi nodes not available due to error in importing them: {e}")
OVI_NODE_CLASS_MAPPINGS = {}
OVI_NODE_DISPLAY_NAME_MAPPINGS = {}
NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(FANTASYTALKING_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(FANTASYPORTRAIT_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(FUN_CAMERA_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(UNI3C_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(CONTROLNET_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(ATI_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(MULTITALK_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(MODEL_LOADING_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(UTILITY_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(NODE_CACHE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(DEPRECATED_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(QWEN_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(MTV_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(S2V_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(HUMO_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(SAMPLER_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(LYNX_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(OVI_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(FLASHVSR_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(MOCHA_NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(SKYREELS_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(FANTASYTALKING_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(FANTASYPORTRAIT_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(FUN_CAMERA_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(UNI3C_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(CONTROLNET_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(ATI_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(MULTITALK_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(MODEL_LOADING_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(UTILITY_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(NODE_CACHE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(QWEN_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(MTV_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(S2V_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(HUMO_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(SAMPLER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(LYNX_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(OVI_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(FLASHVSR_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(MOCHA_NODE_DISPLAY_NAME_MAPPINGS)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+1 -13
View File
@@ -28,19 +28,7 @@ aggressive values this can happen and the motion suffers. Starting later can hel
When NOT using coefficients, the threshold value should be
about 10 times smaller than the value used with coefficients.
Official recommended values https://github.com/ali-vilab/TeaCache/tree/main/TeaCache4Wan2.1:
<pre style='font-family:monospace'>
+-------------------+--------+---------+--------+
| Model | Low | Medium | High |
+-------------------+--------+---------+--------+
| Wan2.1 t2v 1.3B | 0.05 | 0.07 | 0.08 |
| Wan2.1 t2v 14B | 0.14 | 0.15 | 0.20 |
| Wan2.1 i2v 480P | 0.13 | 0.19 | 0.26 |
| Wan2.1 i2v 720P | 0.18 | 0.20 | 0.30 |
+-------------------+--------+---------+--------+
</pre>
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"):
+12 -57
View File
@@ -10,18 +10,13 @@ from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscal
from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.transformers.transformer_wan import (
WanTimeTextImageEmbedding,
WanRotaryPosEmbed,
WanTimeTextImageEmbedding,
WanRotaryPosEmbed,
WanTransformerBlock
)
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
def zero_module(module):
for p in module.parameters():
nn.init.zeros_(p)
return module
class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin):
r"""
@@ -69,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,
@@ -100,10 +95,10 @@ 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"),
@@ -122,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)
@@ -153,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,
@@ -187,7 +181,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
# 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
@@ -216,7 +210,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
# 4. Transformer blocks
controlnet_hidden_states = ()
if torch.is_grad_enabled() and self.gradient_checkpointing:
@@ -239,43 +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,
"vae_channels": 16
}
controlnet = WanControlnet(**parameters)
hidden_states = torch.rand(1, 16, 13, 60, 90)
timestep = torch.tensor([1000]).repeat(17550).unsqueeze(0) #torch.randint(low=0, high=1000, size=(1,), dtype=torch.long)
encoder_hidden_states = torch.rand(1, 512, 4096)
controlnet_states = torch.rand(1, 3, 49, 480, 720)
controlnet_hidden_states = controlnet(
hidden_states=hidden_states,
timestep=timestep,
encoder_hidden_states=encoder_hidden_states,
controlnet_states=controlnet_states,
return_dict=False
)
print("Output states count", len(controlnet_hidden_states[0]))
for out_hidden_states in controlnet_hidden_states[0]:
print(out_hidden_states.shape)
+153 -38
View File
@@ -1,14 +1,52 @@
import torch
import torch.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):
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():
@@ -16,13 +54,22 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
allow_compile = compile_args.get("allow_unmerged_lora_compile", False)
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)
_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:
in_features = state_dict[module_prefix + "weight"].shape[1]
out_features = state_dict[module_prefix + "weight"].shape[0]
if scale_weights is not None:
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(
@@ -30,8 +77,9 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
out_features,
module.bias is not None,
compute_dtype=compute_dtype,
scale_weight=scale_weights.get(scale_key) if scale_weights else None,
allow_compile=allow_compile
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)
@@ -71,8 +119,8 @@ def set_lora_params(module, patches, module_prefix="", device=torch.device("cpu"
continue
lora_strengths = [p[0] for p in patch]
module.set_lora_diffs(lora_diffs, device=device)
module.lora_strengths = lora_strengths
module.step = 0 # Initialize step for LoRA scheduling
module.set_lora_strengths(lora_strengths, device=device)
module._step.fill_(0) # Initialize step for LoRA scheduling
class CustomLinear(nn.Linear):
@@ -84,19 +132,56 @@ class CustomLinear(nn.Linear):
compute_dtype=None,
device=None,
scale_weight=None,
allow_compile=False
allow_compile=False,
is_gguf=False
) -> None:
super().__init__(in_features, out_features, bias, device)
self.compute_dtype = compute_dtype
self.lora_diffs = []
self.step = 0
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._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora)
self._apply_lora_impl = self._apply_lora_custom_op
self._apply_single_lora_impl = self._apply_single_lora_custom_op
self._linear_forward_impl = self._linear_forward_custom_op
else:
self._apply_lora_impl = self._apply_lora_direct
self._apply_single_lora_impl = self._apply_single_lora_direct
self._linear_forward_impl = self._linear_forward_direct
# Direct implementations (no custom ops)
def _apply_lora_direct(self, weight, lora_diff_0, lora_diff_1, lora_diff_2, lora_strength):
patch_diff = torch.mm(
lora_diff_0.flatten(start_dim=1),
lora_diff_1.flatten(start_dim=1)
).reshape(weight.shape) + 0
alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 != 0.0 else 1.0
scale = lora_strength * alpha
return weight + patch_diff * scale
def _apply_single_lora_direct(self, weight, lora_diff, lora_strength):
return weight + lora_diff * lora_strength
def _linear_forward_direct(self, input, weight, bias):
return torch.nn.functional.linear(input, weight, bias)
# Custom op implementations
def _apply_lora_custom_op(self, weight, lora_diff_0, lora_diff_1, lora_diff_2, lora_strength):
return torch.ops.wanvideo.apply_lora(weight, lora_diff_0, lora_diff_1,
float(lora_diff_2) if lora_diff_2 is not None else 0.0, lora_strength
)
def _apply_single_lora_custom_op(self, weight, lora_diff, lora_strength):
return torch.ops.wanvideo.apply_single_lora(weight, lora_diff, lora_strength)
def _linear_forward_custom_op(self, input, weight, bias):
return torch.ops.wanvideo.linear_forward(input, weight, bias)
def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")):
self.lora_diffs = []
for i, diff in enumerate(lora_diffs):
@@ -109,51 +194,81 @@ class CustomLinear(nn.Linear):
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
self.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 outside compiled region"""
"""Apply LoRA using custom ops to avoid graph breaks"""
if not hasattr(self, "lora_diff_0_0"):
return weight
for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths):
if isinstance(lora_strength, list):
lora_strength = lora_strength[self.step]
if lora_strength == 0.0:
continue
elif lora_strength == 0.0:
continue
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])
patch_diff = torch.mm(
lora_diff_0.flatten(start_dim=1),
lora_diff_1.flatten(start_dim=1)
).reshape(weight.shape) + 0
alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 is not None else 1.0
scale = lora_strength * alpha
weight = weight.add(patch_diff, alpha=scale)
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 = weight.add(lora_diff, alpha=lora_strength)
weight = self._apply_single_lora_impl(weight, lora_diff, lora_strength)
return weight
def _prepare_weight(self, input):
"""Prepare weight tensor - handles both regular and GGUF weights"""
if self.is_gguf:
weight = dequantize_gguf_tensor(self.weight).to(self.compute_dtype)
else:
weight = self.weight.to(input)
return weight
def forward(self, input):
weight = self._prepare_weight(input)
if self.bias is not None:
bias = self.bias.to(input)
bias = self.bias.to(input if not self.is_gguf else self.compute_dtype)
else:
bias = None
weight = self.weight.to(input)
if self.scale_weight is not None:
# Only apply scale_weight for non-GGUF models
if not self.is_gguf and self.scale_weight is not None:
if weight.numel() < input.numel():
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)
return torch.nn.functional.linear(input, weight, bias)
def remove_lora_from_module(module):
for name, submodule in module.named_modules():
if hasattr(submodule, "lora_diffs"):
File diff suppressed because it is too large Load Diff
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -1,780 +0,0 @@
{
"id": "8b7a9a57-2303-4ef5-9fc2-bf41713bd1fc",
"revision": 0,
"last_node_id": 46,
"last_link_id": 58,
"nodes": [
{
"id": 33,
"type": "Note",
"pos": [
227.3764190673828,
-205.28524780273438
],
"size": [
351.70458984375,
88
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {},
"widgets_values": [
"Models:\nhttps://huggingface.co/Kijai/WanVideo_comfy/tree/main"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 11,
"type": "LoadWanVideoT5TextEncoder",
"pos": [
224.15325927734375,
-34.481563568115234
],
"size": [
377.1661376953125,
130
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "wan_t5_model",
"type": "WANTEXTENCODER",
"slot_index": 0,
"links": [
15
]
}
],
"properties": {
"Node name for S&R": "LoadWanVideoT5TextEncoder"
},
"widgets_values": [
"umt5-xxl-enc-bf16.safetensors",
"bf16",
"offload_device",
"disabled"
]
},
{
"id": 28,
"type": "WanVideoDecode",
"pos": [
1692.973876953125,
-404.8614501953125
],
"size": [
315,
174
],
"flags": {},
"order": 12,
"mode": 0,
"inputs": [
{
"name": "vae",
"type": "WANVAE",
"link": 43
},
{
"name": "samples",
"type": "LATENT",
"link": 33
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"slot_index": 0,
"links": [
48
]
}
],
"properties": {
"Node name for S&R": "WanVideoDecode"
},
"widgets_values": [
true,
272,
272,
144,
128
]
},
{
"id": 38,
"type": "WanVideoVAELoader",
"pos": [
1687.4093017578125,
-582.2750854492188
],
"size": [
416.25482177734375,
82
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "vae",
"type": "WANVAE",
"slot_index": 0,
"links": [
43
]
}
],
"properties": {
"Node name for S&R": "WanVideoVAELoader"
},
"widgets_values": [
"wanvideo\\Wan2_1_VAE_bf16.safetensors",
"bf16"
]
},
{
"id": 42,
"type": "GetImageSizeAndCount",
"pos": [
1708.7301025390625,
-140.99705505371094
],
"size": [
277.20001220703125,
86
],
"flags": {},
"order": 13,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 48
}
],
"outputs": [
{
"name": "image",
"type": "IMAGE",
"slot_index": 0,
"links": [
56
]
},
{
"label": "832 width",
"name": "width",
"type": "INT",
"links": null
},
{
"label": "480 height",
"name": "height",
"type": "INT",
"links": null
},
{
"label": "257 count",
"name": "count",
"type": "INT",
"links": null
}
],
"properties": {
"Node name for S&R": "GetImageSizeAndCount"
},
"widgets_values": []
},
{
"id": 16,
"type": "WanVideoTextEncode",
"pos": [
675.8850708007812,
-36.032100677490234
],
"size": [
420.30511474609375,
261.5306701660156
],
"flags": {},
"order": 10,
"mode": 0,
"inputs": [
{
"name": "t5",
"type": "WANTEXTENCODER",
"link": 15
},
{
"name": "model_to_offload",
"shape": 7,
"type": "WANVIDEOMODEL",
"link": null
}
],
"outputs": [
{
"name": "text_embeds",
"type": "WANVIDEOTEXTEMBEDS",
"slot_index": 0,
"links": [
30
]
}
],
"properties": {
"Node name for S&R": "WanVideoTextEncode"
},
"widgets_values": [
"high quality nature video featuring a red panda balancing on a bamboo stem while a bird lands on it's head, on the background there is a waterfall",
"色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
true
]
},
{
"id": 30,
"type": "VHS_VideoCombine",
"pos": [
2127.120849609375,
-511.9014587402344
],
"size": [
873.2135620117188,
840.2385864257812
],
"flags": {},
"order": 14,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 56
},
{
"name": "audio",
"shape": 7,
"type": "AUDIO",
"link": null
},
{
"name": "meta_batch",
"shape": 7,
"type": "VHS_BatchManager",
"link": null
},
{
"name": "vae",
"shape": 7,
"type": "VAE",
"link": null
}
],
"outputs": [
{
"name": "Filenames",
"type": "VHS_FILENAMES",
"links": null
}
],
"properties": {
"Node name for S&R": "VHS_VideoCombine"
},
"widgets_values": {
"frame_rate": 16,
"loop_count": 0,
"filename_prefix": "WanVideo2_1_T2V",
"format": "video/h264-mp4",
"pix_fmt": "yuv420p",
"crf": 19,
"save_metadata": true,
"trim_to_audio": false,
"pingpong": false,
"save_output": true,
"videopreview": {
"hidden": false,
"paused": false,
"params": {
"filename": "WanVideo2_1_T2V_00412.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
"frame_rate": 16,
"workflow": "WanVideo2_1_T2V_00412.png",
"fullpath": "N:\\AI\\ComfyUI\\output\\WanVideo2_1_T2V_00412.mp4"
}
}
}
},
{
"id": 37,
"type": "WanVideoEmptyEmbeds",
"pos": [
1305.26708984375,
-571.7843627929688
],
"size": [
315,
106
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "image_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"links": [
42
]
}
],
"properties": {
"Node name for S&R": "WanVideoEmptyEmbeds"
},
"widgets_values": [
832,
480,
257
]
},
{
"id": 35,
"type": "WanVideoTorchCompileSettings",
"pos": [
193.47103881835938,
-614.6900024414062
],
"size": [
390.5999755859375,
178
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "torch_compile_args",
"type": "WANCOMPILEARGS",
"slot_index": 0,
"links": []
}
],
"properties": {
"Node name for S&R": "WanVideoTorchCompileSettings"
},
"widgets_values": [
"inductor",
false,
"default",
false,
64,
true
]
},
{
"id": 45,
"type": "WanVideoTeaCache",
"pos": [
931.4036865234375,
-792.5159912109375
],
"size": [
315,
154
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "cache_args",
"type": "CACHEARGS",
"links": [
58
]
}
],
"properties": {
"Node name for S&R": "WanVideoTeaCache"
},
"widgets_values": [
0.10000000000000002,
1,
-1,
"offload_device",
true
]
},
{
"id": 36,
"type": "Note",
"pos": [
796.0189208984375,
-521.5020751953125
],
"size": [
298.2554016113281,
108.62744140625
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {},
"widgets_values": [
"sdpa should work too, haven't tested flaash\n\nfp8_fast seems to cause huge quality degradation"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 46,
"type": "Note",
"pos": [
937.9556274414062,
-940.750244140625
],
"size": [
297.4364013671875,
88
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {},
"widgets_values": [
"TeaCache with context windows is VERY experimental and lower values than normal should be used."
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 27,
"type": "WanVideoSampler",
"pos": [
1315.2401123046875,
-401.48028564453125
],
"size": [
315,
574.1923217773438
],
"flags": {},
"order": 11,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "WANVIDEOMODEL",
"link": 29
},
{
"name": "text_embeds",
"type": "WANVIDEOTEXTEMBEDS",
"link": 30
},
{
"name": "image_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"link": 42
},
{
"name": "samples",
"shape": 7,
"type": "LATENT",
"link": null
},
{
"name": "feta_args",
"shape": 7,
"type": "FETAARGS",
"link": null
},
{
"name": "context_options",
"shape": 7,
"type": "WANVIDCONTEXT",
"link": 57
},
{
"name": "cache_args",
"shape": 7,
"type": "CACHEARGS",
"link": 58
},
{
"name": "flowedit_args",
"shape": 7,
"type": "FLOWEDITARGS",
"link": null
},
{
"name": "slg_args",
"shape": 7,
"type": "SLGARGS",
"link": null
},
{
"name": "loop_args",
"shape": 7,
"type": "LOOPARGS",
"link": null
}
],
"outputs": [
{
"name": "samples",
"type": "LATENT",
"slot_index": 0,
"links": [
33
]
}
],
"properties": {
"Node name for S&R": "WanVideoSampler"
},
"widgets_values": [
30,
6,
5,
1057359483639288,
"fixed",
true,
"unipc",
0,
1,
"",
"comfy"
]
},
{
"id": 43,
"type": "WanVideoContextOptions",
"pos": [
1307.9542236328125,
-855.8865356445312
],
"size": [
315,
226
],
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "vae",
"shape": 7,
"type": "WANVAE",
"link": null
}
],
"outputs": [
{
"name": "context_options",
"type": "WANVIDCONTEXT",
"slot_index": 0,
"links": [
57
]
}
],
"properties": {
"Node name for S&R": "WanVideoContextOptions"
},
"widgets_values": [
"uniform_standard",
81,
4,
16,
true,
false,
6,
2
]
},
{
"id": 22,
"type": "WanVideoModelLoader",
"pos": [
620.3950805664062,
-357.8426818847656
],
"size": [
477.4410095214844,
226.43276977539062
],
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "compile_args",
"shape": 7,
"type": "WANCOMPILEARGS",
"link": null
},
{
"name": "block_swap_args",
"shape": 7,
"type": "BLOCKSWAPARGS",
"link": null
},
{
"name": "lora",
"shape": 7,
"type": "WANVIDLORA",
"link": null
},
{
"name": "vram_management_args",
"shape": 7,
"type": "VRAM_MANAGEMENTARGS",
"link": null
}
],
"outputs": [
{
"name": "model",
"type": "WANVIDEOMODEL",
"slot_index": 0,
"links": [
29
]
}
],
"properties": {
"Node name for S&R": "WanVideoModelLoader"
},
"widgets_values": [
"WanVideo\\wan2.1_t2v_1.3B_fp16.safetensors",
"fp16",
"disabled",
"offload_device",
"sdpa"
]
}
],
"links": [
[
15,
11,
0,
16,
0,
"WANTEXTENCODER"
],
[
29,
22,
0,
27,
0,
"WANVIDEOMODEL"
],
[
30,
16,
0,
27,
1,
"WANVIDEOTEXTEMBEDS"
],
[
33,
27,
0,
28,
1,
"LATENT"
],
[
42,
37,
0,
27,
2,
"WANVIDIMAGE_EMBEDS"
],
[
43,
38,
0,
28,
0,
"VAE"
],
[
48,
28,
0,
42,
0,
"IMAGE"
],
[
56,
42,
0,
30,
0,
"IMAGE"
],
[
57,
43,
0,
27,
5,
"WANVIDCONTEXT"
],
[
58,
45,
0,
27,
6,
"TEACACHEARGS"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.8140274938684471,
"offset": [
-122.25834160503663,
993.5739491626379
]
},
"node_versions": {
"ComfyUI-WanVideoWrapper": "5a2383621a05825d0d0437781afcb8552d9590fd",
"ComfyUI-KJNodes": "a5bd3c86c8ed6b83c55c2d0e7a59515b15a0137f",
"ComfyUI-VideoHelperSuite": "0a75c7958fe320efcb052f1d9f8451fd20c730a8"
},
"VHS_latentpreview": true,
"VHS_latentpreviewrate": 0,
"VHS_MetadataImage": true,
"VHS_KeepIntermediate": true
},
"version": 0.4
}
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+3 -1
View File
@@ -220,6 +220,7 @@ class WanVideoAddFantasyPortrait:
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the portrait embedding"}),
"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"}),
"portrait_cfg": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 20.0, "step": 0.01, "tooltip": "CFG scale for the portrait embedding"}),
}
}
@@ -228,12 +229,13 @@ class WanVideoAddFantasyPortrait:
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, portrait_embeds, strength, start_percent=0.0, end_percent=1.0):
def add(self, embeds, portrait_embeds, strength, start_percent=0.0, end_percent=1.0, portrait_cfg=1.0):
new_entry = {
"adapter_proj": portrait_embeds,
"strength": strength,
"start_percent": start_percent,
"end_percent": end_percent,
"cfg_scale": portrait_cfg,
}
updated = dict(embeds)
+7 -143
View File
@@ -1,15 +1,11 @@
import torch
import torch.nn as nn
import numpy as np
import gguf
from accelerate import init_empty_weights
from .gguf_utils import GGUFParameter, dequantize_gguf_tensor
from ..utils import log
from .gguf_utils import GGUFParameter
def load_gguf(model_path):
from gguf import GGUFReader
reader = GGUFReader(model_path)
reader = gguf.GGUFReader(model_path)
parsed_parameters = {}
for tensor in reader.tensors:
# if the tensor is a torch supported dtype do not use GGUFParameter
@@ -18,144 +14,12 @@ def load_gguf(model_path):
parsed_parameters[tensor.name] = GGUFParameter(meta_tensor, quant_type=tensor.tensor_type) if is_gguf_quant else meta_tensor
return parsed_parameters, reader
#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py
from ..custom_linear import _replace_linear, set_lora_params, CustomLinear
def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modules_to_not_convert=[], patches=None, compile_args=None):
def _should_convert_to_gguf(state_dict, prefix):
weight_key = prefix + "weight"
return weight_key in state_dict and isinstance(state_dict[weight_key], GGUFParameter)
has_children = list(model.children())
if not has_children:
return
allow_compile = False
for name, module in model.named_children():
if compile_args is not None:
allow_compile = compile_args.get("allow_unmerged_lora_compile", False)
module_prefix = prefix + name + "."
_replace_with_gguf_linear(module, compute_dtype, state_dict, module_prefix, modules_to_not_convert, patches, compile_args)
if (
isinstance(module, nn.Linear)
and not isinstance(module, GGUFLinear)
and _should_convert_to_gguf(state_dict, module_prefix)
and name not in modules_to_not_convert
):
in_features = state_dict[module_prefix + "weight"].shape[1]
out_features = state_dict[module_prefix + "weight"].shape[0]
with init_empty_weights():
model._modules[name] = GGUFLinear(
in_features,
out_features,
module.bias is not None,
compute_dtype=compute_dtype,
allow_compile=allow_compile
)
model._modules[name].source_cls = type(module)
model._modules[name].requires_grad_(False)
return model
return _replace_linear(model, compute_dtype, state_dict, prefix, patches, None, compile_args, modules_to_not_convert)
def set_lora_params_gguf(module, patches, module_prefix="", device=torch.device("cpu")):
# Recursively set lora_diffs and lora_strengths for all GGUFLinear 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_gguf(child, patches, child_prefix, device)
if isinstance(module, GGUFLinear):
key = f"diffusion_model.{module_prefix}weight"
patch = patches.get(key, [])
#print(f"Processing LoRA patches for {key}: {len(patch)} patches found")
if len(patch) == 0:
key = key.replace("_orig_mod.", "")
patch = patches.get(key, [])
if len(patch) != 0:
lora_diffs = []
for p in patch:
lora_obj = p[1]
if "head" in key:
continue # For now skip LoRA for head layers
elif hasattr(lora_obj, "weights"):
lora_diffs.append(lora_obj.weights)
elif isinstance(lora_obj, tuple) and lora_obj[0] == "diff":
lora_diffs.append(lora_obj[1])
else:
continue
module.lora_strengths = [p[0] for p in patch]
module.set_lora_diffs(lora_diffs, device=device)
module.step = 0 # Initialize step for LoRA scheduling
return set_lora_params(module, patches, module_prefix, device)
class GGUFLinear(nn.Linear):
def __init__(
self,
in_features,
out_features,
bias=False,
compute_dtype=None,
device=None,
allow_compile=False
) -> None:
super().__init__(in_features, out_features, bias, device)
self.compute_dtype = compute_dtype
self.lora_diffs = []
self.lora_strengths = []
self.step = 0
self.allow_compile = allow_compile
if not allow_compile:
self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora)
def forward(self, inputs):
weight = dequantize_gguf_tensor(self.weight).to(self.compute_dtype)
bias = self.bias.to(self.compute_dtype) if self.bias is not None else None
weight = self._get_weight_with_lora(weight)#.to(self.compute_dtype)
return torch.nn.functional.linear(inputs, weight, bias)
def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")):
self.lora_diffs = []
for i, diff in enumerate(lora_diffs):
if len(diff) > 1:
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device, self.compute_dtype))
setattr(self, f"lora_diff_{i}_2", diff[2])
self.lora_diffs.append((f"lora_diff_{i}_0", f"lora_diff_{i}_1", f"lora_diff_{i}_2"))
else:
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
self.lora_diffs.append(f"lora_diff_{i}_0")
def _get_weight_with_lora(self, weight):
"""Apply LoRA outside compiled region"""
if not hasattr(self, "lora_diff_0_0"):
return weight
for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths):
if isinstance(lora_strength, list):
lora_strength = lora_strength[self.step]
if lora_strength == 0.0:
continue
elif lora_strength == 0.0:
continue
if isinstance(lora_diff_names, tuple):
lora_diff_0 = getattr(self, lora_diff_names[0])
lora_diff_1 = getattr(self, lora_diff_names[1])
lora_diff_2 = getattr(self, lora_diff_names[2])
patch_diff = torch.mm(
lora_diff_0.flatten(start_dim=1),
lora_diff_1.flatten(start_dim=1)
).reshape(weight.shape) + 0
alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 is not None else 1.0
scale = lora_strength * alpha
weight = weight.add(patch_diff, alpha=scale)
else:
lora_diff = getattr(self, lora_diff_names)
weight = weight.add(lora_diff, alpha=lora_strength)
return weight
GGUFLinear = CustomLinear
+483
View File
@@ -0,0 +1,483 @@
import torch
import os
import gc
from PIL import Image
import numpy as np
from ..latent_preview import prepare_callback
from ..wanvideo.schedulers import get_scheduler
from .multitalk import timestep_transform, add_noise
from ..utils import log, print_memory, temporal_score_rescaling, offload_transformer, init_blockswap
from comfy.utils import load_torch_file
from ..nodes_model_loading import load_weights
from ..HuMo.nodes import get_audio_emb_window
import comfy.model_management as mm
from tqdm import tqdm
import copy
VAE_STRIDE = (4, 8, 8)
PATCH_SIZE = (1, 2, 2)
vae_upscale_factor = 8
script_directory = os.path.dirname(os.path.abspath(__file__))
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
def multitalk_loop(self, **kwargs):
# Unpack kwargs into local variables
(latent, total_steps, steps, start_step, end_step, shift, cfg, denoise_strength,
sigmas, weight_dtype, transformer, patcher, block_swap_args, model, vae, dtype,
scheduler, scheduler_step_args, text_embeds, image_embeds, multitalk_embeds,
multitalk_audio_embeds, unianim_data, dwpose_data, unianimate_poses, uni3c_embeds,
humo_image_cond, humo_image_cond_neg, humo_audio, humo_reference_count,
add_noise_to_samples, audio_stride, use_tsr, tsr_k, tsr_sigma, fantasy_portrait_input,
noise, timesteps, force_offload, add_cond, control_latents, audio_proj,
control_camera_latents, samples, masks, seed_g, gguf_reader, predict_func
) = (kwargs.get(k) for k in (
'latent', 'total_steps', 'steps', 'start_step', 'end_step', 'shift', 'cfg',
'denoise_strength', 'sigmas', 'weight_dtype', 'transformer', 'patcher',
'block_swap_args', 'model', 'vae', 'dtype', 'scheduler', 'scheduler_step_args',
'text_embeds', 'image_embeds', 'multitalk_embeds', 'multitalk_audio_embeds',
'unianim_data', 'dwpose_data', 'unianimate_poses', 'uni3c_embeds',
'humo_image_cond', 'humo_image_cond_neg', 'humo_audio', 'humo_reference_count',
'add_noise_to_samples', 'audio_stride', 'use_tsr', 'tsr_k', 'tsr_sigma',
'fantasy_portrait_input', 'noise', 'timesteps', 'force_offload', 'add_cond',
'control_latents', 'audio_proj', 'control_camera_latents', 'samples', 'masks',
'seed_g', 'gguf_reader', 'predict_with_cfg'
))
mode = image_embeds.get("multitalk_mode", "multitalk")
if mode == "auto":
mode = transformer.multitalk_model_type.lower()
log.info(f"Multitalk mode: {mode}")
cond_frame = None
offload = image_embeds.get("force_offload", False)
offloaded = False
tiled_vae = image_embeds.get("tiled_vae", False)
frame_num = clip_length = image_embeds.get("frame_window_size", 81)
clip_embeds = image_embeds.get("clip_context", None)
if clip_embeds is not None:
clip_embeds = clip_embeds.to(dtype)
colormatch = image_embeds.get("colormatch", "disabled")
motion_frame = image_embeds.get("motion_frame", 25)
target_w = image_embeds.get("target_w", None)
target_h = image_embeds.get("target_h", None)
original_images = cond_image = image_embeds.get("multitalk_start_image", None)
if original_images is None:
original_images = torch.zeros([noise.shape[0], 1, target_h, target_w], device=device)
output_path = image_embeds.get("output_path", "")
img_counter = 0
if len(multitalk_embeds['audio_features'])==2 and (multitalk_embeds['ref_target_masks'] is None):
face_scale = 0.1
x_min, x_max = int(target_h * face_scale), int(target_h * (1 - face_scale))
lefty_min, lefty_max = int((target_w//2) * face_scale), int((target_w//2) * (1 - face_scale))
righty_min, righty_max = int((target_w//2) * face_scale + (target_w//2)), int((target_w//2) * (1 - face_scale) + (target_w//2))
human_mask1, human_mask2 = (torch.zeros([target_h, target_w]) for _ in range(2))
human_mask1[x_min:x_max, lefty_min:lefty_max] = 1
human_mask2[x_min:x_max, righty_min:righty_max] = 1
background_mask = torch.where((human_mask1 + human_mask2) > 0, torch.tensor(0), torch.tensor(1))
human_masks = [human_mask1, human_mask2, background_mask]
ref_target_masks = torch.stack(human_masks, dim=0)
multitalk_embeds['ref_target_masks'] = ref_target_masks
gen_video_list = []
is_first_clip = True
arrive_last_frame = False
cur_motion_frames_num = 1
audio_start_idx = iteration_count = step_iteration_count = 0
audio_end_idx = (audio_start_idx + clip_length) * audio_stride
indices = (torch.arange(4 + 1) - 2) * 1
current_condframe_index = 0
audio_embedding = multitalk_audio_embeds
human_num = len(audio_embedding)
audio_embs = None
cond_frame = None
uni3c_data = None
if uni3c_embeds is not None:
transformer.controlnet = uni3c_embeds["controlnet"]
uni3c_data = uni3c_embeds.copy()
encoded_silence = None
try:
silence_path = os.path.join(script_directory, "encoded_silence.safetensors")
encoded_silence = load_torch_file(silence_path)["audio_emb"].to(dtype)
except:
log.warning("No encoded silence file found, padding with end of audio embedding instead.")
total_frames = len(audio_embedding[0])
estimated_iterations = total_frames // (frame_num - motion_frame) + 1
callback = prepare_callback(patcher, estimated_iterations)
if frame_num >= total_frames:
arrive_last_frame = True
estimated_iterations = 1
log.info(f"Sampling {total_frames} frames in {estimated_iterations} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
while True: # start video generation iteratively
self.cache_state = [None, None]
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
if mode == "infinitetalk":
cond_image = original_images[:, :, current_condframe_index:current_condframe_index+1] if cond_image is not None else None
if multitalk_embeds is not None:
audio_embs = []
# split audio with window size
for human_idx in range(human_num):
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
center_indices = torch.clamp(center_indices, min=0, max=audio_embedding[human_idx].shape[0]-1)
audio_emb = audio_embedding[human_idx][center_indices].unsqueeze(0).to(device)
audio_embs.append(audio_emb)
audio_embs = torch.concat(audio_embs, dim=0).to(dtype)
h, w = (cond_image.shape[-2], cond_image.shape[-1]) if cond_image is not None else (target_h, target_w)
lat_h, lat_w = h // VAE_STRIDE[1], w // VAE_STRIDE[2]
latent_frame_num = (frame_num - 1) // 4 + 1
noise = torch.randn(
16, latent_frame_num,
lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
# Calculate the correct latent slice based on current iteration
if is_first_clip:
latent_start_idx = 0
latent_end_idx = noise.shape[1]
else:
new_frames_per_iteration = frame_num - motion_frame
new_latent_frames_per_iteration = ((new_frames_per_iteration - 1) // 4 + 1)
latent_start_idx = iteration_count * new_latent_frames_per_iteration
latent_end_idx = latent_start_idx + noise.shape[1]
if samples is not None:
noise_mask = samples.get("noise_mask", None)
input_samples = samples["samples"]
if input_samples is not None:
input_samples = input_samples.squeeze(0).to(noise)
# Check if we have enough frames in input_samples
if latent_end_idx > input_samples.shape[1]:
# We need more frames than available - pad the input_samples at the end
pad_length = latent_end_idx - input_samples.shape[1]
last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
input_samples = torch.cat([input_samples, last_frame], dim=1)
input_samples = input_samples[:, latent_start_idx:latent_end_idx]
if noise_mask is not None:
original_image = input_samples.to(device)
assert input_samples.shape[1] == noise.shape[1], f"Slice mismatch: {input_samples.shape[1]} vs {noise.shape[1]}"
if add_noise_to_samples:
latent_timestep = timesteps[0]
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
else:
noise = input_samples
# diff diff prep
if noise_mask is not None:
if len(noise_mask.shape) == 4:
noise_mask = noise_mask.squeeze(1)
if audio_end_idx > noise_mask.shape[0]:
noise_mask = noise_mask.repeat(audio_end_idx // noise_mask.shape[0], 1, 1)
noise_mask = noise_mask[audio_start_idx:audio_end_idx]
noise_mask = torch.nn.functional.interpolate(
noise_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
size=(noise.shape[1], noise.shape[2], noise.shape[3]),
mode='trilinear',
align_corners=False
).repeat(1, noise.shape[0], 1, 1, 1)
thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps)
thresholds = thresholds.reshape(-1, 1, 1, 1, 1).to(device)
masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds
# zero padding and vae encode for img cond
if cond_image is not None or cond_frame is not None:
cond_ = cond_image if (is_first_clip or humo_image_cond is None) else cond_frame
cond_frame_num = cond_.shape[2]
video_frames = torch.zeros(1, 3, frame_num-cond_frame_num, target_h, target_w, device=device, dtype=vae.dtype)
padding_frames_pixels_values = torch.concat([cond_.to(device, vae.dtype), video_frames], dim=2)
# encode
vae.to(device)
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
if mode == "multitalk":
latent_motion_frames = y[:, :cur_motion_frames_latent_num] # C T H W
else:
cond_ = cond_image if is_first_clip else cond_frame
latent_motion_frames = vae.encode(cond_.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
vae.to(offload_device)
#motion_frame_index = cur_motion_frames_latent_num if mode == "infinitetalk" else 1
msk = torch.zeros(4, latent_frame_num, lat_h, lat_w, device=device, dtype=dtype)
msk[:, :1] = 1
y = torch.cat([msk, y]) # 4+C T H W
mm.soft_empty_cache()
else:
y = None
latent_motion_frames = noise[:, :1]
partial_humo_cond_input = partial_humo_cond_neg_input = partial_humo_audio = partial_humo_audio_neg = None
if humo_image_cond is not None:
partial_humo_cond_input = humo_image_cond[:, :latent_frame_num]
partial_humo_cond_neg_input = humo_image_cond_neg[:, :latent_frame_num]
if y is not None:
partial_humo_cond_input[:, :1] = y[:, :1]
if humo_reference_count > 0:
partial_humo_cond_input[:, -humo_reference_count:] = humo_image_cond[:, -humo_reference_count:]
partial_humo_cond_neg_input[:, -humo_reference_count:] = humo_image_cond_neg[:, -humo_reference_count:]
if humo_audio is not None:
if is_first_clip:
audio_embs = None
partial_humo_audio, _ = get_audio_emb_window(humo_audio, frame_num, frame0_idx=audio_start_idx)
#zero_audio_pad = torch.zeros(humo_reference_count, *partial_humo_audio.shape[1:], device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
partial_humo_audio[-humo_reference_count:] = 0
partial_humo_audio_neg = torch.zeros_like(partial_humo_audio, device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
if scheduler == "multitalk":
timesteps = list(np.linspace(1000, 1, steps, dtype=np.float32))
timesteps.append(0.)
timesteps = [torch.tensor([t], device=device) for t in timesteps]
timesteps = [timestep_transform(t, shift=shift, num_timesteps=1000) for t in timesteps]
else:
if isinstance(scheduler, dict):
sample_scheduler = copy.deepcopy(scheduler["sample_scheduler"])
timesteps = scheduler["timesteps"]
else:
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, denoise_strength, sigmas=sigmas)
timesteps = [torch.tensor([float(t)], device=device) for t in timesteps] + [torch.tensor([0.], device=device)]
# sample videos
latent = noise
# injecting motion frames
if not is_first_clip and mode == "multitalk":
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[0])
latent[:, :add_latent.shape[1]] = add_latent
if offloaded:
# Load weights
if transformer.patched_linear and gguf_reader is None:
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args)
elif gguf_reader is not None: #handle GGUF
load_weights(transformer, patcher.model["sd"], base_dtype=dtype, transformer_load_device=device, patcher=patcher, gguf=True, reader=gguf_reader, block_swap_args=block_swap_args)
#blockswap init
init_blockswap(transformer, block_swap_args, model)
# Use the appropriate prompt for this section
if len(text_embeds["prompt_embeds"]) > 1:
prompt_index = min(iteration_count, len(text_embeds["prompt_embeds"]) - 1)
positive = [text_embeds["prompt_embeds"][prompt_index]]
log.info(f"Using prompt index: {prompt_index}")
else:
positive = text_embeds["prompt_embeds"]
# uni3c slices
if uni3c_embeds is not None:
vae.to(device)
# Pad original_images if needed
num_frames = original_images.shape[2]
if audio_end_idx > num_frames:
pad_len = audio_end_idx - num_frames
last_frame = original_images[:, :, -1:].repeat(1, 1, pad_len, 1, 1)
padded_images = torch.cat([original_images, last_frame], dim=2)
else:
padded_images = original_images
render_latent = vae.encode(
padded_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype),
device=device, tiled=tiled_vae
).to(dtype)
vae.to(offload_device)
uni3c_data['render_latent'] = render_latent
# unianimate slices
partial_unianim_data = None
if unianim_data is not None:
partial_dwpose = dwpose_data[:, :, latent_start_idx:latent_end_idx]
partial_unianim_data = {
"dwpose": partial_dwpose,
"random_ref": unianim_data["random_ref"],
"strength": unianimate_poses["strength"],
"start_percent": unianimate_poses["start_percent"],
"end_percent": unianimate_poses["end_percent"]
}
# fantasy portrait slices
partial_fantasy_portrait_input = None
if fantasy_portrait_input is not None:
adapter_proj = fantasy_portrait_input["adapter_proj"]
if latent_end_idx > adapter_proj.shape[1]:
pad_len = latent_end_idx - adapter_proj.shape[1]
last_frame = adapter_proj[:, -1:, :, :].repeat(1, pad_len, 1, 1)
padded_proj = torch.cat([adapter_proj, last_frame], dim=1)
else:
padded_proj = adapter_proj
partial_fantasy_portrait_input = fantasy_portrait_input.copy()
partial_fantasy_portrait_input["adapter_proj"] = padded_proj[:, latent_start_idx:latent_end_idx]
mm.soft_empty_cache()
gc.collect()
# sampling loop
sampling_pbar = tqdm(total=len(timesteps)-1, desc=f"Sampling audio indices {audio_start_idx}-{audio_end_idx}", position=0, leave=True)
for i in range(len(timesteps)-1):
timestep = timesteps[i]
latent_model_input = latent.to(device)
if mode == "infinitetalk":
if humo_image_cond is None or not is_first_clip:
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
noise_pred, _, self.cache_state = predict_func(
latent_model_input, cfg[min(i, len(timesteps)-1)], positive, text_embeds["negative_prompt_embeds"],
timestep, i, y, clip_embeds, control_latents, None, partial_unianim_data, audio_proj, control_camera_latents, add_cond,
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs, fantasy_portrait_input=partial_fantasy_portrait_input,
humo_image_cond=partial_humo_cond_input, humo_image_cond_neg=partial_humo_cond_neg_input, humo_audio=partial_humo_audio, humo_audio_neg=partial_humo_audio_neg,
uni3c_data = uni3c_data)
if callback is not None:
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * timestep.to(device) / 1000).detach().permute(1,0,2,3)
callback(step_iteration_count, callback_latent, None, estimated_iterations*(len(timesteps)-1))
del callback_latent
sampling_pbar.update(1)
step_iteration_count += 1
# update latent
if use_tsr:
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
if scheduler == "multitalk":
noise_pred = -noise_pred
dt = (timesteps[i] - timesteps[i + 1]) / 1000
latent = latent + noise_pred * dt[:, None, None, None]
else:
latent = sample_scheduler.step(noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0).to(noise_pred.device), **scheduler_step_args)[0].squeeze(0)
del noise_pred, latent_model_input, timestep
# differential diffusion inpaint
if masks is not None:
if i < len(timesteps) - 1:
image_latent = add_noise(original_image.to(device), noise.to(device), timesteps[i+1])
mask = masks[i].to(latent)
latent = image_latent * mask + latent * (1-mask)
# injecting motion frames
if not is_first_clip and mode == "multitalk":
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1])
latent[:, :add_latent.shape[1]] = add_latent
else:
if humo_image_cond is None or not is_first_clip:
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
del noise, latent_motion_frames
if offload:
offload_transformer(transformer, remove_lora=False)
offloaded = True
if humo_image_cond is not None and humo_reference_count > 0:
latent = latent[:,:-humo_reference_count]
vae.to(device)
videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
vae.to(offload_device)
sampling_pbar.close()
# optional color correction (less relevant for InfiniteTalk)
if colormatch != "disabled":
videos = videos.permute(1, 2, 3, 0).float().numpy()
from color_matcher import ColorMatcher
cm = ColorMatcher()
cm_result_list = []
for img in videos:
if mode == "multitalk":
cm_result = cm.transfer(src=img, ref=original_images[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
else:
cm_result = cm.transfer(src=img, ref=cond_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
cm_result_list.append(torch.from_numpy(cm_result).to(vae.dtype))
videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2)
# optionally save generated samples to disk
if output_path:
video_np = videos.clamp(-1.0, 1.0).add(1.0).div(2.0).mul(255).cpu().float().numpy().transpose(1, 2, 3, 0).astype('uint8')
num_frames_to_save = video_np.shape[0] if is_first_clip else video_np.shape[0] - cur_motion_frames_num
log.info(f"Saving {num_frames_to_save} generated frames to {output_path}")
start_idx = 0 if is_first_clip else cur_motion_frames_num
for i in range(start_idx, video_np.shape[0]):
im = Image.fromarray(video_np[i])
im.save(os.path.join(output_path, f"frame_{img_counter:05d}.png"))
img_counter += 1
else:
gen_video_list.append(videos if is_first_clip else videos[:, cur_motion_frames_num:])
current_condframe_index += 1
iteration_count += 1
# decide whether is done
if arrive_last_frame:
break
# update next condition frames
is_first_clip = False
cur_motion_frames_num = motion_frame
cond_ = videos[:, -cur_motion_frames_num:].unsqueeze(0)
if mode == "infinitetalk":
cond_frame = cond_
else:
cond_image = cond_
del videos, latent
# Repeat audio emb
if multitalk_embeds is not None:
audio_start_idx += (frame_num - cur_motion_frames_num - humo_reference_count)
audio_end_idx = audio_start_idx + clip_length
if audio_end_idx >= len(audio_embedding[0]):
arrive_last_frame = True
miss_lengths = []
source_frames = []
for human_inx in range(human_num):
source_frame = len(audio_embedding[human_inx])
source_frames.append(source_frame)
if audio_end_idx >= len(audio_embedding[human_inx]):
log.warning(f"Audio embedding for subject {human_inx} not long enough: {len(audio_embedding[human_inx])}, need {audio_end_idx}, padding...")
miss_length = audio_end_idx - len(audio_embedding[human_inx]) + 3
log.warning(f"Padding length: {miss_length}")
if encoded_silence is not None:
add_audio_emb = encoded_silence[-1*miss_length:]
else:
add_audio_emb = torch.flip(audio_embedding[human_inx][-1*miss_length:], dims=[0])
audio_embedding[human_inx] = torch.cat([audio_embedding[human_inx], add_audio_emb.to(device, dtype)], dim=0)
miss_lengths.append(miss_length)
else:
miss_lengths.append(0)
if mode == "infinitetalk" and current_condframe_index >= original_images.shape[2]:
last_frame = original_images[:, :, -1:, :, :]
miss_length = 1
original_images = torch.cat([original_images, last_frame.repeat(1, 1, miss_length, 1, 1)], dim=2)
if not output_path:
gen_video_samples = torch.cat(gen_video_list, dim=1)
else:
gen_video_samples = torch.zeros(3, 1, 64, 64) # dummy output
if force_offload:
if not model["auto_cpu_offload"]:
offload_transformer(transformer)
try:
print_memory(device)
torch.cuda.reset_peak_memory_stats(device)
except:
pass
return {"video": gen_video_samples.permute(1, 2, 3, 0), "output_path": output_path},
+19 -1
View File
@@ -7,6 +7,8 @@ from ..utils import log, set_module_tensor_to_device
import os
import json
import datetime
import scipy.signal as ss
import numpy as np
script_directory = os.path.dirname(os.path.abspath(__file__))
folder_paths.add_model_folder_path("wav2vec2", os.path.join(folder_paths.models_dir, "wav2vec2"))
@@ -134,6 +136,15 @@ def loudness_norm(audio_array, sr=16000, lufs=-23):
return audio_array
normalized_audio = pyloudnorm.normalize.loudness(audio_array, loudness, lufs)
return normalized_audio
def _add_noise_floor(audio, noise_db=-45):
noise_amp = 10 ** (noise_db / 20)
noise = np.random.randn(len(audio)) * noise_amp
return audio + noise
def _smooth_transients(audio, sr=16000):
b, a = ss.butter(3, 3000 / (sr/2))
return ss.lfilter(b, a, audio)
class MultiTalkWav2VecEmbeds:
@classmethod
@@ -153,6 +164,8 @@ class MultiTalkWav2VecEmbeds:
"audio_3": ("AUDIO",),
"audio_4": ("AUDIO",),
"ref_target_masks": ("MASK", {"tooltip": "Per-speaker semantic mask(s) in pixel space. Supply one mask per speaker (plus optional background) to guide mouth assignment"}),
"add_noise_floor": ("BOOLEAN", {"default": False, "tooltip": "Add a low-level noise floor to the audio to reduce silent gaps"}),
"smooth_transients": ("BOOLEAN", {"default": False, "tooltip": "Apply a low-pass filter to the audio to smooth out transients"}),
}
}
@@ -161,7 +174,8 @@ class MultiTalkWav2VecEmbeds:
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, wav2vec_model, normalize_loudness, fps, num_frames, audio_1, audio_scale, audio_cfg_scale, multi_audio_type, audio_2=None, audio_3=None, audio_4=None, ref_target_masks=None):
def process(self, wav2vec_model, normalize_loudness, fps, num_frames, audio_1, audio_scale, audio_cfg_scale, multi_audio_type, audio_2=None, audio_3=None, audio_4=None,
ref_target_masks=None, add_noise_floor=False, smooth_transients=False):
model_type = wav2vec_model["model_type"]
if not "tencent" in model_type.lower():
raise ValueError("Only tencent wav2vec2 models supported by MultiTalk")
@@ -207,6 +221,10 @@ class MultiTalkWav2VecEmbeds:
if normalize_loudness:
audio_segment = loudness_norm(audio_segment, sr=sr)
if add_noise_floor:
audio_segment = _add_noise_floor(audio_segment, noise_db=-45)
if smooth_transients:
audio_segment = _smooth_transients(audio_segment, sr=sr)
audio_feature = np.squeeze(
wav2vec2_feature_extractor(audio_segment, sampling_rate=sr).input_values
+346 -267
View File
@@ -1,11 +1,8 @@
import os, gc, math
import torch
import torch.nn.functional as F
import numpy as np
import hashlib
from .wanvideo.schedulers import get_scheduler, scheduler_list
from .utils import(log, clip_encode_image_tiled, add_noise_to_reference_video, set_module_tensor_to_device)
from .taehv import TAEHV
@@ -22,30 +19,6 @@ offload_device = mm.unet_offload_device()
VAE_STRIDE = (4, 8, 8)
PATCH_SIZE = (1, 2, 2)
class WanVideoAddVideoPromptEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"image_embeds": ("WANVIDIMAGE_EMBEDS",),
"video_prompt_embeds": ("WANVIDIMAGE_EMBEDS",),
"video_prompt_latents": ("LATENT", ),
"text_embeds": ("WANVIDEOTEXTEMBEDS", ),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
def add(self, image_embeds, video_prompt_embeds, video_prompt_latents, text_embeds):
updated = dict(image_embeds)
updated["video_prompt_embeds"] = video_prompt_embeds
updated["video_prompt_embeds"]["video_prompt_latents"] = video_prompt_latents["samples"][0]
updated["video_prompt_embeds"]["text_embeds"] = text_embeds
return (updated,)
class WanVideoEnhanceAVideo:
@classmethod
@@ -789,6 +762,98 @@ class WanVideoAddStandInLatent:
updated = dict(embeds)
updated["standin_input"] = new_entry
return (updated,)
class WanVideoAddBindweaveEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"reference_latents": ("LATENT", {"tooltip": "Reference image to encode"}),
},
"optional": {
"ref_masks": ("MASK", {"tooltip": "Reference mask to encode"}),
"qwenvl_embeds_pos": ("QWENVL_EMBEDS", {"tooltip": "Qwen-VL image embeddings for the reference image"}),
"qwenvl_embeds_neg": ("QWENVL_EMBEDS", {"tooltip": "Qwen-VL image embeddings for the reference image"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "LATENT", "MASK",)
RETURN_NAMES = ("image_embeds", "image_embed_preview", "mask_preview",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, reference_latents, ref_masks=None, qwenvl_embeds_pos=None, qwenvl_embeds_neg=None):
updated = dict(embeds)
image_embeds = embeds["image_embeds"]
max_refs = 4
num_refs = reference_latents["samples"].shape[0]
pad = torch.zeros(image_embeds.shape[0], max_refs-num_refs, image_embeds.shape[2], image_embeds.shape[3], device=image_embeds.device, dtype=image_embeds.dtype)
if num_refs < max_refs:
image_embeds = torch.cat([pad, image_embeds], dim=1)
ref_latents = [ref_latent for ref_latent in reference_latents["samples"]]
image_embeds = torch.cat([*ref_latents, image_embeds], dim=1)
mask = embeds.get("mask", None)
if mask is not None:
mask_pad = torch.zeros(mask.shape[0], max_refs-num_refs, mask.shape[2], mask.shape[3], device=mask.device, dtype=mask.dtype)
if num_refs < max_refs:
mask = torch.cat([mask_pad, mask], dim=1)
if ref_masks is not None:
ref_mask_ = common_upscale(ref_masks.unsqueeze(1), mask.shape[3], mask.shape[2], "nearest", "disabled").movedim(0,1)
ref_mask_ = torch.cat([ref_mask_, torch.zeros(3, ref_mask_.shape[1], ref_mask_.shape[2], ref_mask_.shape[3], device=ref_mask_.device, dtype=ref_mask_.dtype)])
mask = torch.cat([ref_mask_, mask], dim=1)
else:
mask = torch.cat([torch.ones(mask.shape[0], num_refs, mask.shape[2], mask.shape[3], device=mask.device, dtype=mask.dtype), mask], dim=1)
updated["mask"] = mask
clip_embeds = updated.get("clip_context", None)
if clip_embeds is not None:
B, T, C = clip_embeds.shape
target_len = max_refs * 257 # 4 * 257 = 1028
if T < target_len:
pad = torch.zeros(B, target_len - T, C, device=clip_embeds.device, dtype=clip_embeds.dtype)
padded_embeds = torch.cat([clip_embeds, pad], dim=1)
log.info(f"Padded clip embeds from {clip_embeds.shape} to {padded_embeds.shape} for Bindweave")
updated["clip_context"] = padded_embeds
else:
updated["clip_context"] = clip_embeds
updated["image_embeds"] = image_embeds
updated["qwenvl_embeds_pos"] = qwenvl_embeds_pos
updated["qwenvl_embeds_neg"] = qwenvl_embeds_neg
return (updated, {"samples": image_embeds.unsqueeze(0)}, mask[0].float())
class TextImageEncodeQwenVL():
@classmethod
def INPUT_TYPES(s):
return {"required": {
"clip": ("CLIP",),
"prompt": ("STRING", {"default": "", "multiline": True}),
},
"optional": {
"image": ("IMAGE", ),
}
}
RETURN_TYPES = ("QWENVL_EMBEDS",)
RETURN_NAMES = ("qwenvl_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(cls, clip, prompt, image=None):
if image is None:
input_images = []
llama_template = None
else:
input_images = [image[:, :, :, :3]]
llama_template = "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>{}<|im_end|>\n<|im_start|>assistant\n"
tokens = clip.tokenize(prompt, images=input_images, llama_template=llama_template)
conditioning = clip.encode_from_tokens_scheduled(tokens)
print("Qwen-VL embeds shape:", conditioning[0][0].shape)
return (conditioning[0][0],)
class WanVideoAddMTVMotion:
@classmethod
@@ -823,6 +888,83 @@ class WanVideoAddMTVMotion:
updated["mtv_crafter_motion"] = new_entry
return (updated,)
class WanVideoAddStoryMemLatents:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"vae": ("WANVAE",),
"embeds": ("WANVIDIMAGE_EMBEDS",),
"memory_images": ("IMAGE",),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, vae, embeds, memory_images):
updated = dict(embeds)
story_mem_latents, = WanVideoEncodeLatentBatch().encode(vae, memory_images)
updated["story_mem_latents"] = story_mem_latents["samples"].squeeze(2).permute(1, 0, 2, 3) # [C, T, H, W]
return (updated,)
class WanVideoSVIProEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"anchor_samples": ("LATENT", {"tooltip": "Initial start image encoded"}),
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}),
},
"optional": {
"prev_samples": ("LATENT", {"tooltip": "Last latent from previous generation"}),
"motion_latent_count": ("INT", {"default": 1, "min": 0, "max": 100, "step": 1, "tooltip": "Number of latents used to continue"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, anchor_samples, num_frames, prev_samples=None, motion_latent_count=1):
anchor_latent = anchor_samples["samples"][0].clone()
C, T, H, W = anchor_latent.shape
total_latents = (num_frames - 1) // 4 + 1
device = anchor_latent.device
dtype = anchor_latent.dtype
if prev_samples is None or motion_latent_count == 0:
padding_size = total_latents - anchor_latent.shape[1]
padding = torch.zeros(C, padding_size, H, W, dtype=dtype, device=device)
y = torch.concat([anchor_latent, padding], dim=1)
else:
prev_latent = prev_samples["samples"][0].clone()
motion_latent = prev_latent[:, -motion_latent_count:]
padding_size = total_latents - anchor_latent.shape[1] - motion_latent.shape[1]
padding = torch.zeros(C, padding_size, H, W, dtype=dtype, device=device)
y = torch.concat([anchor_latent, motion_latent, padding], dim=1)
msk = torch.ones(1, num_frames, H, W, device=device, dtype=dtype)
msk[:, 1:] = 0
msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
msk = msk.view(1, msk.shape[1] // 4, 4, H, W)
msk = msk.transpose(1, 2)[0]
image_embeds = {
"image_embeds": y,
"num_frames": num_frames,
"lat_h": H,
"lat_w": W,
"mask": msk
}
return (image_embeds,)
#region I2V encode
class WanVideoImageToVideoEncode:
@classmethod
@@ -847,6 +989,8 @@ class WanVideoImageToVideoEncode:
"extra_latents": ("LATENT", {"tooltip": "Extra latents to add to the input front, used for Skyreels A2 reference images"}),
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
"add_cond_latents": ("ADD_COND_LATENTS", {"advanced": True, "tooltip": "Additional cond latents WIP"}),
"augment_empty_frames": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "EXPERIMENTAL: Augment empty frames with the difference to the start image to force more motion"}),
"empty_frame_pad_image": ("IMAGE", {"tooltip": "Use this image to pad empty frames instead of gray, used with SVI-shot and SVI 2.0 LoRAs"}),
}
}
@@ -856,18 +1000,14 @@ class WanVideoImageToVideoEncode:
CATEGORY = "WanVideoWrapper"
def process(self, width, height, num_frames, force_offload, noise_aug_strength,
start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False,
temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None, vae=None):
if start_image is None and end_image is None and add_cond_latents is None:
return WanVideoEmptyEmbeds().process(
num_frames, width, height, control_embeds=control_embeds, extra_latents=extra_latents,
)
start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False,
temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None, vae=None, augment_empty_frames=0.0, empty_frame_pad_image=None):
if vae is None:
raise ValueError("VAE is required for image encoding.")
H = height
W = width
lat_h = H // vae.upsampling_factor
lat_w = W // vae.upsampling_factor
@@ -892,6 +1032,8 @@ class WanVideoImageToVideoEncode:
mask = torch.cat([mask, torch.zeros(base_frames - mask.shape[0], lat_h, lat_w, device=device)])
mask = mask.unsqueeze(0).to(device, vae.dtype)
pixel_mask = mask.clone()
# Repeat first frame and optionally end frame
start_mask_repeated = torch.repeat_interleave(mask[:, 0:1], repeats=4, dim=1) # T, C, H, W
if end_image is not None and not fun_or_fl2v_model:
@@ -914,7 +1056,7 @@ class WanVideoImageToVideoEncode:
resized_start_image = resized_start_image * 2 - 1
if noise_aug_strength > 0.0:
resized_start_image = add_noise_to_reference_video(resized_start_image, ratio=noise_aug_strength)
if end_image is not None:
end_image = end_image[..., :3]
if end_image.shape[1] != H or end_image.shape[2] != W:
@@ -924,30 +1066,46 @@ class WanVideoImageToVideoEncode:
resized_end_image = resized_end_image * 2 - 1
if noise_aug_strength > 0.0:
resized_end_image = add_noise_to_reference_video(resized_end_image, ratio=noise_aug_strength)
# Concatenate image with zero frames and encode
if temporal_mask is None:
if start_image is not None and end_image is None:
zero_frames = torch.zeros(3, num_frames-start_image.shape[0], H, W, device=device, dtype=vae.dtype)
concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames], dim=1)
del resized_start_image, zero_frames
elif start_image is None and end_image is not None:
zero_frames = torch.zeros(3, num_frames-end_image.shape[0], H, W, device=device, dtype=vae.dtype)
concatenated = torch.cat([zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1)
del zero_frames
elif start_image is None and end_image is None:
concatenated = torch.zeros(3, num_frames, H, W, device=device, dtype=vae.dtype)
else:
if fun_or_fl2v_model:
zero_frames = torch.zeros(3, num_frames-(start_image.shape[0]+end_image.shape[0]), H, W, device=device, dtype=vae.dtype)
else:
zero_frames = torch.zeros(3, num_frames-1, H, W, device=device, dtype=vae.dtype)
concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1)
del resized_start_image, zero_frames
if start_image is not None and end_image is None:
zero_frames = torch.zeros(3, num_frames-start_image.shape[0], H, W, device=device, dtype=vae.dtype)
concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames], dim=1)
del resized_start_image, zero_frames
elif start_image is None and end_image is not None:
zero_frames = torch.zeros(3, num_frames-end_image.shape[0], H, W, device=device, dtype=vae.dtype)
concatenated = torch.cat([zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1)
del zero_frames
elif start_image is None and end_image is None:
concatenated = torch.zeros(3, num_frames, H, W, device=device, dtype=vae.dtype)
else:
temporal_mask = common_upscale(temporal_mask.unsqueeze(1), W, H, "nearest", "disabled").squeeze(1)
concatenated = resized_start_image[:,:num_frames].to(vae.dtype)# * temporal_mask[:num_frames].unsqueeze(0).to(vae.dtype)
del resized_start_image, temporal_mask
if fun_or_fl2v_model:
zero_frames = torch.zeros(3, num_frames-(start_image.shape[0]+end_image.shape[0]), H, W, device=device, dtype=vae.dtype)
else:
zero_frames = torch.zeros(3, num_frames-1, H, W, device=device, dtype=vae.dtype)
concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1)
del resized_start_image, zero_frames
if empty_frame_pad_image is not None:
pad_img = empty_frame_pad_image.clone()[..., :3]
if pad_img.shape[1] != H or pad_img.shape[2] != W:
pad_img = common_upscale(pad_img.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1)
pad_img = (pad_img.movedim(-1, 0) * 2 - 1).to(device, dtype=vae.dtype)
num_pad_frames = pad_img.shape[1]
num_target_frames = concatenated.shape[1]
if num_pad_frames < num_target_frames:
pad_img = torch.cat([pad_img, pad_img[:, -1:].expand(-1, num_target_frames - num_pad_frames, -1, -1)], dim=1)
else:
pad_img = pad_img[:, :num_target_frames]
frame_is_empty = (pixel_mask[0].mean(dim=(-2, -1)) < 0.5)[:concatenated.shape[1]].clone()
if start_image is not None:
frame_is_empty[:start_image.shape[0]] = False
if end_image is not None:
frame_is_empty[-end_image.shape[0]:] = False
concatenated[:, frame_is_empty] = pad_img[:, frame_is_empty]
mm.soft_empty_cache()
gc.collect()
@@ -965,6 +1123,9 @@ class WanVideoImageToVideoEncode:
has_ref = True
y[:, :1] *= start_latent_strength
y[:, -1:] *= end_latent_strength
if augment_empty_frames > 0.0:
frame_is_empty = (mask[0].mean(dim=(-2, -1)) < 0.5).view(1, -1, 1, 1)
y = y[:, :1] + (y - y[:, :1]) * ((augment_empty_frames+1) * frame_is_empty + ~frame_is_empty)
# Calculate maximum sequence length
patches_per_frame = lat_h * lat_w // (PATCH_SIZE[1] * PATCH_SIZE[2])
@@ -973,14 +1134,14 @@ class WanVideoImageToVideoEncode:
if add_cond_latents is not None:
add_cond_latents["ref_latent_neg"] = vae.encode(torch.zeros(1, 3, 1, H, W, device=device, dtype=vae.dtype), device)
if force_offload:
vae.model.to(offload_device)
mm.soft_empty_cache()
gc.collect()
image_embeds = {
"image_embeds": y,
"image_embeds": y.cpu(),
"clip_context": clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None,
"negative_clip_context": clip_embeds.get("negative_clip_embeds", None) if clip_embeds is not None else None,
"max_seq_len": max_seq_len,
@@ -992,11 +1153,11 @@ class WanVideoImageToVideoEncode:
"fun_or_fl2v_model": fun_or_fl2v_model,
"has_ref": has_ref,
"add_cond_latents": add_cond_latents,
"mask": mask
"mask": mask.cpu()
}
return (image_embeds,)
# region WanAnimate
class WanVideoAnimateEmbeds:
@classmethod
@@ -1007,15 +1168,15 @@ class WanVideoAnimateEmbeds:
"height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}),
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}),
"force_offload": ("BOOLEAN", {"default": True}),
"frame_window_size": ("INT", {"default": 77, "min": 1, "max": 1000, "step": 1, "tooltip": "Number of frames to use for temporal attention window"}),
"frame_window_size": ("INT", {"default": 77, "min": 1, "max": 10000, "step": 1, "tooltip": "Number of frames to use for temporal attention window"}),
"colormatch": (
[
[
'disabled',
'mkl',
'hm',
'reinhard',
'mvgd',
'hm-mvgd-hm',
'hm',
'reinhard',
'mvgd',
'hm-mvgd-hm',
'hm-mkl-hm',
], {
"default": 'disabled', "tooltip": "Color matching method to use between the windows"
@@ -1182,6 +1343,46 @@ class WanVideoAnimateEmbeds:
}
return (image_embeds,)
# region UniLumos
class WanVideoUniLumosEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"width": ("INT", {"default": 832, "min": 64, "max": 8096, "step": 8, "tooltip": "Width of the image to encode"}),
"height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}),
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}),
},
"optional": {
"foreground_latents": ("LATENT", {"tooltip": "Video foreground latents"}),
"background_latents": ("LATENT", {"tooltip": "Video background latents"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", )
RETURN_NAMES = ("image_embeds",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, num_frames, width, height, foreground_latents=None, background_latents=None):
target_shape = (16, (num_frames - 1) // VAE_STRIDE[0] + 1,
height // VAE_STRIDE[1],
width // VAE_STRIDE[2])
embeds = {
"target_shape": target_shape,
"num_frames": num_frames,
}
if foreground_latents is not None:
embeds["foreground_latents"] = foreground_latents["samples"][0]
else:
embeds["foreground_latents"] = torch.zeros(target_shape[0], target_shape[1], target_shape[2], target_shape[3], device=torch.device("cpu"), dtype=torch.float32)
if background_latents is not None:
embeds["background_latents"] = background_latents["samples"][0]
else:
embeds["background_latents"] = torch.zeros(target_shape[0], target_shape[1], target_shape[2], target_shape[3], device=torch.device("cpu"), dtype=torch.float32)
return (embeds,)
class WanVideoEmptyEmbeds:
@classmethod
@@ -1702,33 +1903,7 @@ class WanVideoContextOptions:
}
return (context_options,)
class WanVideoFlowEdit:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"source_embeds": ("WANVIDEOTEXTEMBEDS", ),
"skip_steps": ("INT", {"default": 4, "min": 0}),
"drift_steps": ("INT", {"default": 0, "min": 0}),
"drift_flow_shift": ("FLOAT", {"default": 3.0, "min": 1.0, "max": 30.0, "step": 0.01}),
"source_cfg": ("FLOAT", {"default": 6.0, "min": 0.0, "max": 30.0, "step": 0.01}),
"drift_cfg": ("FLOAT", {"default": 6.0, "min": 0.0, "max": 30.0, "step": 0.01}),
},
"optional": {
"source_image_embeds": ("WANVIDIMAGE_EMBEDS", ),
}
}
RETURN_TYPES = ("FLOWEDITARGS", )
RETURN_NAMES = ("flowedit_args",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Flowedit options for WanVideo"
def process(self, **kwargs):
return (kwargs,)
class WanVideoLoopArgs:
@classmethod
def INPUT_TYPES(s):
@@ -1800,155 +1975,6 @@ class WanVideoFreeInitArgs:
def process(self, **kwargs):
return (kwargs,)
class WanVideoScheduler: #WIP
@classmethod
def INPUT_TYPES(s):
return {"required": {
"scheduler": (scheduler_list, {"default": "unipc"}),
"steps": ("INT", {"default": 30, "min": 1, "tooltip": "Number of steps for the scheduler"}),
"shift": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 1000.0, "step": 0.01}),
"start_step": ("INT", {"default": 0, "min": 0, "tooltip": "Starting step for the scheduler"}),
"end_step": ("INT", {"default": -1, "min": -1, "tooltip": "Ending step for the scheduler"})
},
"optional": {
"sigmas": ("SIGMAS", ),
},
"hidden": {
"unique_id": "UNIQUE_ID",
},
}
RETURN_TYPES = ("SIGMAS", "INT", "FLOAT", scheduler_list, "INT", "INT",)
RETURN_NAMES = ("sigmas", "steps", "shift", "scheduler", "start_step", "end_step")
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
def process(self, scheduler, steps, start_step, end_step, shift, unique_id, sigmas=None):
sample_scheduler, timesteps, start_idx, end_idx = get_scheduler(
scheduler,
steps,
start_step, end_step, shift,
device,
sigmas=sigmas,
log_timesteps=True)
scheduler_dict = {
"sample_scheduler": sample_scheduler,
"timesteps": timesteps,
}
try:
from server import PromptServer
import io
import base64
import matplotlib.pyplot as plt
except:
PromptServer = None
if unique_id and PromptServer is not None:
try:
# Plot sigmas and save to a buffer
sigmas_np = sample_scheduler.full_sigmas.cpu().numpy()
if not np.isclose(sigmas_np[-1], 0.0, atol=1e-6):
sigmas_np = np.append(sigmas_np, 0.0)
buf = io.BytesIO()
fig = plt.figure(facecolor='#353535')
ax = fig.add_subplot(111)
ax.set_facecolor('#353535') # Set axes background color
x_values = range(0, len(sigmas_np))
ax.plot(x_values, sigmas_np)
# Annotate each sigma value
ax.scatter(x_values, sigmas_np, color='white', s=20, zorder=3) # Small dots at each sigma
for x, y in zip(x_values, sigmas_np):
# Show all annotations if few steps, or just show split step annotations
show_annotation = len(sigmas_np) <= 10
is_split_step = (start_idx > 0 and x == start_idx) or (end_idx != -1 and x == end_idx + 1)
if show_annotation or is_split_step:
color = 'orange'
if is_split_step:
color = 'yellow'
ax.annotate(f"{y:.3f}", (x, y), textcoords="offset points", xytext=(10, 1), ha='center', color=color, fontsize=12)
ax.set_xticks(x_values)
ax.set_title("Sigmas", color='white') # Title font color
ax.set_xlabel("Step", color='white') # X label font color
ax.set_ylabel("Sigma Value", color='white') # Y label font color
ax.tick_params(axis='x', colors='white', labelsize=10) # X tick color
ax.tick_params(axis='y', colors='white', labelsize=10) # Y tick color
# Add split point if end_step is defined
end_idx += 1
if end_idx != -1 and 0 <= end_idx < len(sigmas_np) - 1:
ax.axvline(end_idx, color='red', linestyle='--', linewidth=2, label='end_step split')
# Add split point if start_step is defined
if start_idx > 0 and 0 <= start_idx < len(sigmas_np):
ax.axvline(start_idx, color='green', linestyle='--', linewidth=2, label='start_step split')
if (end_idx != -1 and 0 <= end_idx < len(sigmas_np)) or (start_idx > 0 and 0 <= start_idx < len(sigmas_np)):
handles, labels = ax.get_legend_handles_labels()
if labels:
ax.legend()
if start_idx < end_idx and 0 <= start_idx < len(sigmas_np) and 0 < end_idx < len(sigmas_np):
ax.axvspan(start_idx, end_idx, color='lightblue', alpha=0.1, label='Sampled Range')
plt.tight_layout()
plt.savefig(buf, format='png')
plt.close(fig)
buf.seek(0)
img_base64 = base64.b64encode(buf.read()).decode('utf-8')
buf.close()
# Send as HTML img tag with base64 data
html_img = f"<img src='data:image/png;base64,{img_base64}' alt='Sigmas Plot' style='max-width:100%; height:100%; overflow:hidden; display:block;'>"
PromptServer.instance.send_progress_text(html_img, unique_id)
except Exception as e:
print("Failed to send sigmas plot:", e)
pass
return (sigmas, steps, shift, scheduler_dict, start_step, end_step)
class WanVideoSchedulerSA_ODE:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"use_adaptive_order": ("BOOLEAN", {"default": False, "tooltip": "Use adaptive order"}),
"use_velocity_smoothing": ("BOOLEAN", {"default": True, "tooltip": "Use velocity smoothing"}),
"convergence_threshold": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.001, "tooltip": "Convergence threshold for velocity smoothing"}),
"smoothing_factor": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 1.0, "step": 0.001, "tooltip": "Smoothing factor for velocity smoothing"}),
"steps": ("INT", {"default": 30, "min": 1, "tooltip": "Number of steps for the scheduler"}),
"shift": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 1000.0, "step": 0.01}),
"start_step": ("INT", {"default": 0, "min": 0, "tooltip": "Starting step for the scheduler"}),
"end_step": ("INT", {"default": -1, "min": -1, "tooltip": "Ending step for the scheduler"})
},
"optional": {
"sigmas": ("SIGMAS", ),
},
}
RETURN_TYPES = ("SIGMAS", "INT", "FLOAT", scheduler_list, "INT", "INT",)
RETURN_NAMES = ("sigmas", "steps", "shift", "scheduler", "start_step", "end_step")
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
def process(self, steps, start_step, end_step, shift, use_adaptive_order, use_velocity_smoothing, convergence_threshold, smoothing_factor, sigmas=None):
sample_scheduler, timesteps, _, _ = get_scheduler(
scheduler="sa_ode_stable/lowstep",
steps=steps,
start_step=start_step, end_step=end_step, shift=shift,
device=device,
sigmas=sigmas,
log_timesteps=True,
use_adaptive_order=use_adaptive_order,
use_velocity_smoothing=use_velocity_smoothing,
convergence_threshold=convergence_threshold,
smoothing_factor=smoothing_factor
)
scheduler_dict = {
"sample_scheduler": sample_scheduler,
"timesteps": timesteps,
}
return (sigmas, steps, shift, scheduler_dict, start_step, end_step)
rope_functions = ["default", "comfy", "comfy_chunked"]
class WanVideoRoPEFunction:
@@ -1979,6 +2005,53 @@ class WanVideoRoPEFunction:
return (rope_func_dict,)
return (rope_function,)
#region TTM
class WanVideoAddTTMLatents:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"reference_latents": ("LATENT", {"tooltip": "Latents used as reference for TTM"}),
"mask": ("MASK", {"tooltip": "Mask used for TTM"}),
"start_step": ("INT", {"default": 0, "min": -1, "max": 1000, "step": 1, "tooltip": "Start step for whole denoising process"}),
"end_step": ("INT", {"default": 1, "min": 1, "max": 1000, "step": 1, "tooltip": "The step to stop applying TTM"}),
},
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", )
RETURN_NAMES = ("image_embeds", )
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "https://github.com/time-to-move/TTM"
def add(self, embeds, reference_latents, mask, start_step, end_step):
if end_step < max(0, start_step):
raise ValueError(f"`end_step` ({end_step}) must be >= `start_step` ({start_step}).")
mask_sampled = mask[::4]
mask_sampled = mask_sampled.unsqueeze(1).unsqueeze(0) # [1, T, 1, H, W]
vae_upscale_factor = 8
if reference_latents["samples"].shape[1] == 48:
vae_upscale_factor = 16
# Upsample spatially to latent resolution
H_latent = mask_sampled.shape[-2] // vae_upscale_factor
W_latent = mask_sampled.shape[-1] // vae_upscale_factor
mask_latent = F.interpolate(
mask_sampled.float(),
size=(mask_sampled.shape[2], H_latent, W_latent),
mode="nearest"
)
updated = dict(embeds)
updated["ttm_reference_latents"] = reference_latents["samples"].squeeze(0)
updated["ttm_mask"] = mask_latent.squeeze(0).movedim(1, 0) # [T, 1, H, W]
updated["ttm_start_step"] = start_step
updated["ttm_end_step"] = end_step
return (updated,)
#region VideoDecode
class WanVideoDecode:
@@ -2000,7 +2073,7 @@ class WanVideoDecode:
"tile_stride_y": ("INT", {"default": 128, "min": 32, "max": 2040, "step": 8, "tooltip": "Tile stride height in pixels. Smaller values use less VRAM but will introduce more seams."}),
},
"optional": {
"normalization": (["default", "minmax"], {"advanced": True}),
"normalization": (["default", "minmax", "none"], {"advanced": True}),
}
}
@@ -2024,7 +2097,7 @@ class WanVideoDecode:
video.clamp_(-1.0, 1.0)
video.add_(1.0).div_(2.0)
return video.cpu().float(),
latents = samples["samples"]
latents = samples["samples"].clone()
end_image = samples.get("end_image", None)
has_ref = samples.get("has_ref", False)
drop_last = samples.get("drop_last", False)
@@ -2043,25 +2116,24 @@ class WanVideoDecode:
if drop_last:
latents = latents[:, :, :-1]
if type(vae).__name__ == "TAEHV":
if type(vae).__name__ == "TAEHV":
images = vae.decode_video(latents.permute(0, 2, 1, 3, 4), cond=flashvsr_LQ_images.to(vae.dtype) if flashvsr_LQ_images is not None else None)[0].permute(1, 0, 2, 3)
images = torch.clamp(images, 0.0, 1.0)
images = images.permute(1, 2, 3, 0).cpu().float()
return (images,)
else:
if end_image is not None:
enable_vae_tiling = False
images = vae.decode(latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//8, tile_y//8), tile_stride=(tile_stride_x//8, tile_stride_y//8))[0]
images = images.cpu().float()
if normalization == "minmax":
images.sub_(images.min()).div_(images.max() - images.min())
else:
images.clamp_(-1.0, 1.0)
images.add_(1.0).div_(2.0)
if normalization != "none":
if normalization == "minmax":
images.sub_(images.min()).div_(images.max() - images.min())
else:
images.clamp_(-1.0, 1.0)
images.add_(1.0).div_(2.0)
if is_looped:
temp_latents = torch.cat([latents[:, :, -3:]] + [latents[:, :, :2]], dim=2)
temp_images = vae.decode(temp_latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))[0]
@@ -2072,7 +2144,7 @@ class WanVideoDecode:
if end_image is not None:
images = images[:, 0:-1]
vae.to(offload_device)
mm.soft_empty_cache()
@@ -2101,7 +2173,7 @@ class WanVideoEncodeLatentBatch:
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Encodes a batch of images individually to create a latent video batch where each video is a single frame, useful for I2V init purposes, for example as multiple context window inits"
def encode(self, vae, images, enable_vae_tiling, tile_x, tile_y, tile_stride_x, tile_stride_y, latent_strength=1.0):
def encode(self, vae, images, enable_vae_tiling=False, tile_x=272, tile_y=272, tile_stride_x=144, tile_stride_y=128, latent_strength=1.0):
vae.to(device)
images = images.clone()
@@ -2123,7 +2195,7 @@ class WanVideoEncodeLatentBatch:
latent = vae.encode(img.unsqueeze(0).unsqueeze(0).permute(0, 4, 1, 2, 3), device=device, tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))
else:
latent = vae.encode(img.unsqueeze(0).unsqueeze(0).permute(0, 4, 1, 2, 3), device=device, tiled=enable_vae_tiling)
if latent_strength != 1.0:
latent *= latent_strength
latent_list.append(latent.squeeze(0).cpu())
@@ -2183,14 +2255,16 @@ class WanVideoEncode:
latents = latents.permute(0, 2, 1, 3, 4)
else:
latents = vae.encode(image * 2.0 - 1.0, device=device, tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))
vae.to(offload_device)
if latent_strength != 1.0:
latents *= latent_strength
latents = latents.cpu()
log.info(f"WanVideoEncode: Encoded latents shape {latents.shape}")
mm.soft_empty_cache()
return ({"samples": latents, "noise_mask": mask},)
NODE_CLASS_MAPPINGS = {
@@ -2205,7 +2279,6 @@ NODE_CLASS_MAPPINGS = {
"WanVideoEnhanceAVideo": WanVideoEnhanceAVideo,
"WanVideoContextOptions": WanVideoContextOptions,
"WanVideoTextEmbedBridge": WanVideoTextEmbedBridge,
"WanVideoFlowEdit": WanVideoFlowEdit,
"WanVideoControlEmbeds": WanVideoControlEmbeds,
"WanVideoSLG": WanVideoSLG,
"WanVideoLoopArgs": WanVideoLoopArgs,
@@ -2221,7 +2294,6 @@ NODE_CLASS_MAPPINGS = {
"WanVideoBlockList": WanVideoBlockList,
"WanVideoTextEncodeCached": WanVideoTextEncodeCached,
"WanVideoAddExtraLatent": WanVideoAddExtraLatent,
"WanVideoScheduler": WanVideoScheduler,
"WanVideoAddStandInLatent": WanVideoAddStandInLatent,
"WanVideoAddControlEmbeds": WanVideoAddControlEmbeds,
"WanVideoAddMTVMotion": WanVideoAddMTVMotion,
@@ -2229,8 +2301,12 @@ NODE_CLASS_MAPPINGS = {
"WanVideoAddPusaNoise": WanVideoAddPusaNoise,
"WanVideoAnimateEmbeds": WanVideoAnimateEmbeds,
"WanVideoAddLucyEditLatents": WanVideoAddLucyEditLatents,
"WanVideoSchedulerSA_ODE": WanVideoSchedulerSA_ODE,
"WanVideoAddVideoPromptEmbeds": WanVideoAddVideoPromptEmbeds,
"WanVideoAddBindweaveEmbeds": WanVideoAddBindweaveEmbeds,
"TextImageEncodeQwenVL": TextImageEncodeQwenVL,
"WanVideoUniLumosEmbeds": WanVideoUniLumosEmbeds,
"WanVideoAddTTMLatents": WanVideoAddTTMLatents,
"WanVideoAddStoryMemLatents": WanVideoAddStoryMemLatents,
"WanVideoSVIProEmbeds": WanVideoSVIProEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -2246,7 +2322,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoEnhanceAVideo": "WanVideo Enhance-A-Video",
"WanVideoContextOptions": "WanVideo Context Options",
"WanVideoTextEmbedBridge": "WanVideo TextEmbed Bridge",
"WanVideoFlowEdit": "WanVideo FlowEdit",
"WanVideoControlEmbeds": "WanVideo Control Embeds",
"WanVideoSLG": "WanVideo SLG",
"WanVideoLoopArgs": "WanVideo Loop Args",
@@ -2269,5 +2344,9 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddPusaNoise": "WanVideo Add Pusa Noise",
"WanVideoAnimateEmbeds": "WanVideo Animate Embeds",
"WanVideoAddLucyEditLatents": "WanVideo Add LucyEdit Latents",
"WanVideoSchedulerSA_ODE": "WanVideo Scheduler SA-ODE",
"WanVideoAddBindweaveEmbeds": "WanVideo Add Bindweave Embeds",
"WanVideoUniLumosEmbeds": "WanVideo UniLumos Embeds",
"WanVideoAddTTMLatents": "WanVideo Add TTMLatents",
"WanVideoAddStoryMemLatents": "WanVideo Add StoryMem Latents",
"WanVideoSVIProEmbeds": "WanVideo SVIPro Embeds",
}
+353 -135
View File
File diff suppressed because it is too large Load Diff
+723 -1004
View File
File diff suppressed because it is too large Load Diff
+190 -69
View File
@@ -1,6 +1,9 @@
import torch
import torch.nn.functional as F
import numpy as np
from comfy.utils import common_upscale
from comfy import model_management
from tqdm import tqdm
from .utils import log
from einops import rearrange
@@ -12,6 +15,9 @@ except:
VAE_STRIDE = (4, 8, 8)
PATCH_SIZE = (1, 2, 2)
main_device = model_management.get_torch_device()
offload_device = model_management.unet_offload_device()
class WanVideoImageResizeToClosest:
@classmethod
def INPUT_TYPES(s):
@@ -30,7 +36,7 @@ class WanVideoImageResizeToClosest:
DESCRIPTION = "Resizes image to the closest supported resolution based on aspect ratio and max pixels, according to the original code"
def process(self, image, generation_width, generation_height, aspect_ratio_preservation ):
H, W = image.shape[1], image.shape[2]
max_area = generation_width * generation_height
@@ -42,7 +48,7 @@ class WanVideoImageResizeToClosest:
aspect_ratio = generation_height / generation_width
if aspect_ratio_preservation == "crop_to_new":
crop = "center"
lat_h = round(
np.sqrt(max_area * aspect_ratio) // VAE_STRIDE[1] //
PATCH_SIZE[1] * PATCH_SIZE[1])
@@ -130,27 +136,27 @@ class WanVideoVACEStartToEndFrame:
# Convert negative end_index to positive
if end_index < 0:
end_index = num_frames + end_index
# Create output batch with empty frames
out_batch = torch.ones((num_frames, H, W, 3), device=device) * empty_frame_level
# Create mask tensor with proper dimensions
masks = torch.ones((num_frames, H, W), device=device)
# Pre-process all images at once to avoid redundant work
if end_image is not None and (end_image.shape[1] != H or end_image.shape[2] != W):
end_image = common_upscale(end_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1)
if control_images is not None and (control_images.shape[1] != H or control_images.shape[2] != W):
control_images = common_upscale(control_images.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1)
# Place start image at start_index
if start_image is not None:
frames_to_copy = min(start_image.shape[0], num_frames - start_index)
if frames_to_copy > 0:
out_batch[start_index:start_index + frames_to_copy] = start_image[:frames_to_copy]
masks[start_index:start_index + frames_to_copy] = 0
# Place end image at end_index
if end_image is not None:
# Calculate where to start placing end images
@@ -158,28 +164,28 @@ class WanVideoVACEStartToEndFrame:
if end_start < 0: # Handle case where end images won't all fit
end_image = end_image[abs(end_start):]
end_start = 0
frames_to_copy = min(end_image.shape[0], num_frames - end_start)
if frames_to_copy > 0:
out_batch[end_start:end_start + frames_to_copy] = end_image[:frames_to_copy]
masks[end_start:end_start + frames_to_copy] = 0
# Apply control images to remaining frames that don't have start or end images
if control_images is not None:
# Create a mask of frames that are still empty (mask == 1)
empty_frames = masks.sum(dim=(1, 2)) > 0.5 * H * W
if empty_frames.any():
# Only apply control images where they exist
control_length = control_images.shape[0]
for frame_idx in range(num_frames):
if empty_frames[frame_idx] and frame_idx < control_length:
out_batch[frame_idx] = control_images[frame_idx]
# Apply inpaint mask if provided
if inpaint_mask is not None:
inpaint_mask = common_upscale(inpaint_mask.unsqueeze(1), W, H, "nearest-exact", "disabled").squeeze(1).to(device)
# Handle different mask lengths efficiently
if inpaint_mask.shape[0] > num_frames:
inpaint_mask = inpaint_mask[:num_frames]
@@ -221,31 +227,31 @@ class CreateCFGScheduleFloatList:
cfg_list = [1.0] * steps
start_idx = min(int(steps * start_percent), steps - 1)
end_idx = min(int(steps * end_percent), steps - 1)
for i in range(start_idx, end_idx + 1):
if i >= steps:
break
if end_idx == start_idx:
t = 0
else:
t = (i - start_idx) / (end_idx - start_idx)
if interpolation == "linear":
factor = t
elif interpolation == "ease_in":
factor = t * t
elif interpolation == "ease_out":
factor = t * (2 - t)
cfg_list[i] = round(cfg_scale_start + factor * (cfg_scale_end - cfg_scale_start), 2)
# If start_percent > 0, always include the first step
if start_percent > 0:
cfg_list[0] = 1.0
if unique_id and PromptServer is not None:
try:
try:
PromptServer.instance.send_progress_text(
f"{cfg_list}",
unique_id
@@ -254,7 +260,7 @@ class CreateCFGScheduleFloatList:
pass
return (cfg_list,)
class CreateScheduleFloatList:
@classmethod
def INPUT_TYPES(s):
@@ -284,16 +290,16 @@ class CreateScheduleFloatList:
cfg_list = [default_value] * steps
start_idx = min(int(steps * start_percent), steps - 1)
end_idx = min(int(steps * end_percent), steps - 1)
for i in range(start_idx, end_idx + 1):
if i >= steps:
break
if end_idx == start_idx:
t = 0
else:
t = (i - start_idx) / (end_idx - start_idx)
if interpolation == "linear":
factor = t
elif interpolation == "ease_in":
@@ -308,7 +314,7 @@ class CreateScheduleFloatList:
cfg_list[0] = default_value
if unique_id and PromptServer is not None:
try:
try:
PromptServer.instance.send_progress_text(
f"{cfg_list}",
unique_id
@@ -317,7 +323,7 @@ class CreateScheduleFloatList:
pass
return (cfg_list,)
class DummyComfyWanModelObject:
@classmethod
@@ -343,7 +349,7 @@ class DummyComfyWanModelObject:
return model_sampling
return None
return (DummyModel(),)
class WanVideoLatentReScale:
@classmethod
def INPUT_TYPES(s):
@@ -400,7 +406,7 @@ class WanVideoLatentReScale:
samples["samples"] = latents
return (samples,)
class WanVideoSigmaToStep:
@classmethod
def INPUT_TYPES(s):
@@ -417,7 +423,7 @@ class WanVideoSigmaToStep:
def convert(self, sigma):
return (sigma,)
class NormalizeAudioLoudness:
@classmethod
def INPUT_TYPES(s):
@@ -432,11 +438,11 @@ class NormalizeAudioLoudness:
FUNCTION = "normalize"
CATEGORY = "WanVideoWrapper"
def normalize(self, audio, lufs):
def normalize(self, audio, lufs):
audio_input = audio["waveform"]
sample_rate = audio["sample_rate"]
if audio_input.dim() == 3:
audio_input = audio_input.squeeze(0)
audio_input = audio_input.squeeze(0)
audio_input_np = audio_input.detach().transpose(0, 1).numpy().astype(np.float32)
audio_input_np = np.ascontiguousarray(audio_input_np)
normalized_audio = self.loudness_norm(audio_input_np, sr=sample_rate, lufs=lufs)
@@ -444,7 +450,7 @@ class NormalizeAudioLoudness:
out_audio = {"waveform": torch.from_numpy(normalized_audio).transpose(0, 1).unsqueeze(0).float(), "sample_rate": sample_rate}
return (out_audio, )
def loudness_norm(self, audio_array, sr=16000, lufs=-23):
try:
import pyloudnorm
@@ -456,7 +462,7 @@ class NormalizeAudioLoudness:
return audio_array
normalized_audio = pyloudnorm.normalize.loudness(audio_array, loudness, lufs)
return normalized_audio
class WanVideoPassImagesFromSamples:
@classmethod
def INPUT_TYPES(s):
@@ -500,15 +506,15 @@ class FaceMaskFromPoseKeypoints:
for i, pose_frame in enumerate(pose_frames):
selected_idx, prev_center = self.select_closest_person(pose_frame, person_index if i == 0 else prev_center)
np_frames.append(self.draw_kps(pose_frame, selected_idx))
if not np_frames:
# Handle case where no frames were processed
log.warning("No valid pose frames found, returning empty mask")
return (torch.zeros((1, 64, 64), dtype=torch.float32),)
np_frames = np.stack(np_frames, axis=0)
tensor = torch.from_numpy(np_frames).float() / 255.
print("tensor.shape:", tensor.shape)
log.info(f"tensor.shape: {tensor.shape}")
tensor = tensor[:, :, :, 0]
return (tensor,)
@@ -516,41 +522,41 @@ class FaceMaskFromPoseKeypoints:
people = pose_frame["people"]
if not people:
return -1, None
centers = []
valid_people_indices = []
for idx, person in enumerate(people):
# Check if face keypoints exist and are valid
if "face_keypoints_2d" not in person or not person["face_keypoints_2d"]:
continue
kps = np.array(person["face_keypoints_2d"])
if len(kps) == 0:
continue
n = len(kps) // 3
if n == 0:
continue
facial_kps = rearrange(kps, "(n c) -> n c", n=n, c=3)[:, :2]
# Check if we have valid coordinates (not all zeros)
if np.all(facial_kps == 0):
continue
center = facial_kps.mean(axis=0)
# Check if center is valid (not NaN or infinite)
if np.isnan(center).any() or np.isinf(center).any():
continue
centers.append(center)
valid_people_indices.append(idx)
if not centers:
return -1, None
if isinstance(prev_center_or_index, (int, np.integer)):
# First frame: use person_index, but map to valid people
if 0 <= prev_center_or_index < len(valid_people_indices):
@@ -582,58 +588,58 @@ class FaceMaskFromPoseKeypoints:
width, height = pose_frame["canvas_width"], pose_frame["canvas_height"]
canvas = np.zeros((height, width, 3), dtype=np.uint8)
people = pose_frame["people"]
if person_index < 0 or person_index >= len(people):
return canvas # Out of bounds, return blank
person = people[person_index]
# Check if face keypoints exist and are valid
if "face_keypoints_2d" not in person or not person["face_keypoints_2d"]:
return canvas # No face keypoints, return blank
face_kps_data = person["face_keypoints_2d"]
if len(face_kps_data) == 0:
return canvas # Empty keypoints, return blank
n = len(face_kps_data) // 3
if n < 17: # Need at least 17 points for outer contour
return canvas # Not enough keypoints, return blank
facial_kps = rearrange(np.array(face_kps_data), "(n c) -> n c", n=n, c=3)[:, :2]
# Check if we have valid coordinates (not all zeros)
if np.all(facial_kps == 0):
return canvas # All keypoints are zero, return blank
# Check for NaN or infinite values
if np.isnan(facial_kps).any() or np.isinf(facial_kps).any():
return canvas # Invalid coordinates, return blank
# Check for negative coordinates or coordinates that would create streaks
if np.any(facial_kps < 0):
return canvas # Negative coordinates, likely bad detection
# Check if coordinates are reasonable (not too close to edges which might indicate bad detection)
min_margin = 5 # Minimum distance from edges
if (np.any(facial_kps[:, 0] < min_margin) or
np.any(facial_kps[:, 1] < min_margin) or
np.any(facial_kps[:, 0] > width - min_margin) or
if (np.any(facial_kps[:, 0] < min_margin) or
np.any(facial_kps[:, 1] < min_margin) or
np.any(facial_kps[:, 0] > width - min_margin) or
np.any(facial_kps[:, 1] > height - min_margin)):
# Check if this looks like a streak to corner (many points near 0,0)
corner_points = np.sum((facial_kps[:, 0] < min_margin) & (facial_kps[:, 1] < min_margin))
if corner_points > 3: # Too many points near corner, likely bad detection
return canvas
facial_kps = facial_kps.astype(np.int32)
# Ensure coordinates are within canvas bounds
facial_kps[:, 0] = np.clip(facial_kps[:, 0], 0, width - 1)
facial_kps[:, 1] = np.clip(facial_kps[:, 1], 0, height - 1)
part_color = (255, 255, 255)
outer_contour = facial_kps[:17]
# Additional validation for the contour before drawing
# Check if contour points are too spread out (indicating bad detection)
if len(outer_contour) >= 3:
@@ -642,11 +648,11 @@ class FaceMaskFromPoseKeypoints:
max_x, max_y = np.max(outer_contour, axis=0)
contour_width = max_x - min_x
contour_height = max_y - min_y
# If contour spans more than 80% of canvas, likely bad detection
if (contour_width > 0.8 * width or contour_height > 0.8 * height):
return canvas
# Check if we have a valid contour (at least 3 unique points)
unique_points = np.unique(outer_contour, axis=0)
if len(unique_points) >= 3:
@@ -654,13 +660,124 @@ class FaceMaskFromPoseKeypoints:
# Calculate area to see if it's too large or too small
contour_area = cv2.contourArea(outer_contour)
canvas_area = width * height
# If contour is less than 0.1% or more than 50% of canvas, skip
if 0.001 * canvas_area <= contour_area <= 0.5 * canvas_area:
cv2.fillPoly(canvas, pts=[outer_contour], color=part_color)
return canvas
class DrawGaussianNoiseOnImage:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"image": ("IMAGE", ),
"mask": ("MASK", ),
},
"optional": {
"device": (["cpu", "gpu"], {"default": "cpu", "tooltip": "Device to use for processing"}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
}
}
RETURN_TYPES = ("IMAGE", )
RETURN_NAMES = ("images",)
FUNCTION = "apply"
CATEGORY = "KJNodes/masking"
DESCRIPTION = "Fills the background (masked area) with Gaussian noise sampled using the mean and variance of the subject (unmasked) region."
def apply(self, image, mask, device="cpu", seed=0):
B, H, W, C = image.shape
BM, HM, WM = mask.shape
processing_device = main_device if device == "gpu" else torch.device("cpu")
in_masks = mask.clone().to(processing_device)
in_images = image.clone().to(processing_device)
# Resize mask to match image dimensions
if HM != H or WM != W:
in_masks = F.interpolate(mask.unsqueeze(1), size=(H, W), mode='nearest-exact').squeeze(1)
# Match batch sizes
if B > BM:
in_masks = in_masks.repeat((B + BM - 1) // BM, 1, 1)[:B]
elif BM > B:
in_masks = in_masks[:B]
output_images = []
# Set random seed for reproducibility
generator = torch.Generator(device=processing_device).manual_seed(seed)
for i in tqdm(range(B), desc="DrawGaussianNoiseOnImage batch"):
curr_mask = in_masks[i]
img_idx = min(i, B - 1)
curr_image = in_images[img_idx]
# Expand mask to 3 channels
mask_expanded = curr_mask.unsqueeze(-1).expand(-1, -1, 3)
# Calculate mean and std per channel from the subject region (where mask is 1)
subject_mask = mask_expanded > 0.5
# Initialize noise tensor
noise = torch.zeros_like(curr_image)
for c in range(C):
channel = curr_image[:, :, c]
channel_mask = subject_mask[:, :, c]
if channel_mask.sum() > 0:
# Get subject pixels
subject_pixels = channel[channel_mask]
# Calculate statistics
mean = subject_pixels.mean()
std = subject_pixels.std()
# Generate Gaussian noise for this channel
noise[:, :, c] = torch.normal(mean=mean.item(), std=std.item(),
size=(H, W), generator=generator,
device=processing_device)
# Clamp noise to valid range
noise = torch.clamp(noise, 0.0, 1.0)
# Apply: keep subject, fill background with noise
masked_image = curr_image * mask_expanded + noise * (1 - mask_expanded)
output_images.append(masked_image)
# If no masks were processed, return empty tensor
if not output_images:
return (torch.zeros((0, H, W, 3), dtype=image.dtype),)
out_rgb = torch.stack(output_images, dim=0).cpu()
return (out_rgb, )
class WanVideoPreviewEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
}
}
RETURN_TYPES = ("LATENT", "MASK")
RETURN_NAMES = ("image_embeds", "mask",)
FUNCTION = "get"
CATEGORY = "WanVideoWrapper"
def get(self, embeds):
latents = embeds.get("image_embeds", None)
mask = embeds.get("mask", None)
if mask is not None:
mask = mask[0].float().cpu()
return ({"samples": latents.unsqueeze(0)}, mask)
NODE_CLASS_MAPPINGS = {
"WanVideoImageResizeToClosest": WanVideoImageResizeToClosest,
"WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame,
@@ -673,6 +790,8 @@ NODE_CLASS_MAPPINGS = {
"NormalizeAudioLoudness": NormalizeAudioLoudness,
"WanVideoPassImagesFromSamples": WanVideoPassImagesFromSamples,
"FaceMaskFromPoseKeypoints": FaceMaskFromPoseKeypoints,
"DrawGaussianNoiseOnImage": DrawGaussianNoiseOnImage,
"WanVideoPreviewEmbeds": WanVideoPreviewEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest",
@@ -686,4 +805,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"NormalizeAudioLoudness": "Normalize Audio Loudness",
"WanVideoPassImagesFromSamples": "WanVideo Pass Images From Samples",
"FaceMaskFromPoseKeypoints": "Face Mask From Pose Keypoints",
}
"DrawGaussianNoiseOnImage": "Draw Gaussian Noise On Image",
"WanVideoPreviewEmbeds": "WanVideo Preview Embeds",
}
+440
View File
@@ -0,0 +1,440 @@
import torch
from torch import nn
from torch.nn import functional as F
from einops import rearrange
import numpy as np
from typing import Tuple
from .unet_causal_3d_blocks import get_down_block3d, CausalConv3d
class ControlNetCausalConditioningEmbedding(nn.Module):
def __init__(self, conditioning_embedding_channels: int, conditioning_channels: int = 3, block_out_channels: Tuple[int, ...] = (16, 32, 96, 256)):
super().__init__()
self.conv_in = CausalConv3d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1)
self.blocks = nn.ModuleList([])
for i in range(len(block_out_channels) - 1):
channel_in = block_out_channels[i]
channel_out = block_out_channels[i + 1]
self.blocks.append(nn.Conv2d(channel_in, channel_in, kernel_size=3, padding=1))
self.blocks.append(nn.Conv2d(channel_in, channel_out, kernel_size=3, padding=1, stride=2))
self.conv_out = nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1)
def forward(self, conditioning):
embedding = self.conv_in(conditioning)
embedding = F.silu(embedding)
for block in self.blocks:
embedding = block(embedding)
embedding = F.silu(embedding)
embedding = self.conv_out(embedding)
return embedding
class MiniHunyuanEncoder(nn.Module):
'''
a direct copy of hunyuan encoder
'''
def __init__(
self,
in_channels = 3,
out_channels = 3,
down_block_types = ['DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D'],
block_out_channels = [128, 256, 512, 512],
layers_per_block = 2,
norm_num_groups = 32,
act_fn: str = "silu",
time_compression_ratio: int = 4,
spatial_compression_ratio: int = 8,
):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(
in_channels, block_out_channels[0], kernel_size=3, stride=1)
self.mid_block = None
self.down_blocks = nn.ModuleList([])
# down
output_channel = block_out_channels[0]
for i, down_block_type in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_downsample_layers = int(
np.log2(spatial_compression_ratio))
num_time_downsample_layers = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i >= (
len(block_out_channels) - 1 - num_time_downsample_layers) and not is_final_block)
elif time_compression_ratio == 8:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i < num_time_downsample_layers)
else:
raise ValueError(
f"Unsupported time_compression_ratio: {time_compression_ratio}")
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
downsample_stride_T = (2, ) if add_time_downsample else (1, )
downsample_stride = tuple(
downsample_stride_T + downsample_stride_HW)
down_block = get_down_block3d(
down_block_type,
num_layers=self.layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=bool(
add_spatial_downsample or add_time_downsample),
downsample_stride=downsample_stride,
resnet_eps=1e-6,
downsample_padding=0,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
)
self.down_blocks.append(down_block)
self.conv_out = CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
def forward(self, sample):
assert len(sample.shape) == 5, "The input tensor should have 5 dimensions"
sample = self.conv_in(sample)
# down
for down_block in self.down_blocks:
sample = down_block(sample)
sample = self.conv_out(sample)
return sample
class ControlNetConditioningEmbedding(nn.Module):
"""
Quoting from https://arxiv.org/abs/2302.05543: "Stable Diffusion uses a pre-processing method similar to VQ-GAN
[11] to convert the entire dataset of 512 × 512 images into smaller 64 × 64 “latent images” for stabilized
training. This requires ControlNets to convert image-based conditions to 64 × 64 feature space to match the
convolution size. We use a tiny network E(·) of four convolution layers with 4 × 4 kernels and 2 × 2 strides
(activated by ReLU, channels are 16, 32, 64, 128, initialized with Gaussian weights, trained jointly with the full
model) to encode image-space conditions ... into feature maps ..."
"""
def __init__(
self,
conditioning_embedding_channels: int,
conditioning_channels: int = 3,
block_out_channels: Tuple[int, ...] = (16, 32, 96, 256),
):
super().__init__()
self.conv_in = nn.Conv2d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1)
self.blocks = nn.ModuleList([])
for i in range(len(block_out_channels) - 1):
channel_in = block_out_channels[i]
channel_out = block_out_channels[i + 1]
self.blocks.append(nn.Conv2d(channel_in, channel_in, kernel_size=3, padding=1))
self.blocks.append(nn.Conv2d(channel_in, channel_out, kernel_size=3, padding=1, stride=2))
self.conv_out = nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1)
def forward(self, conditioning):
embedding = self.conv_in(conditioning)
embedding = F.silu(embedding)
for block in self.blocks:
embedding = block(embedding)
embedding = F.silu(embedding)
embedding = self.conv_out(embedding)
return embedding
class InflatedGroupNorm(nn.GroupNorm):
def forward(self, x):
video_length = x.shape[2]
x = rearrange(x, "b c f h w -> (b f) c h w")
x = super().forward(x)
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
return x
class InflatedConv3d(nn.Conv2d):
def forward(self, x):
video_length = x.shape[2]
x = rearrange(x, "b c f h w -> (b f) c h w")
x = super().forward(x)
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
return x
class ResnetBlockInflated(nn.Module):
def __init__(self, *, in_channels, out_channels=None, dropout=0.0, groups=32, groups_out=None, pre_norm=True, eps=1e-6, non_linearity="swish", output_scale_factor=1.0):
super().__init__()
self.pre_norm = pre_norm
self.pre_norm = True
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.output_scale_factor = output_scale_factor
if groups_out is None:
groups_out = groups
self.norm1 = InflatedGroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
self.conv1 = InflatedConv3d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
self.norm2 = InflatedGroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
self.dropout = torch.nn.Dropout(dropout)
self.conv2 = InflatedConv3d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
if non_linearity == "swish":
self.nonlinearity = lambda x: F.silu(x)
elif non_linearity == "silu":
self.nonlinearity = nn.SiLU()
def forward(self, input_tensor, temb):
if temb is not None:
print("Warning: temb is None in ResnetBlockInflated")
hidden_states = input_tensor
hidden_states = self.norm1(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.conv1(hidden_states)
if temb is not None:
hidden_states = hidden_states + temb
hidden_states = self.norm2(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.dropout(hidden_states)
hidden_states = self.conv2(hidden_states)
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
return output_tensor
class DownEncoderBlockInflated(nn.Module):
def __init__(self, *, num_layers: int, in_channels: int, out_channels: int, add_downsample: bool, downsample_stride: tuple = (1, 2, 2),
resnet_eps: float = 1e-6, resnet_act_fn: str = "silu", resnet_groups: int = 32):
super().__init__()
self.resnets = nn.ModuleList([ResnetBlockInflated(
in_channels=in_channels if i == 0 else out_channels,
out_channels=out_channels,
eps=resnet_eps,
non_linearity=resnet_act_fn,
groups=resnet_groups,
) for i in range(num_layers)])
self.downsamplers = nn.ModuleList()
if add_downsample:
self.downsamplers.append(
InflatedConv3d(
out_channels,
out_channels,
kernel_size=3,
stride=2,
padding=1,
)
)
self.down_stride = downsample_stride
else:
self.down_stride = (1, 1, 1)
def forward(self, x, temb=None):
for resnet in self.resnets:
x = resnet(x, temb)
for down in self.downsamplers:
x = down(x)
return x
class SFT(nn.Module): # 2D SFT
def __init__(
self, in_channels, out_channels, intermediate_channels=128, groups=32, eps=1e-6):
super().__init__()
self.out_channels = out_channels
self.norm = InflatedGroupNorm(groups, out_channels, eps, affine=True)
self.mlp_shared = nn.Sequential(InflatedConv3d(in_channels, intermediate_channels, kernel_size=3, stride=1, padding=1), nn.SiLU())
self.mlp_gamma = InflatedConv3d(intermediate_channels, out_channels, kernel_size=3, stride=1, padding=1)
self.mlp_beta = InflatedConv3d(intermediate_channels, out_channels, kernel_size=3, stride=1, padding=1)
def forward(self, hidden_state, condition):
"""
hidden_state : (B, Cout, T, H, W)
condition : (B, Cin, 1, H, W)
"""
hidden_state = self.norm(hidden_state) #2D SFT 2D Norm
actv = self.mlp_shared(condition)
gamma = self.mlp_gamma(actv)
beta = self.mlp_beta(actv)
return torch.addcmul(beta, hidden_state, 1 + gamma)
class MiniEncoder2D(nn.Module):
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
down_block_types: list = (
"DownEncoderBlockInflated",
"DownEncoderBlockInflated",
"DownEncoderBlockInflated",
"DownEncoderBlockInflated",
),
block_out_channels: list = (128, 256, 512, 512),
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
spatial_compression_ratio: int = 8,
):
super().__init__()
# -------------------------------------------------------------------
# conv in
# -------------------------------------------------------------------
self.conv_in = InflatedConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1, padding=1)
self.down_blocks = nn.ModuleList()
output_channel = block_out_channels[0]
num_spatial_down_layers = int(np.log2(spatial_compression_ratio))
for i, block_type in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
# is_final_block = i == len(block_out_channels) - 1
add_spatial_downsample = bool(i < num_spatial_down_layers)
downsample_stride = (1, 2, 2) if add_spatial_downsample else (1, 1, 1)
down_block = DownEncoderBlockInflated(
num_layers=layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=add_spatial_downsample,
downsample_stride=downsample_stride,
resnet_eps=1e-6,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
)
self.down_blocks.append(down_block)
self.conv_out = InflatedConv3d(output_channel, out_channels, kernel_size=3, stride=1, padding=1)
def forward(self, x):
# (B,C,1,H,W)
x = self.conv_in(x)
for block in self.down_blocks:
x = block(x)
return self.conv_out(x)
class Driven_Ref_PoseEncoder(nn.Module):
def __init__(
self, in_channels = 3, out_channels = 3,
down_block_types = ['DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D'],
block_out_channels = [128, 256, 512, 512], layers_per_block = 2, norm_num_groups = 32,
act_fn: str = "silu", time_compression_ratio: int = 4, spatial_compression_ratio: int = 8,
):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1)
self.mid_block = None
self.down_blocks = nn.ModuleList([])
# down
output_channel = block_out_channels[0]
for i, down_block_type in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_downsample_layers = int(
np.log2(spatial_compression_ratio))
num_time_downsample_layers = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i >= (
len(block_out_channels) - 1 - num_time_downsample_layers) and not is_final_block)
elif time_compression_ratio == 8:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i < num_time_downsample_layers)
else:
raise ValueError(
f"Unsupported time_compression_ratio: {time_compression_ratio}")
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
downsample_stride_T = (2, ) if add_time_downsample else (1, )
downsample_stride = tuple(
downsample_stride_T + downsample_stride_HW)
down_block = get_down_block3d(
down_block_type,
num_layers=self.layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=bool(
add_spatial_downsample or add_time_downsample),
downsample_stride=downsample_stride,
resnet_eps=1e-6,
downsample_padding=0,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
attention_head_dim=output_channel,
)
self.down_blocks.append(down_block)
self.conv_out = CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
self.ref_pose_encoder = MiniEncoder2D(
in_channels = in_channels,
out_channels = out_channels,
block_out_channels = block_out_channels,
norm_num_groups = norm_num_groups,
layers_per_block = layers_per_block,
spatial_compression_ratio = spatial_compression_ratio,
)
self.sft_layers = nn.ModuleList()
for i, ch in enumerate(block_out_channels):
if i == 0: # 0 层 (H/2,W/2) 不做 SFT
self.sft_layers.append(None)
else: # H/4、H/8、H/16 做 SFT
self.sft_layers.append(
SFT(
in_channels=ch,
out_channels=ch,
intermediate_channels=max(8, ch // 2),
groups=norm_num_groups,
)
)
def forward(self, driven_pose, ref_pose):
# driven_pose b c t h w
# ref_pose b c 1 h w
ref_pose_cond, ref_feats = self.ref_pose_encoder(ref_pose)
x = self.conv_in(driven_pose)
for i, down_block in enumerate(self.down_blocks):
x = down_block(x)
if self.sft_layers[i] is not None:
cond_feat = ref_feats[i]
x = self.sft_layers[i](x, cond_feat)
driven_pose_cond = self.conv_out(x)
return driven_pose_cond, ref_pose_cond
+152
View File
@@ -0,0 +1,152 @@
import torch
from ..utils import log
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoAddOneToAllReferenceEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"vae": ("WANVAE", {"tooltip": "VAE model"}),
"ref_image": ("IMAGE",),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
},
"optional": {
"ref_mask": ("MASK",),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, vae, ref_image, strength, start_percent, end_percent, ref_mask=None):
updated = dict(embeds)
ref_latent = ref_latent_empty = None
vae.to(device)
ref_image_in = (ref_image[..., :3].permute(3, 0, 1, 2) * 2 - 1).to(device, vae.dtype)
ref_latent = vae.encode([ref_image_in], device, tiled=False)
ref_mask_in = None
if ref_mask is not None:
ref_mask_in = (ref_mask.unsqueeze(0).repeat(3, 1, 1, 1) * 2 - 1.).to(device, vae.dtype)
else:
ref_mask_in = torch.zeros_like(ref_image_in)-1
ref_mask_latent = vae.encode([ref_mask_in], device, tiled=False)
if ref_mask is not None and not torch.all(ref_mask == 0):
ref_latent_empty = vae.encode([torch.zeros_like(ref_image_in)-1], device, tiled=False)
else:
ref_latent_empty = ref_mask_latent
vae.to(offload_device)
updated.setdefault("one_to_all_embeds", {})
updated["one_to_all_embeds"]["ref_latent_pos"] = torch.cat([ref_latent, ref_latent_empty], dim=1)
updated["one_to_all_embeds"]["ref_latent_neg"] = torch.cat([ref_latent_empty, ref_latent_empty], dim=1)
updated["one_to_all_embeds"]["ref_strength"] = strength
updated["one_to_all_embeds"]["ref_start_percent"] = start_percent
updated["one_to_all_embeds"]["ref_end_percent"] = end_percent
return (updated,)
class WanVideoAddOneToAllPoseEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"pose_images": ("IMAGE", {"tooltip": "Pose images for the entire video"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the pose control"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the pose control application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the pose control application"}),
},
"optional": {
"pose_prefix_image": ("IMAGE",),
"pose_cfg_scale": ("FLOAT", {"default": 1.5, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "CFG scale for the pose control, has no effect if main cfg scale is 1.0"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, pose_images, strength, pose_prefix_image=None, start_percent=0.0, end_percent=1.0, pose_cfg_scale=1.5):
updated = dict(embeds)
updated.setdefault("one_to_all_embeds", {})
pose_images_in = pose_images[..., :3].unsqueeze(0).permute(0, 4, 1, 2, 3) * 2 - 1 # 1 B H W C -> B C 1 H W
updated["one_to_all_embeds"]["pose_images"] = pose_images_in
if pose_prefix_image is not None:
updated["one_to_all_embeds"]["pose_prefix_image"] = pose_prefix_image.unsqueeze(0).permute(0, 4, 1, 2, 3) * 2 - 1 # 1 B H W C -> B C 1 H W
else:
updated["one_to_all_embeds"]["pose_prefix_image"] = pose_images_in[:, :, :1]
updated["one_to_all_embeds"]["controlnet_strength"] = strength
updated["one_to_all_embeds"]["controlnet_start_percent"] = start_percent
updated["one_to_all_embeds"]["controlnet_end_percent"] = end_percent
updated["one_to_all_embeds"]["pose_cfg_scale"] = pose_cfg_scale
return (updated,)
class WanVideoAddOneToAllExtendEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"prev_latents": ("LATENT", {"tooltip": "Previous latents to be used to continue generation"}),
"window_size": ("INT", {"default": 81, "min": 1, "max": 256, "step": 1, "tooltip": "Number of new frames to generate" }),
"overlap": ("INT", {"default": 5, "min": 0, "max": 64, "step": 1, "tooltip": "Number of overlapping frames between previous and new frames" }),
"frames_processed": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1, "tooltip": "Number of frames already processed in the video" }),
"if_not_enough_frames": (["pad_with_last", "error"], {"default": "pad_with_last", "tooltip": "What to do if there are not enough frames in pose_images for the window"}),
},
"optional": {
"pose_images": ("IMAGE", {"tooltip": "Pose images for the entire video"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "IMAGE",)
RETURN_NAMES = ("image_embeds", "pose_slice",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, prev_latents, if_not_enough_frames, window_size=81, overlap=5, frames_processed=0, pose_images=None):
updated = dict(embeds)
updated.setdefault("one_to_all_embeds", {})
updated["one_to_all_embeds"]["prev_latents"] = prev_latents["samples"][0]
if pose_images is not None:
pose_images_in = pose_images.clone()[..., :3]
start = max(0, frames_processed - overlap)
end = start + window_size
log.info(f"Extracting pose images from {start} to {end}")
if start >= pose_images_in.shape[0]:
raise ValueError(f"start index {start} exceeds pose images length {pose_images_in.shape[0]}")
if end > pose_images_in.shape[0]:
if if_not_enough_frames == "pad_with_last":
padding_needed = end - pose_images_in.shape[0]
pose_images_in = torch.cat([pose_images_in, pose_images_in[-1:].repeat(padding_needed, 1, 1, 1)], dim=0)
log.info(f"Not enough frames, padding with {padding_needed} frames to reach {end} total frames")
else:
raise ValueError(f"end index {end} exceeds pose images length {pose_images.shape[0]}")
pose_slice = pose_images_in[start:end]
else:
pose_slice = torch.zeros((1, 64, 64, 3))
return (updated, pose_slice)
NODE_CLASS_MAPPINGS = {
"WanVideoAddOneToAllReferenceEmbeds": WanVideoAddOneToAllReferenceEmbeds,
"WanVideoAddOneToAllPoseEmbeds": WanVideoAddOneToAllPoseEmbeds,
"WanVideoAddOneToAllExtendEmbeds": WanVideoAddOneToAllExtendEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddOneToAllReferenceEmbeds": "WanVideo Add OneToAll Reference Embeds",
"WanVideoAddOneToAllPoseEmbeds": "WanVideo Add OneToAll Pose Embeds",
"WanVideoAddOneToAllExtendEmbeds": "WanVideo Add OneToAll Extend Embeds",
}
+220
View File
@@ -0,0 +1,220 @@
from typing import Dict, Union
import torch
import torch.nn as nn
from ..wanvideo.modules.model import WanLayerNorm, WanSelfAttention, EmbedND_RifleX, sinusoidal_embedding_1d, apply_rotary_emb_split, apply_rope_comfy1
class WanAttentionBlock(nn.Module):
def __init__(self, in_features, out_features, ffn_dim, ffn2_dim, num_heads, qk_norm=True, cross_attn_norm=False, eps=1e-6, attention_mode="sdpa", rope_func="comfy", rms_norm_function="default"):
super().__init__()
self.dim = out_features
self.ffn_dim = ffn_dim
self.num_heads = num_heads
self.head_dim = out_features // num_heads
self.qk_norm = qk_norm
self.cross_attn_norm = cross_attn_norm
self.eps = eps
self.attention_mode = attention_mode
self.rope_func = rope_func
# layers
self.norm1 = WanLayerNorm(self.dim, eps)
self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm, eps, self.attention_mode, rms_norm_function=rms_norm_function, head_norm=False)
self.norm2 = WanLayerNorm(self.dim, eps)
self.ffn = nn.Sequential(nn.Linear(in_features, ffn_dim), nn.GELU(approximate='tanh'), nn.Linear(ffn2_dim, out_features))
self.modulation = nn.Parameter(torch.randn(1, 6, out_features) / in_features**0.5)
def get_mod(self, e, modulation):
if e.dim() == 3:
if e.shape[-1] == 512:
e = self.modulation(e)
return e.unsqueeze(2).chunk(6, dim=-1)
return (modulation + e).chunk(6, dim=1) # 1, 6, dim
elif e.dim() == 4:
e_mod = modulation.unsqueeze(2) + e
return [ei.squeeze(1) for ei in e_mod.unbind(dim=1)]
def modulate(self, norm_x, shift_msa, scale_msa):
return torch.addcmul(shift_msa, norm_x, 1 + scale_msa)
def ffn_chunked(self, mod_x, num_chunks=4):
seq_len = mod_x.shape[1]
if seq_len <= 8192 or num_chunks <= 1:
return self.ffn(mod_x)
return torch.cat([self.ffn(chunk.contiguous()) for chunk in mod_x.chunk(num_chunks, dim=1)], dim=1)
#region attention forward
def forward(self, x, e, seq_lens, freqs, split_rope=True, e_tr=None, tr_start=0, tr_num=0):
use_token_replace = False
if e_tr is not None and tr_num > 0:
tr_shift_msa, tr_scale_msa, tr_gate_msa, tr_shift_mlp, tr_scale_mlp, tr_gate_mlp = self.get_mod(e_tr.to(x.device), self.modulation)
use_token_replace = True
tr_start = tr_start or 0
tr_end = tr_start + (tr_num or 0)
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device), self.modulation)
del e
input_dtype = x.dtype
if use_token_replace:
norm_x = self.norm1(x.to(shift_msa.dtype))
input_x = torch.cat([
torch.addcmul(shift_msa, norm_x[:, :tr_start], 1 + scale_msa), # before replace → T
torch.addcmul(tr_shift_msa, norm_x[:, tr_start:tr_end], 1 + tr_scale_msa), # replace segment → t=0
torch.addcmul(shift_msa, norm_x[:, tr_end:], 1 + scale_msa) # after replace → T
], dim=1).to(input_dtype)
else:
input_x = self.modulate(self.norm1(x.to(shift_msa.dtype)), shift_msa, scale_msa).to(input_dtype)
del shift_msa, scale_msa
b, s, n, d = *x.shape[:2], self.self_attn.num_heads, self.self_attn.head_dim
h_dim = w_dim = 2 * (self.head_dim // 6)
t_dim = self.head_dim - h_dim - w_dim
q = self.self_attn.norm_q(self.self_attn.q(input_x)).to(self.self_attn.norm_q.weight.dtype).view(b, s, n, d)
if split_rope:
q = apply_rotary_emb_split(q, freqs, t_dim) # Apply split rotary embedding (only to H/W dimensions, leaving T unchanged)
else:
q = apply_rope_comfy1(q, freqs)
k = self.self_attn.norm_k(self.self_attn.k(input_x).to(self.self_attn.norm_k.weight.dtype)).to(input_x.dtype).view(b, s, n, d)
if split_rope:
k = apply_rotary_emb_split(k, freqs, t_dim)
else:
k = apply_rope_comfy1(k, freqs)
v = self.self_attn.v(input_x).view(b, s, n, d)
del input_x
y = self.self_attn.forward(q, k, v, seq_lens)
del q, k, v
if use_token_replace:
x = x + torch.cat([
y[:, :tr_start] * gate_msa,
y[:, tr_start:tr_end] * tr_gate_msa,
y[:, tr_end:] * gate_msa
], dim=1).to(input_dtype)
else:
x = x.addcmul(y, gate_msa)
del y, gate_msa
# ffn
if use_token_replace:
norm2_x = self.norm2(x.to(shift_mlp.dtype))
mod_x = torch.cat([
torch.addcmul(shift_mlp, norm2_x[:, :tr_start], 1 + scale_mlp),
torch.addcmul(tr_shift_mlp, norm2_x[:, tr_start:tr_end], 1 + tr_scale_mlp),
torch.addcmul(shift_mlp, norm2_x[:, tr_end:], 1 + scale_mlp)
], dim=1)
else:
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
del shift_mlp, scale_mlp
x_ffn = self.ffn_chunked(mod_x.to(input_dtype), num_chunks=1)
del mod_x
# gate_mlp
if use_token_replace:
x = x + torch.cat([
x_ffn[:, :tr_start] * gate_mlp,
x_ffn[:, tr_start:tr_end] * tr_gate_mlp,
x_ffn[:, tr_end:] * gate_mlp
], dim=1).to(input_dtype)
else:
x = x.addcmul(x_ffn.to(gate_mlp.dtype), gate_mlp).to(input_dtype)
del gate_mlp
return x
class WanRefextractor(nn.Module):
def __init__(self, patch_size=(1, 2, 2), in_dim=16, dim=5120, in_features=5120, out_features=5120, ffn_dim=8192, ffn2_dim=8192,
freq_dim=256, num_heads=16, num_layers=32, eps=1e-6,
qk_norm=True, cross_attn_norm=True,
attention_mode='sdpa', rope_func='comfy', rms_norm_function='default',
main_device=torch.device('cuda'), offload_device=torch.device('cpu'), dtype=torch.float16):
super().__init__()
self.patch_size = patch_size
self.freq_dim = freq_dim
self.dim = dim
self.main_device = main_device
self.base_dtype = dtype
self.attention_mode = attention_mode
self.patch_embedding = nn.Conv3d(in_dim, dim, kernel_size=patch_size, stride=patch_size)
self.time_embedding = nn.Sequential(nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
self.time_projection = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6))
self.blocks = nn.ModuleList([
WanAttentionBlock(in_features, out_features, ffn_dim, ffn2_dim, num_heads,
qk_norm, cross_attn_norm, eps, attention_mode="sdpa", rope_func=rope_func, rms_norm_function=rms_norm_function)
for i in range(num_layers)
])
self.ref_blocks = nn.ModuleList([])
for _ in range(len(self.blocks)+1):
self.ref_blocks.append(nn.Linear(in_features, out_features))
d = dim // num_heads
self.rope_embedder = EmbedND_RifleX(d,10000.0, [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)], num_frames=1, k=0)
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None):
patch_size = self.patch_size
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
h_len = ((h + (patch_size[1] // 2)) // patch_size[1])
w_len = ((w + (patch_size[2] // 2)) // patch_size[2])
if steps_t is None:
steps_t = t_len
if steps_h is None:
steps_h = h_len
if steps_w is None:
steps_w = w_len
img_ids = torch.zeros((steps_t, steps_h, steps_w, 3), device=device, dtype=dtype)
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1)
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(freq_offset, h_len - 1, steps=steps_h, device=device, dtype=dtype).reshape(1, -1, 1)
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(freq_offset, w_len - 1, steps=steps_w, device=device, dtype=dtype).reshape(1, 1, -1)
img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
freqs = self.rope_embedder(img_ids, ntk_alphas).movedim(1, 2)
return freqs
def forward(
self,
x: torch.Tensor,
timestep: torch.LongTensor,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
B, C, F, H, W = x.shape
freqs = self.rope_encode_comfy(F, H, W, device=x.device, dtype=x.dtype)
self.patch_embedding.to(self.main_device)
x = self.patch_embedding(x.float()).to(x.dtype).flatten(2).transpose(1, 2).to(self.base_dtype)
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.int32)
time_embed_dtype = self.time_embedding[0].weight.dtype
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
time_embed_dtype = self.base_dtype
e = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, timestep.flatten()).to(time_embed_dtype)) # b, dim
e0 = self.time_projection(e).unflatten(1, (6, self.dim)).to(self.base_dtype) # b, 6, dim
del e
# 4. Transformer blocks
block_samples = ()
for block in self.blocks:
block_samples = block_samples + (x, )
x = block(x, e0, seq_lens, freqs)
block_samples = block_samples + (x, )
ref_block_samples = ()
for block_sample, ref_block in zip(block_samples, self.ref_blocks):
block_sample = ref_block(block_sample)
ref_block_samples = ref_block_samples + (block_sample, )
return ref_block_samples, freqs
+144
View File
@@ -0,0 +1,144 @@
# Copyright 2024 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
#
# Modified from diffusers==0.29.2
#
# ==============================================================================
from typing import Optional
import torch
import torch.nn.functional as F
from torch import nn
import comfy.ops
ops = comfy.ops.disable_weight_init
def prepare_causal_attention_mask(n_frame: int, n_hw: int, dtype, device, batch_size: int = None):
seq_len = n_frame * n_hw
mask = torch.full((seq_len, seq_len), float(
"-inf"), dtype=dtype, device=device)
for i in range(seq_len):
i_frame = i // n_hw
mask[i, : (i_frame + 1) * n_hw] = 0
if batch_size is not None:
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
return mask
class CausalConv3d(nn.Module):
def __init__(self, chan_in, chan_out, kernel_size, stride = 1, dilation = 1, pad_mode='replicate', **kwargs):
super().__init__()
self.pad_mode = pad_mode
padding = (kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size - 1, 0) # W, H, T
self.time_causal_padding = padding
self.conv = ops.Conv3d(chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs)
def forward(self, x):
x = F.pad(x, self.time_causal_padding, mode=self.pad_mode)
return self.conv(x)
class DownsampleCausal3D(nn.Module):
def __init__(self, channels, use_conv=False, out_channels=None, padding=1, name="conv", kernel_size=3, bias=True, stride=2):
super().__init__()
self.channels, self.out_channels, self.use_conv, self.padding, self.name = channels, out_channels or channels, use_conv, padding, name
self.conv = CausalConv3d(self.channels, self.out_channels, kernel_size=kernel_size, stride=stride, bias=bias)
def forward(self, x, scale=1.0):
return self.conv(x)
class ResnetBlockCausal3D(nn.Module):
def __init__(self, *, in_channels: int, out_channels: Optional[int] = None, groups: int = 32, eps: float = 1e-6, conv_3d_out_channels: Optional[int] = None):
super().__init__()
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
self.norm2 = torch.nn.GroupNorm(num_groups=groups, num_channels=out_channels, eps=eps, affine=True)
self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3, stride=1)
conv_3d_out_channels = conv_3d_out_channels or out_channels
self.conv2 = CausalConv3d(out_channels, conv_3d_out_channels, kernel_size=3, stride=1)
def forward(self, input_tensor: torch.FloatTensor, temb: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
hidden_states = input_tensor
hidden_states = self.conv1(nn.SiLU()(self.norm1(hidden_states)))
if temb is not None:
hidden_states = hidden_states + temb
hidden_states = self.conv2(nn.SiLU()(self.norm2(hidden_states)))
return input_tensor + hidden_states
def get_down_block3d(down_block_type: str, num_layers: int, in_channels: int, out_channels: int,
add_downsample: bool, downsample_stride: int, resnet_eps: float, resnet_act_fn: str, resnet_groups: Optional[int] = None,
downsample_padding: Optional[int] = None, **kwargs):
down_block_type = down_block_type[7:] if down_block_type.startswith(
"UNetRes") else down_block_type
if down_block_type == "DownEncoderBlockCausal3D":
return DownEncoderBlockCausal3D(
num_layers=num_layers,
in_channels=in_channels,
out_channels=out_channels,
add_downsample=add_downsample,
downsample_stride=downsample_stride,
resnet_eps=resnet_eps,
resnet_act_fn=resnet_act_fn,
resnet_groups=resnet_groups,
downsample_padding=downsample_padding,
)
raise ValueError(f"{down_block_type} does not exist.")
class DownEncoderBlockCausal3D(nn.Module):
def __init__(self, in_channels: int, out_channels: int, num_layers: int = 1, resnet_eps: float = 1e-6,
resnet_groups: int = 32, add_downsample: bool = True, downsample_stride: int = 2, downsample_padding: int = 1, **kwargs):
super().__init__()
resnets = []
for i in range(num_layers):
in_channels = in_channels if i == 0 else out_channels
resnets.append(
ResnetBlockCausal3D(
in_channels=in_channels,
out_channels=out_channels,
eps=resnet_eps,
groups=resnet_groups,
)
)
self.resnets = nn.ModuleList(resnets)
if add_downsample:
self.downsamplers = nn.ModuleList([DownsampleCausal3D(
out_channels,
use_conv=True,
out_channels=out_channels,
padding=downsample_padding,
name="op",
stride=downsample_stride,
)])
else:
self.downsamplers = None
def forward(self, hidden_states: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
for resnet in self.resnets:
hidden_states = resnet(hidden_states, temb=None, scale=scale)
if self.downsamplers is not None:
for downsampler in self.downsamplers:
hidden_states = downsampler(hidden_states, scale)
return hidden_states
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "ComfyUI-WanVideoWrapper"
description = "ComfyUI wrapper nodes for WanVideo"
version = "1.3.8"
version = "1.4.5"
license = {file = "LICENSE"}
dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.33.0", "peft >= 0.17.0", "ftfy", "gguf >= 0.17.1", "pyloudnorm"]
+46 -2
View File
@@ -1,7 +1,35 @@
## Note: Due to the stupid amount of bots or people thinking this is some of video generation service, I've blocked new accounts from posting issues for now.
# ComfyUI wrapper nodes for [WanVideo](https://github.com/Wan-Video/Wan2.1) and related models.
## Memory use update (again)
I've made everythign less reliant on torch.compile for VRAM efficiency, so things should work better even without it. Also figured workaround for some issues when using compile that made first run use drastically more VRAM, issue I battled with myself a lot.
## Update notification that can affect memory use in old workflows
In a recent update I changed how unmerged LoRA weights are handled:
Previously mostly due to my laziness they were always loaded from RAM when used, this was of course inefficient and also made using torch.compile for LoRA applying difficult, thus forcing a graph break when using unmerged LoRAs.
Now the LoRA weights are assigned as buffers to the corresponding modules, so they are part of the blocks and obey the block swapping unifying the offloading and allowing LoRA weights to benefit from the prefetch feature for async offoading. Downside is that this means if you did not use block swap, you will see increased memory use as the LoRAs are part of the model and all on VRAM.
If you use block swap, the LoRAs are swapped along the rest of the block, but the block size is now larger, this means you may have to compensate with couple of more blocks swapped.
Example situation: you use 1GB LoRA unmerged and swap 20 blocks on 14B model, we can divide the LoRA size by block count, single block grows by 25MB, 20 blocks grow by 500MB, so your VRAM usage would be 500MB more than before, to compensate you swap 2 more blocks.
### Unrelated other VRAM issue with torch.compile
After any update that modifies the model code and when using torch.compile it's common to run into issues with VRAM, this can be caused by using older pytorch/triton version without latest compile fixes, and/or from old triton caches, mostly in Windows. This manifests in the issue that first run of new input size may have drastically increased memory use, which can clear from simply running it again, and once cached, not manifest again. Again I've only seen this happen in Windows.
To clear your Triton cache you can delete the contents of following (default) folders:
`C:\Users\<username>\.triton`
`C:\Users\<username>\AppData\Local\Temp\torchinductor_<username>`
## Note: Due to the stupid amount of bots or people thinking this is some of video generation service, I've blocked new accounts from posting issues for now.
# WORK IN PROGRESS (perpetually)
# Why should I use custom nodes when WanVideo works natively?
@@ -76,6 +104,22 @@ WanAnimate: https://github.com/Wan-Video/Wan2.2/tree/main/wan/modules/animate
Lynx: https://github.com/bytedance/lynx
MoCha: https://github.com/Orange-3DV-Team/MoCha
UniLumos: https://github.com/alibaba-damo-academy/Lumos-Custom
Bindweave: https://github.com/bytedance/BindWeave
Training free techniques:
TimeToMove: https://github.com/time-to-move/TTM
SteadyDancer: https://github.com/MCG-NJU/SteadyDancer
One-to-all-Animation: https://github.com/ssj9596/One-to-All-Animation
SCAIL: https://github.com/zai-org/SCAIL
Not exactly Wan model, but close enough to work with the code base:
+106
View File
@@ -0,0 +1,106 @@
# Modify from https://github.com/liyunsheng13/dcd/blob/main/models/imagenet/mobilenetv2_dcd.py
import torch
import torch.nn as nn
import torch.nn.functional as F
class Hsigmoid(nn.Module):
def __init__(self, inplace=True):
super(Hsigmoid, self).__init__()
self.inplace = inplace
def forward(self, x):
return F.relu6(x + 3., inplace=self.inplace) / 3.
class DYModule(nn.Module):
def __init__(self, inp, oup, fc_squeeze=8):
super(DYModule, self).__init__()
self.conv = nn.Conv2d(inp, oup, 1, 1, 0, bias=False)
if inp < oup:
self.mul = 4
reduction = 8
self.avg_pool = nn.AdaptiveAvgPool2d(2)
else:
self.mul = 1
reduction = 2
self.avg_pool = nn.AdaptiveAvgPool2d(1)
self.dim = min((inp * self.mul) // reduction, oup // reduction)
while self.dim ** 2 > inp * self.mul * 2:
reduction *= 2
self.dim = min((inp * self.mul) // reduction, oup // reduction)
if self.dim < 4:
self.dim = 4
squeeze = max(inp * self.mul, self.dim ** 2) // fc_squeeze
if squeeze < 4:
squeeze = 4
self.conv_q = nn.Conv2d(inp, self.dim, 1, 1, 0, bias=False)
self.fc = nn.Sequential(
nn.Linear(inp * self.mul, squeeze, bias=False),
SEModule_small(squeeze),
)
self.fc_phi = nn.Linear(squeeze, self.dim ** 2, bias=False)
self.fc_scale = nn.Linear(squeeze, oup, bias=False)
self.hs = Hsigmoid()
self.conv_p = nn.Conv2d(self.dim, oup, 1, 1, 0, bias=False)
# self.bn1 = nn.BatchNorm2d(self.dim)
self.bn1 = nn.GroupNorm(num_groups=4, num_channels=self.dim)
# self.bn2 = nn.BatchNorm1d(self.dim)
self.bn2 = nn.GroupNorm(num_groups=4, num_channels=self.dim)
def forward(self, x):
x_type = x.dtype
r = self.conv(x.to(self.conv.weight.dtype)).to(x_type)
b, c, h, w = x.size()
y = self.avg_pool(x).view(b, c * self.mul)
y = self.fc(y)
dy_phi = self.fc_phi(y).view(b, self.dim, self.dim)
dy_scale = self.hs(self.fc_scale(y)).view(b, -1, 1, 1)
r = dy_scale.expand_as(r) * r
x = self.conv_q(x.to(self.conv_q.weight.dtype)).to(self.bn1.weight.dtype)
x = self.bn1(x)
x = x.view(b, -1, h * w)
x = x + self.bn2(torch.matmul(dy_phi, x.to(dy_phi.dtype)).to(self.bn2.weight.dtype))
x = x.view(b, -1, h, w)
x = self.conv_p(x.to(self.conv_p.weight.dtype)).to(x_type)
return x + r
class SEModule_small(nn.Module):
def __init__(self, channel):
super(SEModule_small, self).__init__()
self.fc = nn.Sequential(
nn.Linear(channel, channel, bias=False),
Hsigmoid()
)
def forward(self, x):
y = self.fc(x)
return x * y
class SEModule(nn.Module):
def __init__(self, channel, reduction=4):
super(SEModule, self).__init__()
self.avg_pool = nn.AdaptiveAvgPool2d(1)
self.fc = nn.Sequential(
nn.Linear(channel, channel // reduction, bias=False),
nn.ReLU(inplace=True),
nn.Linear(channel // reduction, channel, bias=False),
Hsigmoid()
)
def forward(self, x):
b, c, _, _ = x.size()
y = self.avg_pool(x).view(b, c)
y = self.fc(y).view(b, c, 1, 1)
return x * y.expand_as(x)
+62
View File
@@ -0,0 +1,62 @@
import os
import torch
import numpy as np
from ..utils import log
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
import comfy.model_management as mm
from comfy.utils import load_torch_file, ProgressBar
import folder_paths
script_directory = os.path.dirname(os.path.abspath(__file__))
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoAddSteadyDancerEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"pose_latents_positive": ("LATENT",),
"pose_strength_spatial": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the pose embedding"}),
"pose_strength_temporal": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the pose embedding"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
},
"optional": {
"pose_latents_negative": ("LATENT",),
"clip_vision_embeds": ("WANVIDIMAGE_CLIPEMBEDS",),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, pose_latents_positive, pose_strength_spatial, pose_strength_temporal, start_percent=0.0, end_percent=1.0, pose_latents_negative=None, clip_vision_embeds=None):
sdancer_embeds = {
"cond_pos": pose_latents_positive["samples"][0],
"cond_neg": pose_latents_negative["samples"][0] if pose_latents_negative else None,
"pose_strength_spatial": pose_strength_spatial,
"pose_strength_temporal": pose_strength_temporal,
"start_percent": start_percent,
"end_percent": end_percent,
"clip_fea": clip_vision_embeds,
}
updated = dict(embeds)
updated["sdancer_embeds"] = sdancer_embeds
return (updated,)
NODE_CLASS_MAPPINGS = {
"WanVideoAddSteadyDancerEmbeds": WanVideoAddSteadyDancerEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddSteadyDancerEmbeds": "WanVideo Add SteadyDancer Embeds",
}
+133
View File
@@ -0,0 +1,133 @@
import torch
import torch.nn as nn
class FactorConv3d(nn.Module):
"""
(2+1)D decomposition of 3D convolution: 1xHxW spatial convolution → Swish → Tx1x1 temporal convolution
"""
def __init__(self,
in_channels: int,
out_channels: int,
kernel_size,
stride: int = 1,
dilation: int = 1):
super().__init__()
if isinstance(kernel_size, int):
k_t, k_h, k_w = kernel_size, kernel_size, kernel_size
else:
k_t, k_h, k_w = kernel_size
pad_t = (k_t - 1) * dilation // 2
pad_hw = (k_h - 1) * dilation // 2
self.spatial = nn.Conv3d(
in_channels, in_channels,
kernel_size=(1, k_h, k_w),
stride=(1, stride, stride),
padding=(0, pad_hw, pad_hw),
dilation=(1, dilation, dilation),
groups=in_channels,
bias=False
)
self.temporal = nn.Conv3d(
in_channels, out_channels,
kernel_size=(k_t, 1, 1),
stride=(stride, 1, 1),
padding=(pad_t, 0, 0),
dilation=(dilation, 1, 1),
bias=True
)
self.act = nn.SiLU()
def forward(self, x):
out_dtype = x.dtype
x = self.spatial(x.to(self.spatial.weight.dtype)).to(out_dtype)
x = self.act(x)
return self.temporal(x.to(self.temporal.weight.dtype)).to(out_dtype)
class LayerNorm2D(nn.Module):
"""
LayerNorm over C for a 4-D tensor (B, C, H, W)
"""
def __init__(self, num_channels, eps=1e-5, affine=True):
super().__init__()
self.num_channels = num_channels
self.eps = eps
self.affine = affine
if affine:
self.weight = nn.Parameter(torch.ones(1, num_channels, 1, 1))
self.bias = nn.Parameter(torch.zeros(1, num_channels, 1, 1))
def forward(self, x):
# x: (B, C, H, W)
mean = x.mean(dim=1, keepdim=True) # (B, 1, H, W)
var = x.var (dim=1, keepdim=True, unbiased=False)
x = (x - mean) / torch.sqrt(var + self.eps)
if self.affine:
x = x * self.weight + self.bias
return x
class PoseRefNetNoBNV3(nn.Module):
def __init__(self,
in_channels_c: int,
in_channels_x: int,
hidden_dim: int = 256,
num_heads: int = 8,
dropout: float = 0.1):
super().__init__()
self.d_model = hidden_dim
self.nhead = num_heads
self.proj_p = nn.Conv2d(in_channels_c, hidden_dim, kernel_size=1)
self.proj_r = nn.Conv2d(in_channels_x, hidden_dim, kernel_size=1)
self.proj_p_back = nn.Conv2d(hidden_dim, in_channels_c, kernel_size=1)
self.cross_attn = nn.MultiheadAttention(hidden_dim,
num_heads=num_heads,
dropout=dropout)
self.ffn_pose = nn.Sequential(
nn.Conv2d(hidden_dim, hidden_dim, kernel_size=1),
nn.SiLU(),
nn.Conv2d(hidden_dim, hidden_dim, kernel_size=1)
)
self.norm1 = LayerNorm2D(hidden_dim)
self.norm2 = LayerNorm2D(hidden_dim)
def forward(self, pose, ref, mask=None):
"""
pose : (B, C1, T, H, W)
ref : (B, C2, T, H, W)
mask : (B, T*H*W) optional key_padding_mask
return: (B, d_model, T, H, W)
"""
B, _, T, H, W = pose.shape
p_trans = pose.permute(0, 2, 1, 3, 4).contiguous().flatten(0, 1)
r_trans = ref.permute(0, 2, 1, 3, 4).contiguous().flatten(0, 1)
p_trans = self.proj_p(p_trans.to(self.proj_p.weight.dtype)).to(self.cross_attn.in_proj_weight.dtype).flatten(2).transpose(1, 2)
r_trans = self.proj_r(r_trans.to(self.proj_r.weight.dtype)).to(self.cross_attn.in_proj_weight.dtype).flatten(2).transpose(1, 2)
out = self.cross_attn(query=r_trans,
key=p_trans,
value=p_trans,
key_padding_mask=mask)[0]
out = self.norm1(out.transpose(1, 2).contiguous().view(B*T, -1, H, W))
out_type = out.dtype
out = out + self.ffn_pose(out.to(self.ffn_pose[0].weight.dtype)).to(out_type)
out = self.norm2(out)
out = self.proj_p_back(out.to(self.proj_p_back.weight.dtype)).to(out_type)
return out.view(B, T, -1, H, W).contiguous().transpose(1, 2)
+9 -2
View File
@@ -8,6 +8,7 @@ import torch.nn as nn
import torch.nn.functional as F
from tqdm.auto import tqdm
from collections import namedtuple
from ..wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38
DecoderResult = namedtuple("DecoderResult", ("frame", "memory"))
TWorkItem = namedtuple("TWorkItem", ("input_tensor", "block_index"))
@@ -146,7 +147,7 @@ def apply_model_with_memblocks(model, x, parallel, show_progress_bar):
return x
class TAEHV(nn.Module):
def __init__(self, state_dict, parallel=False, decoder_time_upscale=(True, True), decoder_space_upscale=(True, True, True), dtype=torch.float16):
def __init__(self, state_dict, parallel=False, decoder_time_upscale=(True, True), decoder_space_upscale=(True, True, True), dtype=torch.float16, model_name="taehv"):
"""Initialize pretrained TAEHV from the given checkpoint.
Arg:
@@ -161,6 +162,7 @@ class TAEHV(nn.Module):
if self.latent_channels == 48:
self.patch_size = 2
self.dtype = dtype
self.model_name = model_name
self.encoder = nn.Sequential(
conv(self.image_channels*self.patch_size**2, 64), nn.ReLU(inplace=True),
@@ -180,8 +182,11 @@ class TAEHV(nn.Module):
)
if state_dict is not None:
self.load_state_dict(self.patch_tgrow_layers(state_dict))
self.parallel = parallel
orig_vae = WanVideoVAE38() if self.latent_channels == 48 else WanVideoVAE()
self.mean = orig_vae.mean.to(dtype).movedim(1, 2)
self.inv_std = orig_vae.inv_std.to(dtype).movedim(1, 2)
def patch_tgrow_layers(self, sd):
"""Patch TGrow layers to use a smaller kernel if needed.
@@ -221,6 +226,8 @@ class TAEHV(nn.Module):
if False, frames will be processed sequentially.
Returns NTCHW RGB tensor with ~[0, 1] values.
"""
if "light" in self.model_name.lower():
x = x / self.inv_std.to(x) + self.mean.to(x)
x = apply_model_with_memblocks(self.decoder, x, self.parallel, show_progress_bar)
if self.patch_size > 1: x = F.pixel_shuffle(x, self.patch_size)
return x[:, self.frames_to_trim:]
@@ -0,0 +1,230 @@
# https://github.com/thu-ml/DiT-Extrapolation/blob/ultra-wan/sageattn/attn_qk_int8_per_block.py
import torch
import triton
import triton.language as tl
@triton.jit
def _attn_fwd_inner(acc, l_i, m_i, q, q_scale, kv_len, current_flag,
K_ptrs, K_scale_ptr, V_ptrs, stride_kn, stride_vn,
Block_bias_ptrs, stride_bbz, stride_bbh, stride_bm, stride_bn,
Decay_mask_ptrs, stride_dmz, stride_dmh, stride_dm, stride_dn,
start_m,
BLOCK_M: tl.constexpr, HEAD_DIM: tl.constexpr, BLOCK_N: tl.constexpr,
STAGE: tl.constexpr, offs_m: tl.constexpr, offs_n: tl.constexpr,
xpos_xi: tl.constexpr = 0.9999934149894527,
frame_tokens: tl.constexpr = 1560,
sigmoid_a: tl.constexpr = 1.0,
alpha_xpos_xi: tl.constexpr = 0.9999967941742395,
beta_xpos_xi: tl.constexpr = 0.9999860536252945,
sink_width: tl.constexpr = 4,
window_width: tl.constexpr = 16,
multi_factor: tl.constexpr = None,
entropy_factor: tl.constexpr = None,
):
lo, hi = 0, kv_len
for start_n in range(lo, hi, BLOCK_N):
start_n = tl.multiple_of(start_n, BLOCK_N)
k_mask = offs_n[None, :] < (kv_len - start_n)
k = tl.load(K_ptrs, mask = k_mask)
k_scale = tl.load(K_scale_ptr)
m = offs_m[:, None]
n = start_n + offs_n
qk = tl.dot(q, k).to(tl.float32) * q_scale * k_scale
window_th = frame_tokens * window_width / 2
dist2 = tl.abs(m - n).to(tl.int32)
dist_mask = dist2 <= window_th
negative_mask = (qk<0)
qk = tl.where(dist_mask | negative_mask, qk, qk*multi_factor)
window3 = (m <= frame_tokens) & (n > window_width*frame_tokens)
qk = tl.where(window3, -1e4, qk)
m_ij = tl.maximum(m_i, tl.max(qk, 1))
qk = qk - m_ij[:, None]
p = tl.math.exp2(qk)
l_ij = tl.sum(p, 1)
alpha = tl.math.exp2(m_i - m_ij)
l_i = l_i * alpha + l_ij
acc = acc * alpha[:, None]
v = tl.load(V_ptrs, mask = offs_n[:, None] < (kv_len - start_n))
p = p.to(tl.float16)
acc += tl.dot(p, v, out_dtype=tl.float16)
m_i = m_ij
K_ptrs += BLOCK_N * stride_kn
K_scale_ptr += 1
V_ptrs += BLOCK_N * stride_vn
return acc, l_i
@triton.jit
def _attn_fwd(Q, K, V, Q_scale, K_scale, Out,
Block_bias, Decay_mask,
flags, stride_f_b, stride_f_h,
stride_qz, stride_qh, stride_qn,
stride_kz, stride_kh, stride_kn,
stride_vz, stride_vh, stride_vn,
stride_oz, stride_oh, stride_on,
stride_bbz, stride_bbh, stride_bm, stride_bn,
stride_dmz, stride_dmh, stride_dm, stride_dn,
qo_len, kv_len, H: tl.constexpr, num_kv_groups: tl.constexpr,
HEAD_DIM: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
STAGE: tl.constexpr,
xpos_xi: tl.constexpr = 0.9999934149894527,
frame_tokens: tl.constexpr = 1560,
sigmoid_a: tl.constexpr = 1.0,
alpha_xpos_xi: tl.constexpr = 0.9999967941742395,
beta_xpos_xi: tl.constexpr = 0.9999860536252945,
sink_width: tl.constexpr = 4,
window_width: tl.constexpr = 16,
multi_factor: tl.constexpr = None,
entropy_factor: tl.constexpr = None,
):
start_m = tl.program_id(0)
off_z = tl.program_id(2).to(tl.int64)
off_h = tl.program_id(1).to(tl.int64)
q_scale_offset = (off_z * H + off_h) * tl.cdiv(qo_len, BLOCK_M)
k_scale_offset = (off_z * (H // num_kv_groups) + off_h // num_kv_groups) * tl.cdiv(kv_len, BLOCK_N)
flag_ptr = flags + off_z * stride_f_b + off_h * stride_f_h
current_flag = tl.load(flag_ptr)
offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = tl.arange(0, BLOCK_N)
offs_k = tl.arange(0, HEAD_DIM)
Q_ptrs = Q + (off_z * stride_qz + off_h * stride_qh) + offs_m[:, None] * stride_qn + offs_k[None, :]
Q_scale_ptr = Q_scale + q_scale_offset + start_m
K_ptrs = K + (off_z * stride_kz + (off_h // num_kv_groups) * stride_kh) + offs_n[None, :] * stride_kn + offs_k[:, None]
K_scale_ptr = K_scale + k_scale_offset
V_ptrs = V + (off_z * stride_vz + (off_h // num_kv_groups) * stride_vh) + offs_n[:, None] * stride_vn + offs_k[None, :]
O_block_ptr = Out + (off_z * stride_oz + off_h * stride_oh) + offs_m[:, None] * stride_on + offs_k[None, :]
# # 计算block_bias指针
Block_bias_ptrs = Block_bias + off_z * stride_bbz + off_h * stride_bbh
# 计算decay_mask指针
Decay_mask_ptrs = Decay_mask + off_z * stride_dmz + off_h * stride_dmh
m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float("inf")
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
q = tl.load(Q_ptrs, mask = offs_m[:, None] < qo_len)
q_scale = tl.load(Q_scale_ptr)
acc, l_i = _attn_fwd_inner(acc, l_i, m_i, q, q_scale, kv_len, current_flag, K_ptrs, K_scale_ptr, V_ptrs,
stride_kn, stride_vn,
Block_bias_ptrs, stride_bbz, stride_bbh, stride_bm, stride_bn,
Decay_mask_ptrs, stride_dmz, stride_dmh, stride_dm, stride_dn,
start_m,
BLOCK_M, HEAD_DIM, BLOCK_N,
4 - STAGE, offs_m, offs_n,
xpos_xi=xpos_xi,
frame_tokens=frame_tokens,
sigmoid_a=sigmoid_a,
alpha_xpos_xi=alpha_xpos_xi,
beta_xpos_xi=beta_xpos_xi,
sink_width=sink_width,
window_width=window_width,
multi_factor=multi_factor,
entropy_factor=entropy_factor,
)
acc = acc / l_i[:, None]
tl.store(O_block_ptr, acc.to(Out.type.element_ty), mask = (offs_m[:, None] < qo_len))
def forward(q, k, v, flags, block_bias, decay_mask, q_scale, k_scale, tensor_layout="HND", output_dtype=torch.float16,
xpos_xi: tl.constexpr = 0.9999934149894527,
frame_tokens: tl.constexpr = 1560,
sigmoid_a: tl.constexpr = 1.0,
alpha_xpos_xi: tl.constexpr = 0.9999967941742395,
beta_xpos_xi: tl.constexpr = 0.9999860536252945,
BLOCK_M: tl.constexpr = 128,
BLOCK_N: tl.constexpr = 128,
sink_width: tl.constexpr = 4,
window_width: tl.constexpr = 16,
multi_factor: tl.constexpr = None,
entropy_factor: tl.constexpr = None,
):
stage = 1
o = torch.empty(q.shape, dtype=output_dtype, device=q.device)
b, h_qo, qo_len, head_dim = q.shape
if block_bias is None:
block_bias = torch.zeros((b, h_qo, (qo_len + BLOCK_M - 1) // BLOCK_M, (qo_len + BLOCK_N - 1) // BLOCK_N), dtype=torch.float16, device=q.device)
if decay_mask is None:
decay_mask = torch.zeros((b, h_qo, (qo_len + BLOCK_M - 1) // BLOCK_M, (qo_len + BLOCK_N - 1) // BLOCK_N), dtype=torch.bool, device=q.device)
if tensor_layout == "HND":
b, h_qo, qo_len, head_dim = q.shape
_, h_kv, kv_len, _ = k.shape
stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(1), q.stride(2)
stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(1), k.stride(2)
stride_bz_v, stride_h_v, stride_seq_v = v.stride(0), v.stride(1), v.stride(2)
stride_bz_o, stride_h_o, stride_seq_o = o.stride(0), o.stride(1), o.stride(2)
stride_bbz, stride_bbh, stride_bm, stride_bn = block_bias.stride()
stride_dmz, stride_dmh, stride_dm, stride_dn = decay_mask.stride()
# elif tensor_layout == "NHD":
# b, qo_len, h_qo, head_dim = q.shape
# _, kv_len, h_kv, _ = k.shape
# stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(2), q.stride(1)
# stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(2), k.stride(1)
# stride_bz_v, stride_h_v, stride_seq_v = v.stride(0), v.stride(2), v.stride(1)
# stride_bz_o, stride_h_o, stride_seq_o = o.stride(0), o.stride(2), o.stride(1)
# stride_bbz, stride_bbh, stride_bm, stride_bn = block_bias.stride(0), block_bias.stride(2), block_bias.stride(1), block_bias.stride(3)
else:
raise ValueError(f"tensor_layout {tensor_layout} not supported")
stride_f_b, stride_f_h = flags.stride()
HEAD_DIM_K = head_dim
num_kv_groups = h_qo // h_kv
grid = (triton.cdiv(qo_len, BLOCK_M), h_qo, b)
_attn_fwd[grid](
q, k, v, q_scale, k_scale, o,
block_bias, decay_mask,
flags,
stride_f_b, stride_f_h,
stride_bz_q, stride_h_q, stride_seq_q,
stride_bz_k, stride_h_k, stride_seq_k,
stride_bz_v, stride_h_v, stride_seq_v,
stride_bz_o, stride_h_o, stride_seq_o,
stride_bbz, stride_bbh, stride_bm, stride_bn,
stride_dmz, stride_dmh, stride_dm, stride_dn,
qo_len, kv_len,
h_qo, num_kv_groups,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, HEAD_DIM=HEAD_DIM_K,
STAGE=stage,
num_warps=4 if head_dim == 64 else 8,
num_stages=3 if head_dim == 64 else 4,
xpos_xi=xpos_xi,
frame_tokens=frame_tokens,
sigmoid_a=sigmoid_a,
alpha_xpos_xi=alpha_xpos_xi,
beta_xpos_xi=beta_xpos_xi,
sink_width=sink_width,
window_width=window_width,
multi_factor=multi_factor,
entropy_factor=entropy_factor,
)
return o
+64
View File
@@ -0,0 +1,64 @@
# source https://github.com/thu-ml/DiT-Extrapolation/blob/ultra-wan/sageattn/core.py
import torch
import triton.language as tl
from .quant_per_block import per_block_int8
from .attn_qk_int8_per_block import forward as attn_false
from typing import Optional
def sage_attention(
qkv: list[torch.Tensor],
tensor_layout: str ="HND",
is_causal=False,
sm_scale: Optional[float] = None,
smooth_k: bool =True,
xpos_xi: tl.constexpr = 0.9999934149894527,
flags = None,
block_bias = None,
sigmoid_a: float = 1.0,
alpha_xpos_xi: float = 0.97,
beta_xpos_xi: float = 0.8,
decay_mask = None,
sink_width: int = 4,
window_width: int = 21,
multi_factor: Optional[float] = None,
entropy_factor: Optional[float] = None,
block_size : int = 64,
**kwargs
) -> torch.Tensor:
dtype = qkv[0].dtype
q, k, v = qkv[0].transpose(1, 2), qkv[1].transpose(1, 2), qkv[2].transpose(1, 2) # to HND
if flags == None:
flags = torch.zeros([q.shape[0],q.shape[1]], dtype=torch.int32, device=q.device)
seq_dim = 2
if smooth_k:
km = k.mean(dim=seq_dim, keepdim=True)
k -= km
else:
km = None
if dtype == torch.bfloat16 or dtype == torch.float32:
v = v.to(torch.float16)
if q.dtype != k.dtype or q.dtype != v.dtype:
k, v = k.to(q.dtype), v.to(q.dtype)
q_int8, q_scale, k_int8, k_scale = per_block_int8(q, k, sm_scale=sm_scale, tensor_layout=tensor_layout, BLKQ=block_size, BLKK=block_size)
del q, k
o = attn_false(q_int8, k_int8, v, flags, block_bias, decay_mask, q_scale, k_scale,
tensor_layout=tensor_layout, output_dtype=dtype, xpos_xi=xpos_xi, sigmoid_a=sigmoid_a,
alpha_xpos_xi=alpha_xpos_xi, beta_xpos_xi=beta_xpos_xi,
BLOCK_M=block_size, BLOCK_N=block_size,
sink_width=sink_width,
window_width=window_width,
multi_factor=multi_factor,
entropy_factor=entropy_factor,
)
return o.transpose(1, 2).contiguous()
+84
View File
@@ -0,0 +1,84 @@
# https://github.com/thu-ml/DiT-Extrapolation/blob/ultra-wan/sageattn/quant_per_block.py
import torch
import triton
import triton.language as tl
@triton.jit
def quant_per_block_int8_kernel(Input, Output, Scale, L,
stride_iz, stride_ih, stride_in,
stride_oz, stride_oh, stride_on,
stride_sz, stride_sh,
sm_scale,
C: tl.constexpr, BLK: tl.constexpr):
off_blk = tl.program_id(0)
off_h = tl.program_id(1)
off_b = tl.program_id(2)
offs_n = off_blk * BLK + tl.arange(0, BLK)
offs_k = tl.arange(0, C)
input_ptrs = Input + off_b * stride_iz + off_h * stride_ih + offs_n[:, None] * stride_in + offs_k[None, :]
output_ptrs = Output + off_b * stride_oz + off_h * stride_oh + offs_n[:, None] * stride_on + offs_k[None, :]
scale_ptrs = Scale + off_b * stride_sz + off_h * stride_sh + off_blk
x = tl.load(input_ptrs, mask=offs_n[:, None] < L)
x = x.to(tl.float32)
x *= sm_scale
scale = tl.max(tl.abs(x)) / 127.
x_int8 = x / scale
x_int8 += 0.5 * tl.where(x_int8 >= 0, 1, -1)
x_int8 = x_int8.to(tl.int8)
tl.store(output_ptrs, x_int8, mask=offs_n[:, None] < L)
tl.store(scale_ptrs, scale)
def per_block_int8(q, k, BLKQ=128, BLKK=64, sm_scale=None, tensor_layout="HND"):
q_int8 = torch.empty(q.shape, dtype=torch.int8, device=q.device)
k_int8 = torch.empty(k.shape, dtype=torch.int8, device=k.device)
if tensor_layout == "HND":
b, h_qo, qo_len, head_dim = q.shape
_, h_kv, kv_len, _ = k.shape
stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(1), q.stride(2)
stride_bz_qo, stride_h_qo, stride_seq_qo = q_int8.stride(0), q_int8.stride(1), q_int8.stride(2)
stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(1), k.stride(2)
stride_bz_ko, stride_h_ko, stride_seq_ko = k_int8.stride(0), k_int8.stride(1), k_int8.stride(2)
# elif tensor_layout == "NHD":
# b, qo_len, h_qo, head_dim = q.shape
# _, kv_len, h_kv, _ = k.shape
# stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(2), q.stride(1)
# stride_bz_qo, stride_h_qo, stride_seq_qo = q_int8.stride(0), q_int8.stride(2), q_int8.stride(1)
# stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(2), k.stride(1)
# stride_bz_ko, stride_h_ko, stride_seq_ko = k_int8.stride(0), k_int8.stride(2), k_int8.stride(1)
else:
raise ValueError(f"Unknown tensor layout: {tensor_layout}")
q_scale = torch.empty((b, h_qo, (qo_len + BLKQ - 1) // BLKQ, 1), device=q.device, dtype=torch.float32)
k_scale = torch.empty((b, h_kv, (kv_len + BLKK - 1) // BLKK, 1), device=q.device, dtype=torch.float32)
if sm_scale is None:
sm_scale = head_dim**-0.5
grid = ((qo_len + BLKQ - 1) // BLKQ, h_qo, b)
quant_per_block_int8_kernel[grid](
q, q_int8, q_scale, qo_len,
stride_bz_q, stride_h_q, stride_seq_q,
stride_bz_qo, stride_h_qo, stride_seq_qo,
q_scale.stride(0), q_scale.stride(1),
sm_scale=(sm_scale * 1.44269504),
C=head_dim, BLK=BLKQ
)
grid = ((kv_len + BLKK - 1) // BLKK, h_kv, b)
quant_per_block_int8_kernel[grid](
k, k_int8, k_scale, kv_len,
stride_bz_k, stride_h_k, stride_seq_k,
stride_bz_ko, stride_h_ko, stride_seq_ko,
k_scale.stride(0), k_scale.stride(1),
sm_scale=1.0,
C=head_dim, BLK=BLKK
)
return q_int8, q_scale, k_int8, k_scale
+6 -13
View File
@@ -115,16 +115,9 @@ class WanRotaryPosEmbed(nn.Module):
freqs_w = freqs[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1)
freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1).reshape(1, 1, ppf * pph * ppw, -1)
return freqs
from ..wanvideo.modules.attention import sageattn_func
def zero_module(module):
# Zero out the parameters of a module and return it.
for p in module.parameters():
p.detach().zero_()
return module
class SimpleAttnProcessor2_0:
def __init__(self, attention_mode):
self.attention_mode = attention_mode
@@ -278,7 +271,7 @@ class MaskCamEmbed(nn.Module):
mid_channels = controlnet_cfg.get("mid_channels", 64)
self.mask_proj = nn.Sequential(nn.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8)),
nn.GroupNorm(mid_channels // 8, mid_channels), nn.SiLU())
self.mask_zero_proj = zero_module(nn.Conv3d(mid_channels, controlnet_cfg["conv_out_dim"], kernel_size=(1, 2, 2), stride=(1, 2, 2)))
self.mask_zero_proj = nn.Conv3d(mid_channels, controlnet_cfg["conv_out_dim"], kernel_size=(1, 2, 2), stride=(1, 2, 2))
def forward(self, add_inputs: torch.Tensor):
# render_mask.shape [b,c,f,h,w]
@@ -321,7 +314,7 @@ class WanControlNet(ModelMixin):
)
self.proj_out = nn.ModuleList(
[
zero_module(nn.Linear(self.dim, 5120))
nn.Linear(self.dim, 5120)
for _ in range(controlnet_cfg["num_layers"])
]
)
@@ -341,7 +334,7 @@ class WanControlNet(ModelMixin):
self.controlnet_mask_embedding = MaskCamEmbed(controlnet_cfg)
def forward(self, render_latent, render_mask, camera_embedding, temb, device):
def forward(self, render_latent, render_mask, camera_embedding, temb, out_device):
controlnet_rotary_emb = self.controlnet_rope(render_latent)
controlnet_inputs = self.controlnet_patch_embedding(render_latent.to(torch.float32))
if not self.quantized:
@@ -361,7 +354,7 @@ class WanControlNet(ModelMixin):
if add_inputs is not None:
add_inputs = self.controlnet_mask_embedding(add_inputs)
controlnet_inputs = controlnet_inputs + add_inputs
hidden_states = self.proj_in(controlnet_inputs)
controlnet_states = []
@@ -371,6 +364,6 @@ class WanControlNet(ModelMixin):
temb=temb,
rotary_emb=controlnet_rotary_emb
)
controlnet_states.append(self.proj_out[i](hidden_states).to(device))
controlnet_states.append(self.proj_out[i](hidden_states).to(out_device))
return controlnet_states
+29 -33
View File
@@ -10,8 +10,6 @@ from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
import folder_paths
import json
import numpy as np
class WanVideoUni3C_ControlnetLoader:
@classmethod
@@ -22,7 +20,7 @@ class WanVideoUni3C_ControlnetLoader:
"base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}),
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e5m2'], {"default": 'disabled', "tooltip": "optional quantization method"}),
"load_device": (["main_device", "offload_device"], {"default": "main_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
"load_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
"attention_mode": ([
"sdpa",
"sageattn",
@@ -45,17 +43,17 @@ class WanVideoUni3C_ControlnetLoader:
offload_device = mm.unet_offload_device()
transformer_load_device = device if load_device == "main_device" else offload_device
base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[base_precision]
model_path = folder_paths.get_full_path_or_raise("controlnet", model)
sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True)
if not "controlnet_patch_embedding.weight" in sd:
raise ValueError("Invalid ControlNet model")
in_channels = sd["controlnet_patch_embedding.weight"].shape[1]
ffn_dim = sd["controlnet_blocks.0.ffn.0.bias"].shape[0]
@@ -79,7 +77,7 @@ class WanVideoUni3C_ControlnetLoader:
with init_empty_weights():
controlnet = WanControlNet(controlnet_cfg)
controlnet.eval()
if quantization == "disabled":
for k, v in sd.items():
if isinstance(v, torch.Tensor):
@@ -97,18 +95,18 @@ class WanVideoUni3C_ControlnetLoader:
else:
dtype = base_dtype
params_to_keep = {"norm", "head", "time_in", "vector_in", "controlnet_patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "proj_in"}
log.info("Using accelerate to load and assign controlnet model weights to device...")
param_count = sum(1 for _ in controlnet.named_parameters())
for name, param in tqdm(controlnet.named_parameters(),
desc=f"Loading transformer parameters to {transformer_load_device}",
for name, param in tqdm(controlnet.named_parameters(),
desc=f"Loading transformer parameters to {transformer_load_device}",
total=param_count,
leave=True):
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
if "controlnet_patch_embedding" in name:
dtype_to_use = torch.float32
set_module_tensor_to_device(controlnet, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
del sd
if compile_args is not None:
@@ -123,8 +121,8 @@ class WanVideoUni3C_ControlnetLoader:
for i, block in enumerate(controlnet.controlnet_blocks):
controlnet.controlnet_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
else:
controlnet = torch.compile(controlnet, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
controlnet = torch.compile(controlnet, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
if load_device == "offload_device" and controlnet.device != offload_device:
log.info(f"Moving controlnet model from {controlnet.device} to {offload_device}")
@@ -146,6 +144,7 @@ class WanVideoUni3C_embeds:
"optional": {
"render_latent": ("LATENT",),
"render_mask": ("MASK", {"tooltip": "NOT IMPLEMENTED!"}),
"offload": ("BOOLEAN", {"default": True, "tooltip": "If enabled, the controlnet model will be offloaded before main model block processing to save VRAM."}),
},
}
@@ -154,9 +153,7 @@ class WanVideoUni3C_embeds:
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, controlnet, strength, start_percent, end_percent, render_latent=None, render_mask=None):
device = mm.get_torch_device()
def process(self, controlnet, strength, start_percent, end_percent, render_latent=None, render_mask=None, offload=True):
latent_mask = latents = None
if render_latent is not None:
@@ -164,17 +161,17 @@ class WanVideoUni3C_embeds:
# nframe = latents.shape[2] * 4
# height = latents.shape[3] * 8
# width = latents.shape[4] * 8
if render_mask is not None:
raise NotImplementedError("render_mask is not implemented at this time")
mask = torch.nn.functional.interpolate(
render_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
size=(nframe, height, width),
mode='trilinear',
align_corners=False
).squeeze(0)
latent_mask = mask.unsqueeze(0).to(device)
log.info(f"latent mask shape {latent_mask.shape}")
# mask = torch.nn.functional.interpolate(
# render_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
# size=(nframe, height, width),
# mode='trilinear',
# align_corners=False
# ).squeeze(0)
# latent_mask = mask.unsqueeze(0).to(device)
# log.info(f"latent mask shape {latent_mask.shape}")
# # load camera
# cam_info = json.load(open(f"{render_path}/cam_info.json"))
@@ -199,7 +196,7 @@ class WanVideoUni3C_embeds:
# K_inv = K.inverse()
# intrinsic = K[None].repeat(nframe, 1, 1)
# w2c_0, c2w_0 = set_initial_camera(start_elevation, depth_avg)
# w2cs, c2ws, intrinsic = build_cameras(cam_traj=cam_traj,
# w2c_0=w2c_0,
@@ -215,7 +212,7 @@ class WanVideoUni3C_embeds:
# y_offset=y_offset,
# z_offset=z_offset)
# from .camera import get_camera_embedding
# camera_embedding = get_camera_embedding(intrinsic, w2cs, nframe, height, width, normalize=True)
#print("camera embedding shape", camera_embedding.shape)
@@ -227,11 +224,12 @@ class WanVideoUni3C_embeds:
"end": end_percent,
"render_latent": latents,
"render_mask": latent_mask,
"camera_embedding": None
"camera_embedding": None,
"offload": offload,
}
return (uni3c_embeds,)
NODE_CLASS_MAPPINGS = {
"WanVideoUni3C_ControlnetLoader": WanVideoUni3C_ControlnetLoader,
"WanVideoUni3C_embeds": WanVideoUni3C_embeds,
@@ -240,5 +238,3 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoUni3C_ControlnetLoader": "WanVideo Uni3C Controlnet Loader",
"WanVideoUni3C_embeds": "WanVideo Uni3C Embeds",
}
+42 -53
View File
@@ -9,37 +9,28 @@ from ..utils import log
import comfy.model_management as mm
from comfy.utils import ProgressBar
import comfy.ops
ops = comfy.ops.disable_weight_init
def update_transformer(transformer, state_dict):
concat_dim = 4
transformer.dwpose_embedding = nn.Sequential(
nn.Conv3d(3, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)),
nn.SiLU(),
nn.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)),
nn.SiLU(),
nn.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)),
nn.SiLU(),
nn.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,2,2), padding=(1,1,1)),
nn.SiLU(),
nn.Conv3d(concat_dim * 4, concat_dim * 4, 3, stride=(2,2,2), padding=1),
nn.SiLU(),
nn.Conv3d(concat_dim * 4, concat_dim * 4, 3, stride=(2,2,2), padding=1),
nn.SiLU(),
nn.Conv3d(concat_dim * 4, 5120, (1,2,2), stride=(1,2,2), padding=0))
ops.Conv3d(3, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)), nn.SiLU(),
ops.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)), nn.SiLU(),
ops.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)), nn.SiLU(),
ops.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,2,2), padding=(1,1,1)), nn.SiLU(),
ops.Conv3d(concat_dim * 4, concat_dim * 4, 3, stride=(2,2,2), padding=1), nn.SiLU(),
ops.Conv3d(concat_dim * 4, concat_dim * 4, 3, stride=(2,2,2), padding=1), nn.SiLU(),
ops.Conv3d(concat_dim * 4, 5120, (1,2,2), stride=(1,2,2), padding=0))
randomref_dim = 20
transformer.randomref_embedding_pose = nn.Sequential(
nn.Conv2d(3, concat_dim * 4, 3, stride=1, padding=1),
nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=1, padding=1),
nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=1, padding=1),
nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1),
nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1),
nn.SiLU(),
nn.Conv2d(3, concat_dim * 4, 3, stride=1, padding=1), nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=1, padding=1), nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=1, padding=1), nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1), nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1), nn.SiLU(),
nn.Conv2d(concat_dim * 4, randomref_dim, 3, stride=2, padding=1),
)
unianimate_sd = {}
@@ -123,7 +114,7 @@ class DWposeDetector:
body = candidate[:,:18].copy()
body = body.reshape(nums*18, locs)
score = subset[:,:18].copy()
for i in range(len(score)):
for j in range(len(score[i])):
if score[i][j] > score_threshold:
@@ -142,17 +133,17 @@ class DWposeDetector:
else:
bodyfoot_score[i][j] = -1
if -1 not in bodyfoot_score[:,18] and -1 not in bodyfoot_score[:,19]:
bodyfoot_score[:,18] = np.array([18.])
bodyfoot_score[:,18] = np.array([18.])
else:
bodyfoot_score[:,18] = np.array([-1.])
if -1 not in bodyfoot_score[:,21] and -1 not in bodyfoot_score[:,22]:
bodyfoot_score[:,19] = np.array([19.])
bodyfoot_score[:,19] = np.array([19.])
else:
bodyfoot_score[:,19] = np.array([-1.])
bodyfoot_score = bodyfoot_score[:, :20]
bodyfoot = candidate[:,:24].copy()
for i in range(nums):
if -1 not in bodyfoot[i][18] and -1 not in bodyfoot[i][19]:
bodyfoot[i][18] = (bodyfoot[i][18]+bodyfoot[i][19])/2
@@ -162,7 +153,7 @@ class DWposeDetector:
bodyfoot[i][19] = (bodyfoot[i][21]+bodyfoot[i][22])/2
else:
bodyfoot[i][19] = np.array([-1., -1.])
bodyfoot = bodyfoot[:,:20,:]
bodyfoot = bodyfoot.reshape(nums*20, locs)
@@ -172,7 +163,7 @@ class DWposeDetector:
hands = candidate[:,92:113]
hands = np.vstack([hands, candidate[:,113:]])
# bodies = dict(candidate=body, subset=score)
bodies = dict(candidate=bodyfoot, subset=bodyfoot_score, score=bodyfoot_score)
pose = dict(bodies=bodies, hands=hands, faces=faces)
@@ -180,7 +171,7 @@ class DWposeDetector:
# return draw_pose(pose, H, W)
return pose
def draw_pose(pose, H, W, stick_width=4,draw_body=True, draw_hands=True, draw_feet=True,
def draw_pose(pose, H, W, stick_width=4,draw_body=True, draw_hands=True, draw_feet=True,
body_keypoint_size=4, hand_keypoint_size=4, draw_head=True):
from .dwpose.util import draw_body_and_foot, draw_handpose, draw_facepose
bodies = pose['bodies']
@@ -202,7 +193,7 @@ def draw_pose(pose, H, W, stick_width=4,draw_body=True, draw_hands=True, draw_fe
def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_threshold, stick_width,
draw_body=True, draw_hands=True, hand_keypoint_size=4, draw_feet=True,
body_keypoint_size=4, handle_not_detected="repeat", draw_head=True):
results_vis = []
comfy_pbar = ProgressBar(len(pose_images))
@@ -224,7 +215,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
pose = np.zeros_like(img)
results_vis.append(pose)
comfy_pbar.update(1)
bodies = results_vis[0]['bodies']
faces = results_vis[0]['faces']
hands = results_vis[0]['hands']
@@ -268,7 +259,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
results_vis[0]['faces'][:,:,1] *= y_ratio
results_vis[0]['hands'][:,:,0] *= x_ratio
results_vis[0]['hands'][:,:,1] *= y_ratio
########neck########
l_neck_ref = ((ref_candidate[0][0] - ref_candidate[1][0]) ** 2 + (ref_candidate[0][1] - ref_candidate[1][1]) ** 2) ** 0.5
l_neck_0 = ((candidate[0][0] - candidate[1][0]) ** 2 + (candidate[0][1] - candidate[1][1]) ** 2) ** 0.5
@@ -287,7 +278,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
results_vis[0]['bodies']['candidate'][16,1] += y_offset_neck
results_vis[0]['bodies']['candidate'][17,0] += x_offset_neck
results_vis[0]['bodies']['candidate'][17,1] += y_offset_neck
########shoulder2########
l_shoulder2_ref = ((ref_candidate[2][0] - ref_candidate[1][0]) ** 2 + (ref_candidate[2][1] - ref_candidate[1][1]) ** 2) ** 0.5
l_shoulder2_0 = ((candidate[2][0] - candidate[1][0]) ** 2 + (candidate[2][1] - candidate[1][1]) ** 2) ** 0.5
@@ -435,9 +426,9 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
results_vis[0]['bodies']['candidate'][17,0] += x_offset_head17
results_vis[0]['bodies']['candidate'][17,1] += y_offset_head17
########MovingAverage########
########left leg########
l_ll1_ref = ((ref_candidate[8][0] - ref_candidate[9][0]) ** 2 + (ref_candidate[8][1] - ref_candidate[9][1]) ** 2) ** 0.5
l_ll1_0 = ((candidate[8][0] - candidate[9][0]) ** 2 + (candidate[8][1] - candidate[9][1]) ** 2) ** 0.5
@@ -522,7 +513,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
results_vis[i]['bodies']['candidate'][17,1] += y_offset_neck
########shoulder2########
x_offset_shoulder2 = (results_vis[i]['bodies']['candidate'][1][0]-results_vis[i]['bodies']['candidate'][2][0])*(1.-shoulder2_ratio)
y_offset_shoulder2 = (results_vis[i]['bodies']['candidate'][1][1]-results_vis[i]['bodies']['candidate'][2][1])*(1.-shoulder2_ratio)
@@ -676,7 +667,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
results_vis[i]['bodies']['candidate'] += offset[np.newaxis, :]
results_vis[i]['faces'] += offset[np.newaxis, np.newaxis, :]
results_vis[i]['hands'] += offset[np.newaxis, np.newaxis, :]
dwpose_woface_list = []
for i in range(len(results_vis)):
#try:
@@ -724,11 +715,11 @@ class WanVideoUniAnimateDWPoseDetector:
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, pose_images, score_threshold, stick_width, reference_pose_image=None, draw_body=True, body_keypoint_size=4,
def process(self, pose_images, score_threshold, stick_width, reference_pose_image=None, draw_body=True, body_keypoint_size=4,
draw_feet=True, draw_hands=True, hand_keypoint_size=4, colorspace="RGB", handle_not_detected="empty", draw_head=True):
device = mm.get_torch_device()
#model loading
dw_pose_model = "dw-ll_ucoco_384_bs5.torchscript.pt"
yolo_model = "yolox_l.torchscript.pt"
@@ -742,27 +733,27 @@ class WanVideoUniAnimateDWPoseDetector:
if not os.path.exists(model_det):
log.info(f"Downloading yolo model to: {model_base_path}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id="hr16/yolox-onnx",
snapshot_download(repo_id="hr16/yolox-onnx",
allow_patterns=[f"*{yolo_model}*"],
local_dir=model_base_path,
local_dir=model_base_path,
local_dir_use_symlinks=False)
if not os.path.exists(model_pose):
log.info(f"Downloading dwpose model to: {model_base_path}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id="hr16/DWPose-TorchScript-BatchSize5",
snapshot_download(repo_id="hr16/DWPose-TorchScript-BatchSize5",
allow_patterns=[f"*{dw_pose_model}*"],
local_dir=model_base_path,
local_dir=model_base_path,
local_dir_use_symlinks=False)
if not hasattr(self, "det") or not hasattr(self, "pose"):
self.det = torch.jit.load(model_det, map_location=device)
self.pose = torch.jit.load(model_pose, map_location=device)
self.dwpose_detector = DWposeDetector(self.det, self.pose)
self.dwpose_detector = DWposeDetector(self.det, self.pose)
#model inference
height, width = pose_images.shape[1:3]
pose_np = pose_images.cpu().numpy() * 255
ref_np = None
if reference_pose_image is not None:
@@ -772,11 +763,11 @@ class WanVideoUniAnimateDWPoseDetector:
prev_fuser_state = torch._C._jit_texpr_fuser_enabled()
torch._C._jit_set_texpr_fuser_enabled(False) # removes warmup delay, may want to enable later
poses, reference_pose = pose_extract(pose_np, ref_np, self.dwpose_detector, height, width, score_threshold, stick_width=stick_width,
draw_body=draw_body, body_keypoint_size=body_keypoint_size, draw_feet=draw_feet,
draw_body=draw_body, body_keypoint_size=body_keypoint_size, draw_feet=draw_feet,
draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size, handle_not_detected=handle_not_detected, draw_head=draw_head)
poses = poses / 255.0
torch._C._jit_set_texpr_fuser_enabled(prev_fuser_state)
if reference_pose_image is not None:
reference_pose = reference_pose.unsqueeze(0) / 255.0
else:
@@ -828,11 +819,9 @@ class WanVideoUniAnimatePoseInput:
NODE_CLASS_MAPPINGS = {
"WanVideoUniAnimatePoseInput": WanVideoUniAnimatePoseInput,
"WanVideoUniAnimateDWPoseDetector": WanVideoUniAnimateDWPoseDetector,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoUniAnimatePoseInput": "WanVideo UniAnimate Pose Input",
"WanVideoUniAnimateDWPoseDetector": "WanVideo UniAnimate DWPose Detector",
}
+120 -11
View File
@@ -4,17 +4,116 @@ import logging
import math
from tqdm import tqdm
from pathlib import Path
import os
import gc
import types, collections
from comfy.utils import ProgressBar, copy_to_param, set_attr_param
from comfy.model_patcher import get_key_weight, string_to_seed
from comfy.lora import calculate_weight
from comfy.model_management import cast_to_device
from comfy.float import stochastic_rounding
from .custom_linear import remove_lora_from_module
import folder_paths
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
log = logging.getLogger(__name__)
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
try:
from .gguf.gguf import GGUFParameter
except:
pass
COLOR_CODES = {
"reset": "\033[0m",
"red": "\033[31m",
"green": "\033[32m",
"yellow": "\033[33m",
"blue": "\033[34m",
"magenta": "\033[35m",
"cyan": "\033[36m",
"white": "\033[37m",
}
def color_text(text, color):
try:
return f"{COLOR_CODES.get(color, COLOR_CODES['reset'])}{text}{COLOR_CODES['reset']}"
except Exception:
return text
class MetaParameter(torch.nn.Parameter):
def __new__(cls, dtype, quant_type=None):
data = torch.empty(0, dtype=dtype)
self = torch.nn.Parameter(data, requires_grad=False)
self.quant_type = quant_type
return self
def offload_transformer(transformer, remove_lora=True):
transformer.teacache_state.clear_all()
transformer.magcache_state.clear_all()
transformer.easycache_state.clear_all()
if transformer.patched_linear:
for name, param in transformer.named_parameters():
if "loras" in name or "controlnet" in name:
continue
module = transformer
subnames = name.split('.')
for subname in subnames[:-1]:
module = getattr(module, subname)
attr_name = subnames[-1]
if param.data.is_floating_point():
meta_param = torch.nn.Parameter(torch.empty_like(param.data, device='meta'), requires_grad=False)
setattr(module, attr_name, meta_param)
elif isinstance(param.data, GGUFParameter):
quant_type = getattr(param, 'quant_type', None)
setattr(module, attr_name, MetaParameter(param.data.dtype, quant_type))
else:
pass
if remove_lora:
remove_lora_from_module(transformer)
else:
transformer.to(offload_device)
for block in transformer.blocks:
block.kv_cache = None
if transformer.audio_model is not None and hasattr(block, 'audio_block'):
block.audio_block = None
mm.soft_empty_cache()
gc.collect()
def init_blockswap(transformer, block_swap_args, model):
if not transformer.patched_linear:
if block_swap_args is not None:
for name, param in transformer.named_parameters():
if "block" not in name or "control_adapter" in name or "face" in name:
param.data = param.data.to(device)
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
param.data = param.data.to(offload_device)
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
param.data = param.data.to(offload_device)
transformer.block_swap(
block_swap_args["blocks_to_swap"] - 1 ,
block_swap_args["offload_txt_emb"],
block_swap_args["offload_img_emb"],
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
)
elif model["auto_cpu_offload"]:
for module in transformer.modules():
if hasattr(module, "offload"):
module.offload()
if hasattr(module, "onload"):
module.onload()
for block in transformer.blocks:
block.modulation = torch.nn.Parameter(block.modulation.to(device))
transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device))
else:
transformer.to(device)
def check_device_same(first_device, second_device):
if first_device.type != second_device.type:
return False
@@ -108,13 +207,11 @@ def check_diffusers_version():
except importlib.metadata.PackageNotFoundError:
raise AssertionError("diffusers is not installed.")
def print_memory(device):
memory = torch.cuda.memory_allocated(device) / 1024**3
def print_memory(device, process="Sampling"):
max_memory = torch.cuda.max_memory_allocated(device) / 1024**3
max_reserved = torch.cuda.max_memory_reserved(device) / 1024**3
log.info(f"Allocated memory: {memory=:.3f} GB")
log.info(f"Max allocated memory: {max_memory=:.3f} GB")
log.info(f"Max reserved memory: {max_reserved=:.3f} GB")
log.info(f"[{process}] Max allocated memory: {max_memory=:.3f} GB")
log.info(f"[{process}] Max reserved memory: {max_reserved=:.3f} GB")
#memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False)
#log.info(f"Memory Summary:\n{memory_summary}")
@@ -125,6 +222,18 @@ def get_module_memory_mb(module):
memory += param.nelement() * param.element_size()
return memory / (1024 * 1024) # Convert to MB
def get_module_memory_mb_per_device(module):
memory_per_device = {}
memory = 0
for param in module.parameters():
if param.data is not None:
device = str(param.device)
memory += param.nelement() * param.element_size()
memory_per_device[device] = memory_per_device.get(device, 0) + memory
memory_per_device = {dev: mem / (1024 * 1024) for dev, mem in memory_per_device.items()}
return memory_per_device
def get_tensor_memory(tensor):
memory_bytes = tensor.element_size() * tensor.nelement()
return f"{memory_bytes / (1024 * 1024):.2f} MB"
@@ -140,7 +249,7 @@ def patch_weight_to_device(self, key, device_to=None, inplace_update=False, back
self.backup[key] = collections.namedtuple('Dimension', ['weight', 'inplace_update'])(weight.to(device=self.offload_device, copy=inplace_update), inplace_update)
if device_to is not None:
temp_weight = cast_to_device(weight, device_to, torch.float32, copy=True)
temp_weight = mm.cast_to_device(weight, device_to, torch.float32, copy=True)
else:
temp_weight = weight.to(torch.float32, copy=True)
if convert_func is not None:
@@ -584,9 +693,9 @@ def check_duplicate_nodes():
"""Check ComfyUI custom_nodes directory for duplicate installations"""
custom_nodes_dir = Path(folder_paths.folder_names_and_paths["custom_nodes"][0][0])
current_path = Path(__file__).parent
wanvideo_dirs = []
# Check all directories in custom_nodes
for path in custom_nodes_dir.iterdir():
if (path.is_dir() and
@@ -594,7 +703,7 @@ def check_duplicate_nodes():
'wanvideo' in path.name.lower() and
'wrapper' in path.name.lower()):
wanvideo_dirs.append(str(path))
return wanvideo_dirs
#https://github.com/temporalscorerescaling/TSR/
+71 -195
View File
@@ -1,25 +1,21 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import torch
from ...utils import log
# Flash Attention imports
try:
import flash_attn_interface
FLASH_ATTN_3_AVAILABLE = True
except Exception as e:
FLASH_ATTN_3_AVAILABLE = False
from comfy.ldm.modules.attention import optimized_attention
def attention_func_error(*args, **kwargs):
raise ImportError("Selected attention mode not available. Please ensure required packages are installed correctly.")
from .attention_flash import flash_attention
try:
import flash_attn
FLASH_ATTN_2_AVAILABLE = True
except Exception as e:
FLASH_ATTN_2_AVAILABLE = False
# Sage Attention imports
# using custom ops to avoid graph breaks with torch.compile
try:
from sageattention import sageattn
@torch.compiler.disable()
def sageattn_func(q, k, v, attn_mask=None, dropout_p=0, is_causal=False, tensor_layout="HND"):
@torch.library.custom_op("wanvideo::sageattn", mutates_args=())
def sageattn_func(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, attn_mask: torch.Tensor | None = None, dropout_p: float = 0.0, is_causal: bool = False, tensor_layout: str = "HND"
) -> torch.Tensor:
if not (q.dtype == k.dtype == v.dtype):
return sageattn(q, k.to(q.dtype), v.to(q.dtype), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout)
elif q.dtype == torch.float32:
@@ -27,6 +23,13 @@ try:
else:
return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout)
@sageattn_func.register_fake
def _(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=False, tensor_layout="HND"):
# Return tensor with same shape as q
return q.clone()
sageattn_func = torch.ops.wanvideo.sageattn
def sageattn_func_compiled(q, k, v, attn_mask=None, dropout_p=0, is_causal=False, tensor_layout="HND"):
if not (q.dtype == k.dtype == v.dtype):
return sageattn(q, k.to(q.dtype), v.to(q.dtype), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout)
@@ -40,20 +43,14 @@ except Exception as e:
log.warning("sageattention package is not installed, sageattention will not be available")
elif isinstance(e, ImportError) and "DLL" in str(e):
log.warning("sageattention DLL loading error, sageattention will not be available")
sageattn_func = None
sageattn_func = attention_func_error
try:
from sageattn3 import sageattn3_blackwell as sageattn_blackwell
except:
try:
from sageattn import sageattn_blackwell
except:
SAGE3_AVAILABLE = False
try:
from sageattention import sageattn_varlen
@torch.compiler.disable()
def sageattn_varlen_func(q, k, v, q_lens, k_lens, max_seqlen_q, max_seqlen_k, dropout_p=0, is_causal=False):
from typing import List
@torch.library.custom_op("wanvideo::sageattn_varlen", mutates_args=())
def sageattn_varlen_func(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, q_lens: List[int], k_lens: List[int], max_seqlen_q: int, max_seqlen_k: int, dropout_p: float = 0.0, is_causal: bool = False) -> torch.Tensor:
cu_seqlens_q = torch.tensor([0] + list(torch.cumsum(torch.tensor(q_lens), dim=0)), device=q.device, dtype=torch.int32)
cu_seqlens_k = torch.tensor([0] + list(torch.cumsum(torch.tensor(k_lens), dim=0)), device=q.device, dtype=torch.int32)
if not (q.dtype == k.dtype == v.dtype):
@@ -62,180 +59,59 @@ try:
return sageattn_varlen(q.to(torch.float16), k.to(torch.float16), v.to(torch.float16), cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, dropout_p=dropout_p, is_causal=is_causal).to(torch.float32)
else:
return sageattn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, dropout_p=dropout_p, is_causal=is_causal)
except:
sageattn_varlen_func = None
__all__ = [
'flash_attention',
'attention',
]
@sageattn_varlen_func.register_fake
def _(q, k, v, q_lens, k_lens, max_seqlen_q, max_seqlen_k, dropout_p=0.0, is_causal=False):
# Return tensor with same shape as q
return q.clone()
sageattn_varlen_func = torch.ops.wanvideo.sageattn_varlen
except:
sageattn_varlen_func = attention_func_error
# sage3
try:
from sageattn3 import sageattn3_blackwell as sageattn_blackwell
except:
try:
from sageattn import sageattn_blackwell
except:
sageattn_blackwell = attention_func_error
try:
from ...ultravico.sageattn.core import sage_attention as sageattn_ultravico
@torch.library.custom_op("wanvideo::sageattn_ultravico", mutates_args=())
def sageattn_func_ultravico(qkv: List[torch.Tensor], attn_mask: torch.Tensor | None = None, dropout_p: float = 0.0, is_causal: bool = False, multi_factor: float = 0.9, frame_tokens: int = 1536
) -> torch.Tensor:
return sageattn_ultravico(qkv, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, multi_factor=multi_factor, frame_tokens=frame_tokens)
@sageattn_func_ultravico.register_fake
def _(qkv, attn_mask=None, dropout_p=0.0, is_causal=False, multi_factor=0.9):
return torch.empty_like(qkv[0]).contiguous()
sageattn_func_ultravico = torch.ops.wanvideo.sageattn_ultravico
except:
sageattn_func_ultravico = attention_func_error
def flash_attention(
q,
k,
v,
q_lens=None,
k_lens=None,
dropout_p=0.,
softmax_scale=None,
q_scale=None,
causal=False,
window_size=(-1, -1),
deterministic=False,
dtype=torch.bfloat16,
version=None,
):
"""
q: [B, Lq, Nq, C1].
k: [B, Lk, Nk, C1].
v: [B, Lk, Nk, C2]. Nq must be divisible by Nk.
q_lens: [B].
k_lens: [B].
dropout_p: float. Dropout probability.
softmax_scale: float. The scaling of QK^T before applying softmax.
causal: bool. Whether to apply causal attention mask.
window_size: (left right). If not (-1, -1), apply sliding window local attention.
deterministic: bool. If True, slightly slower and uses more memory.
dtype: torch.dtype. Apply when dtype of q/k/v is not float16/bfloat16.
"""
half_dtypes = (torch.float16, torch.bfloat16)
#assert dtype in half_dtypes
#assert q.device.type == 'cuda' and q.size(-1) <= 256
# params
b, lq, lk, out_dtype = q.size(0), q.size(1), k.size(1), q.dtype
def half(x):
return x if x.dtype in half_dtypes else x.to(dtype)
# preprocess query
if q_lens is None:
q = half(q.flatten(0, 1))
q_lens = torch.tensor(
[lq] * b, dtype=torch.int32).to(
device=q.device, non_blocking=True)
else:
q = half(torch.cat([u[:v] for u, v in zip(q, q_lens)]))
# preprocess key, value
if k_lens is None:
k = half(k.flatten(0, 1))
v = half(v.flatten(0, 1))
k_lens = torch.tensor(
[lk] * b, dtype=torch.int32).to(
device=k.device, non_blocking=True)
else:
k = half(torch.cat([u[:v] for u, v in zip(k, k_lens)]))
v = half(torch.cat([u[:v] for u, v in zip(v, k_lens)]))
q = q.to(v.dtype)
k = k.to(v.dtype)
if q_scale is not None:
q = q * q_scale
if version is not None and version == 3 and not FLASH_ATTN_3_AVAILABLE:
log.warning('Flash attention 3 is not available, use flash attention 2 instead.')
# apply attention
if (version is None or version == 3) and FLASH_ATTN_3_AVAILABLE:
# Note: dropout_p, window_size are not supported in FA3 now.
x = flash_attn_interface.flash_attn_varlen_func(
q=q,
k=k,
v=v,
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
seqused_q=None,
seqused_k=None,
max_seqlen_q=lq,
max_seqlen_k=lk,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic)[0].unflatten(0, (b, lq))
else:
assert FLASH_ATTN_2_AVAILABLE
x = flash_attn.flash_attn_varlen_func(
q=q,
k=k,
v=v,
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
max_seqlen_q=lq,
max_seqlen_k=lk,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
causal=causal,
window_size=window_size,
deterministic=deterministic).unflatten(0, (b, lq))
# output
return x.type(out_dtype)
def attention(
q,
k,
v,
q_lens=None,
k_lens=None,
max_seqlen_q=None,
max_seqlen_k=None,
dropout_p=0.,
softmax_scale=None,
q_scale=None,
causal=False,
window_size=(-1, -1),
deterministic=False,
dtype=torch.bfloat16,
attention_mode='sdpa',
attn_mask=None,
):
def attention(q, k, v, q_lens=None, k_lens=None, max_seqlen_q=None, max_seqlen_k=None, dropout_p=0.,
softmax_scale=None, q_scale=None, causal=False, window_size=(-1, -1), deterministic=False, dtype=torch.bfloat16,
attention_mode='sdpa', attn_mask=None, transformer_options={}, frame_tokens=1536, heads=128):
if "flash" in attention_mode:
if attention_mode == 'flash_attn_2':
fa_version = 2
elif attention_mode == 'flash_attn_3':
fa_version = 3
return flash_attention(
q=q,
k=k,
v=v,
q_lens=q_lens,
k_lens=k_lens,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
q_scale=q_scale,
causal=causal,
window_size=window_size,
deterministic=deterministic,
dtype=dtype,
version=fa_version,
return flash_attention(q, k, v, q_lens=q_lens, k_lens=k_lens, dropout_p=dropout_p, softmax_scale=softmax_scale,
q_scale=q_scale, causal=causal, window_size=window_size, deterministic=deterministic, dtype=dtype, version=2 if attention_mode == 'flash_attn_2' else 3,
)
elif attention_mode == 'sdpa':
elif attention_mode == 'sageattn_3':
return sageattn_blackwell(q.transpose(1,2), k.transpose(1,2), v.transpose(1,2), per_block_mean=False).transpose(1,2).contiguous()
elif attention_mode == 'sageattn_varlen':
return sageattn_varlen_func(q,k,v, q_lens=q_lens, k_lens=k_lens, max_seqlen_k=max_seqlen_k, max_seqlen_q=max_seqlen_q)
elif attention_mode == 'sageattn_compiled': # for sage versions that allow torch.compile, may be redundant now as other sageattn ops are wrapper in custom ops
return sageattn_func_compiled(q, k, v, tensor_layout="NHD").contiguous()
elif attention_mode == 'sageattn':
return sageattn_func(q, k, v, tensor_layout="NHD").contiguous()
elif attention_mode == 'sageattn_ultravico':
return sageattn_func_ultravico([q, k, v], multi_factor=transformer_options.get("ultravico_alpha", 0.9), frame_tokens=frame_tokens).contiguous()
elif attention_mode == 'comfy':
return optimized_attention(q.transpose(1,2), k.transpose(1,2), v.transpose(1,2), heads=heads, skip_reshape=True)
else: # sdpa
if not (q.dtype == k.dtype == v.dtype):
return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2).to(q.dtype), v.transpose(1, 2).to(q.dtype), attn_mask=attn_mask).transpose(1, 2).contiguous()
return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), attn_mask=attn_mask).transpose(1, 2).contiguous()
elif attention_mode == 'sageattn_3':
return sageattn_blackwell(
q.transpose(1,2),
k.transpose(1,2),
v.transpose(1,2),
per_block_mean=False #seems necessary for reasonable VRAM usage, not sure of other implications
).transpose(1,2).contiguous()
elif attention_mode == 'sageattn_varlen':
return sageattn_varlen_func(
q,k,v,
q_lens=q_lens,
k_lens=k_lens,
max_seqlen_k=max_seqlen_k,
max_seqlen_q=max_seqlen_q
)
elif attention_mode == 'sageattn_compiled':
return sageattn_func_compiled(q, k, v, tensor_layout="NHD").contiguous()
else:
return sageattn_func(q, k, v, tensor_layout="NHD").contiguous()
+80
View File
@@ -0,0 +1,80 @@
import torch
from ...utils import log
def attention_func_error(*args, **kwargs):
raise ImportError("Selected attention mode not available. Please ensure required packages are installed correctly.")
try:
import flash_attn_interface
FLASH_ATTN_3_AVAILABLE = True
except Exception as e:
FLASH_ATTN_3_AVAILABLE = False
try:
import flash_attn
FLASH_ATTN_2_AVAILABLE = True
except Exception as e:
FLASH_ATTN_2_AVAILABLE = False
if not FLASH_ATTN_2_AVAILABLE and not FLASH_ATTN_3_AVAILABLE:
flash_attention = attention_func_error
else:
def flash_attention(q, k, v, q_lens=None, k_lens=None, dropout_p=0., softmax_scale=None, q_scale=None, causal=False, window_size=(-1, -1), deterministic=False, dtype=torch.bfloat16, version=None):
half_dtypes = (torch.float16, torch.bfloat16)
# params
b, lq, lk, out_dtype = q.size(0), q.size(1), k.size(1), q.dtype
def half(x):
return x if x.dtype in half_dtypes else x.to(dtype)
# preprocess query
if q_lens is None:
q = half(q.flatten(0, 1))
q_lens = torch.tensor(
[lq] * b, dtype=torch.int32).to(
device=q.device, non_blocking=True)
else:
q = half(torch.cat([u[:v] for u, v in zip(q, q_lens)]))
# preprocess key, value
if k_lens is None:
k = half(k.flatten(0, 1))
v = half(v.flatten(0, 1))
k_lens = torch.tensor(
[lk] * b, dtype=torch.int32).to(
device=k.device, non_blocking=True)
else:
k = half(torch.cat([u[:v] for u, v in zip(k, k_lens)]))
v = half(torch.cat([u[:v] for u, v in zip(v, k_lens)]))
q = q.to(v.dtype)
k = k.to(v.dtype)
if q_scale is not None:
q = q * q_scale
if version is not None and version == 3 and not FLASH_ATTN_3_AVAILABLE:
log.warning('Flash attention 3 is not available, use flash attention 2 instead.')
if (version is None or version == 3) and FLASH_ATTN_3_AVAILABLE:
# Note: dropout_p, window_size are not supported in FA3 now.
x = flash_attn_interface.flash_attn_varlen_func(q=q, k=k, v=v,
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
seqused_q=None, seqused_k=None, max_seqlen_q=lq, max_seqlen_k=lk,
softmax_scale=softmax_scale, causal=causal,
deterministic=deterministic).unflatten(0, (b, lq))
else:
assert FLASH_ATTN_2_AVAILABLE
x = flash_attn.flash_attn_varlen_func(q=q, k=k, v=v,
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
0, dtype=torch.int32).to(q.device, non_blocking=True),
max_seqlen_q=lq, max_seqlen_k=lk, dropout_p=dropout_p,
softmax_scale=softmax_scale, causal=causal, window_size=window_size,
deterministic=deterministic).unflatten(0, (b, lq))
return x.type(out_dtype)
+824 -589
View File
File diff suppressed because it is too large Load Diff
+73 -25
View File
@@ -1,12 +1,15 @@
import torch
from .fm_solvers import (FlowDPMSolverMultistepScheduler, get_sampling_sigmas, retrieve_timesteps)
import numpy as np
from .fm_solvers import (FlowDPMSolverMultistepScheduler)
from .fm_solvers_unipc import FlowUniPCMultistepScheduler
from .basic_flowmatch import FlowMatchScheduler
from .flowmatch_pusa import FlowMatchSchedulerPusa
from .flowmatch_res_multistep import FlowMatchSchedulerResMultistep
from .ersde_scheduler import ERSDEScheduler
from .scheduling_flow_match_lcm import FlowMatchLCMScheduler
from .fm_sa_ode import FlowMatchSAODEStableScheduler
from .fm_rcm import rCMFlowMatchScheduler
from .vitb_unipc import ViBTScheduler
from ...utils import log
try:
@@ -24,25 +27,35 @@ scheduler_list = [
"deis",
"lcm", "lcm/beta",
"res_multistep",
"er_sde",
"flowmatch_causvid",
"flowmatch_distill",
"flowmatch_pusa",
"multitalk",
"sa_ode_stable",
"rcm"
"rcm",
"vibt_unipc",
]
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, flowedit_args=None, denoise_strength=1.0, sigmas=None, log_timesteps=False, **kwargs):
def _apply_custom_sigmas(sample_scheduler, sigmas, device):
sample_scheduler.sigmas = sigmas.to(device)
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, denoise_strength=1.0, sigmas=None, log_timesteps=False, enhance_hf=False, **kwargs):
timesteps = None
if 'unipc' in scheduler:
if sigmas is not None:
steps = len(sigmas) - 1
if scheduler == 'vibt_unipc':
sample_scheduler = ViBTScheduler()
sample_scheduler.set_parameters(shift=shift)
sample_scheduler.set_timesteps(steps, device=device)
elif 'unipc' in scheduler:
sample_scheduler = FlowUniPCMultistepScheduler(shift=shift)
if sigmas is None:
sample_scheduler.set_timesteps(steps, device=device, shift=shift, use_beta_sigmas=('beta' in scheduler))
else:
sample_scheduler.sigmas = sigmas.to(device)
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif scheduler in ['euler/beta', 'euler', 'longcat_distill_euler']:
if 'longcat' in scheduler:
num_distill_sample_steps = 50
@@ -57,10 +70,10 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas)
else:
sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta'))
if flowedit_args: #seems to work better
timesteps, _ = retrieve_timesteps(sample_scheduler, device=device, sigmas=get_sampling_sigmas(steps, shift))
if sigmas is None:
sample_scheduler.set_timesteps(steps, device=device)
else:
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif 'dpm' in scheduler:
if 'sde' in scheduler:
algorithm_type = "sde-dpmsolver++"
@@ -70,16 +83,20 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
if sigmas is None:
sample_scheduler.set_timesteps(steps, device=device, use_beta_sigmas=('beta' in scheduler))
else:
sample_scheduler.sigmas = sigmas.to(device)
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif scheduler == 'deis':
sample_scheduler = DEISMultistepScheduler(use_flow_sigmas=True, prediction_type="flow_prediction", flow_shift=shift)
sample_scheduler.set_timesteps(steps, device=device)
sample_scheduler.sigmas[-1] = 1e-6
if sigmas is None:
sample_scheduler.set_timesteps(steps, device=device)
sample_scheduler.sigmas[-1] = 1e-6
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif 'lcm' in scheduler:
sample_scheduler = FlowMatchLCMScheduler(shift=shift, use_beta_sigmas=(scheduler == 'lcm/beta'))
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
if sigmas is None:
sample_scheduler.set_timesteps(steps, device=device)
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif 'flowmatch_causvid' in scheduler:
if sigmas is not None:
raise NotImplementedError("This scheduler does not support custom sigmas")
@@ -112,21 +129,50 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)])
elif 'flowmatch_pusa' in scheduler:
sample_scheduler = FlowMatchSchedulerPusa(shift=shift, sigma_min=0.0, extra_one_step=True)
sample_scheduler.set_timesteps(steps+1, denoising_strength=denoise_strength, shift=shift,
sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
if sigmas is None:
sample_scheduler.set_timesteps(steps+1, denoising_strength=denoise_strength, shift=shift)
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif scheduler == 'res_multistep':
sample_scheduler = FlowMatchSchedulerResMultistep(shift=shift)
sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
if sigmas is None:
sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength)
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif scheduler == 'er_sde':
sample_scheduler = ERSDEScheduler(shift=shift)
if sigmas is None:
sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength)
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif "sa_ode_stable" in scheduler:
sample_scheduler = FlowMatchSAODEStableScheduler(shift=shift, **kwargs)
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
if sigmas is None:
sample_scheduler.set_timesteps(steps, device=device)
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
elif 'rcm' in scheduler:
sample_scheduler = rCMFlowMatchScheduler()
sample_scheduler.set_timesteps(steps, sigma_max=120)
if sigmas is None:
sample_scheduler.set_timesteps(steps, sigma_max=120)
else:
_apply_custom_sigmas(sample_scheduler, sigmas, device)
if timesteps is None:
timesteps = sample_scheduler.timesteps
if enhance_hf:
num_tail_uniform_steps = max(3, min(15, int(len(timesteps) * 0.2))) # Use 20% of steps for uniform tail (minimum 3, maximum 15)
tail_uniform_start = float(timesteps.max()) * 0.5 # Split at 50% of the timestep range
tail_uniform_end = 0
timesteps_uniform_tail = list(np.linspace(tail_uniform_start, tail_uniform_end, num_tail_uniform_steps, dtype=np.float32, endpoint=(tail_uniform_end != 0)))
timesteps_uniform_tail = [torch.tensor(t, device=device).unsqueeze(0) for t in timesteps_uniform_tail]
filtered_timesteps = [timestep.unsqueeze(0).to(device) for timestep in timesteps if timestep > tail_uniform_start]
timesteps = torch.cat(filtered_timesteps + timesteps_uniform_tail)
sample_scheduler.timesteps = timesteps
sample_scheduler.sigmas = torch.cat([timesteps / 1000, torch.zeros(1, device=timesteps.device)])
steps = len(timesteps)
if (isinstance(start_step, int) and end_step != -1 and start_step >= end_step) or (not isinstance(start_step, int) and start_step != -1 and end_step >= start_step):
raise ValueError("start_step must be less than end_step")
@@ -136,7 +182,7 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
end_idx = len(timesteps) - 1
if log_timesteps:
log.info(f"------- Scheduler info -------")
log.info("------- Scheduler info -------")
log.info(f"Total timesteps: {timesteps}")
if isinstance(start_step, float):
@@ -156,6 +202,7 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
end_idx = end_step - 1
# Slice timesteps and sigmas once, based on indices
all_timesteps = timesteps
timesteps = timesteps[start_idx:end_idx+1]
sample_scheduler.full_sigmas = sample_scheduler.sigmas.clone()
sample_scheduler.sigmas = sample_scheduler.sigmas[start_idx:start_idx+len(timesteps)+1] # always one longer
@@ -163,9 +210,10 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
if log_timesteps:
log.info(f"Using timesteps: {timesteps}")
log.info(f"Using sigmas: {sample_scheduler.sigmas}")
log.info(f"------------------------------")
log.info("------------------------------")
if hasattr(sample_scheduler, 'timesteps'):
sample_scheduler.timesteps = timesteps
setattr(sample_scheduler, 'all_timesteps', all_timesteps)
return sample_scheduler, timesteps, start_idx, end_idx
return sample_scheduler, timesteps, start_idx, end_idx
+154
View File
@@ -0,0 +1,154 @@
import torch
class ERSDEScheduler():
"""Extended Reverse-Time SDE solver (VP ER-SDE-Solver-3).
Based on: arXiv: https://arxiv.org/abs/2309.06169
Code reference: https://github.com/QinpengCui/ER-SDE-Solver/blob/main/er_sde_solver.py
"""
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0,
sigma_max=1.0, sigma_min=0.003 / 1.002, max_stage=3, s_noise=1.0,
num_integration_points=200):
self.num_train_timesteps = num_train_timesteps
self.shift = shift
self.sigma_max = sigma_max
self.sigma_min = sigma_min
self.max_stage = max_stage
self.s_noise = s_noise
self.num_integration_points = num_integration_points
self.set_timesteps(num_inference_steps)
self.old_denoised = None
self.old_denoised_d = None
self.step_index = 0
def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, sigmas=None):
"""Generate the full sigma schedule (from max to min)."""
full_sigmas = torch.linspace(self.sigma_max, self.sigma_min, self.num_train_timesteps)
ss = len(full_sigmas) / num_inference_steps
if sigmas is None:
sigmas = []
for x in range(num_inference_steps):
idx = int(round(x * ss))
sigmas.append(float(full_sigmas[idx]))
sigmas.append(0.0)
self.sigmas = torch.FloatTensor(sigmas)
self.sigmas = self.shift * self.sigmas / (1 + (self.shift - 1) * self.sigmas)
self.timesteps = self.sigmas[:-1] * self.num_train_timesteps
self.step_index = 0
self.old_denoised = None
self.old_denoised_d = None
def default_er_sde_noise_scaler(self, x):
return x * ((x ** 0.3).exp() + 10.0)
def step(self, model_output, timestep, sample, generator):
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
self.sigmas = self.sigmas.to(model_output.device)
self.timesteps = self.timesteps.to(model_output.device)
if timestep.ndim == 0:
timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0)
else:
timestep_id = torch.argmin((self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
noise_scaler = self.default_er_sde_noise_scaler
# Get current and next sigma
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
if (timestep_id + 1 >= len(self.sigmas)).any():
sigma_next = torch.zeros_like(sigma)
else:
sigma_next = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
er_lambda_s = sigma
er_lambda_t = sigma_next
# Calculate alpha values
alpha_s = sigma / (er_lambda_s + 1e-10)
alpha_t = sigma_next / (er_lambda_t + 1e-10)
r_alpha = alpha_t / (alpha_s + 1e-10)
# Denoised prediction (x_0 estimate)
denoised = sample - sigma * model_output
# Determine which stage to use
stage_used = min(self.max_stage, self.step_index + 1)
if sigma_next == 0 or (sigma_next == 0.0).all():
# Final step - return denoised
x = denoised
else:
r = noise_scaler(er_lambda_t) / (noise_scaler(er_lambda_s) + 1e-10)
# Stage 1: Euler step
x = r_alpha * r * sample + alpha_t * (1 - r) * denoised
if stage_used >= 2 and self.old_denoised is not None:
dt = er_lambda_t - er_lambda_s
lambda_step_size = -dt / self.num_integration_points
# Create integration points
point_indice = torch.arange(0, self.num_integration_points,
dtype=torch.float32, device=sample.device)
lambda_pos = er_lambda_t + point_indice * lambda_step_size
scaled_pos = noise_scaler(lambda_pos)
# Stage 2: Second-order correction
s = torch.sum(1 / (scaled_pos + 1e-10)) * lambda_step_size
# Get previous sigma for derivative calculation
if timestep_id > 0:
sigma_prev = self.sigmas[timestep_id - 1].reshape(-1, 1, 1, 1)
er_lambda_prev = sigma_prev
else:
er_lambda_prev = er_lambda_s
denoised_d = (denoised - self.old_denoised) / ((er_lambda_s - er_lambda_prev) + 1e-10)
x = x + alpha_t * (dt + s * noise_scaler(er_lambda_t)) * denoised_d
if stage_used >= 3 and self.old_denoised_d is not None:
# Stage 3: Third-order correction
s_u = torch.sum((lambda_pos - er_lambda_s) / (scaled_pos + 1e-10)) * lambda_step_size
# Get sigma from two steps ago
if timestep_id > 1:
sigma_prev_prev = self.sigmas[timestep_id - 2].reshape(-1, 1, 1, 1)
er_lambda_prev_prev = sigma_prev_prev
else:
er_lambda_prev_prev = er_lambda_prev
denoised_u = (denoised_d - self.old_denoised_d) / (((er_lambda_s - er_lambda_prev_prev) / 2) + 1e-10)
x = x + alpha_t * ((dt ** 2) / 2 + s_u * noise_scaler(er_lambda_t)) * denoised_u
self.old_denoised_d = denoised_d
# Add stochastic noise
if self.s_noise > 0:
noise_term = (er_lambda_t ** 2 - er_lambda_s ** 2 * r ** 2).sqrt()
noise_term = torch.nan_to_num(noise_term, nan=0.0)
noise = torch.randn(*x.shape, dtype=torch.float32, device=torch.device("cpu"), generator=generator).to(x)
x = x + alpha_t * noise * self.s_noise * noise_term
# Store current denoised for next iteration
self.old_denoised = denoised
self.step_index += 1
return x
def add_noise(self, original_samples, noise, timestep):
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
self.sigmas = self.sigmas.to(noise.device)
self.timesteps = self.timesteps.to(noise.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
sample = (1 - sigma) * original_samples + sigma * noise
return sample.type_as(noise)
@@ -35,9 +35,7 @@ class FlowMatchSchedulerResMultistep():
self.sigmas = torch.FloatTensor(sigmas)
self.sigmas = self.shift * self.sigmas / \
(1 + (self.shift - 1) * self.sigmas)
self.timesteps = self.sigmas * self.num_train_timesteps
#print(f"Timesteps: {self.timesteps}, Sigmas: {self.sigmas}")
self.timesteps = self.sigmas[:-1] * self.num_train_timesteps
def step(self, model_output, timestep, sample):
if timestep.ndim == 2:
@@ -48,14 +46,14 @@ class FlowMatchSchedulerResMultistep():
timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0)
else:
timestep_id = torch.argmin((self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
sigma_prev = self.sigmas[timestep_id - 1].reshape(-1, 1, 1, 1) if timestep_id > 0 else sigma
if (timestep_id + 1 >= len(self.sigmas)).any():
sigma_next = torch.tensor(0)
else:
sigma_next = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
x0_pred = (sample - sigma * model_output)
if sigma_next == 0 or self.prev_model_output is None:
@@ -73,7 +71,7 @@ class FlowMatchSchedulerResMultistep():
self.old_sigma_next = sigma_next
self.prev_model_output = x0_pred
return x
def add_noise(self, original_samples, noise, timestep):
"""
+41
View File
@@ -0,0 +1,41 @@
from diffusers.schedulers import UniPCMultistepScheduler
import torch
class ViBTScheduler(UniPCMultistepScheduler):
def __init__(self, **kwargs):
super().__init__(**{**kwargs, "use_flow_sigmas": True})
self.set_parameters()
def set_parameters(self, noise_scale=1.0, shift=5.0, seed=None):
self.noise_scale = noise_scale
self.config.flow_shift = shift
def step(self, model_output, timestep, sample, generator, **kwargs):
delta_t = (
max(self.timesteps[self.timesteps < timestep]) - timestep
if any(self.timesteps < timestep)
else -timestep - 1
) / 1000
current_t = (timestep + 1) / 1000.0
eta = (-delta_t * (current_t + delta_t) / current_t) ** 0.5
noise = torch.randn(
sample.shape,
generator=generator,
device=torch.device("cpu"),
dtype=sample.dtype,
).to(sample.device)
latents = sample + delta_t * model_output + eta * self.noise_scale * noise
return (latents,)
@classmethod
def from_scheduler(
cls, scheduler: UniPCMultistepScheduler, noise_scale=1.0, shift_gamma=5.0
):
obj = cls.__new__(cls)
obj.__dict__ = scheduler.__dict__.copy()
obj.set_parameters(noise_scale, shift_gamma)
return obj
+145 -45
View File
@@ -5,6 +5,9 @@ import torch.nn as nn
import torch.nn.functional as F
from tqdm import tqdm
from comfy.utils import ProgressBar
from ..utils import print_memory, log
from comfy import model_management as mm
device = mm.get_torch_device()
import comfy.ops
ops = comfy.ops.disable_weight_init
@@ -254,10 +257,11 @@ class Resample38(Resample):
class ResidualBlock(nn.Module):
def __init__(self, in_dim, out_dim, dropout=0.0):
def __init__(self, in_dim, out_dim, dropout=0.0, cpu_cache=False):
super().__init__()
self.in_dim = in_dim
self.out_dim = out_dim
self.cpu_cache = cpu_cache
# layers
self.residual = nn.Sequential(
@@ -269,6 +273,12 @@ class ResidualBlock(nn.Module):
if in_dim != out_dim else nn.Identity()
def forward(self, x, feat_cache=None, feat_idx=[0]):
if self.cpu_cache:
return self._forward_cpu_cache(x, feat_cache, feat_idx)
else:
return self._forward(x, feat_cache, feat_idx)
def _forward(self, x, feat_cache=None, feat_idx=[0]):
h = self.shortcut(x)
for layer in self.residual:
if check_is_instance(layer, CausalConv3d) and feat_cache is not None:
@@ -288,6 +298,26 @@ class ResidualBlock(nn.Module):
x = layer(x)
return x + h
def _forward_cpu_cache(self, x, feat_cache=None, feat_idx=[0]):
h = self.shortcut(x)
for layer in self.residual:
if check_is_instance(layer, CausalConv3d) and feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
cached_frame = feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device)
cache_x = torch.cat([cached_frame, cache_x], dim=2)
prev_cache = feat_cache[idx].to(x.device) if feat_cache[idx] is not None else None
x = layer(x, prev_cache)
feat_cache[idx] = cache_x.to("cpu", non_blocking=True)
feat_idx[0] += 1
else:
x = layer(x)
return x + h
class AttentionBlock(nn.Module):
"""
@@ -507,7 +537,8 @@ class Encoder3d(nn.Module):
attn_scales=[],
temperal_downsample=[True, True, False],
dropout=0.0,
pruning_rate=0.0):
pruning_rate=0.0,
cpu_cache=False):
super().__init__()
self.dim = dim
self.z_dim = z_dim
@@ -515,6 +546,7 @@ class Encoder3d(nn.Module):
self.num_res_blocks = num_res_blocks
self.attn_scales = attn_scales
self.temperal_downsample = temperal_downsample
self.cpu_cache = cpu_cache
# dimensions
dims = [dim * u for u in [1] + dim_mult]
@@ -529,7 +561,7 @@ class Encoder3d(nn.Module):
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
# residual (+attention) blocks
for _ in range(num_res_blocks):
downsamples.append(ResidualBlock(in_dim, out_dim, dropout))
downsamples.append(ResidualBlock(in_dim, out_dim, dropout, cpu_cache=cpu_cache))
if scale in attn_scales:
downsamples.append(AttentionBlock(out_dim))
in_dim = out_dim
@@ -543,9 +575,9 @@ class Encoder3d(nn.Module):
self.downsamples = nn.Sequential(*downsamples)
# middle blocks
self.middle = nn.Sequential(ResidualBlock(out_dim, out_dim, dropout),
self.middle = nn.Sequential(ResidualBlock(out_dim, out_dim, dropout, cpu_cache=cpu_cache),
AttentionBlock(out_dim),
ResidualBlock(out_dim, out_dim, dropout))
ResidualBlock(out_dim, out_dim, dropout, cpu_cache=cpu_cache))
# output blocks
self.head = nn.Sequential(RMS_norm(out_dim, images=False), nn.SiLU(),
@@ -612,7 +644,8 @@ class Encoder3d_38(nn.Module):
attn_scales=[],
temperal_downsample=[False, True, True],
dropout=0.0,
pruning_rate=0.0):
pruning_rate=0.0,
cpu_cache=False):
super().__init__()
self.dim = dim
self.z_dim = z_dim
@@ -620,6 +653,7 @@ class Encoder3d_38(nn.Module):
self.num_res_blocks = num_res_blocks
self.attn_scales = attn_scales
self.temperal_downsample = temperal_downsample
self.cpu_cache = cpu_cache
# dimensions
dims = [dim * u for u in [1] + dim_mult]
@@ -650,9 +684,9 @@ class Encoder3d_38(nn.Module):
# middle blocks
self.middle = nn.Sequential(
ResidualBlock(out_dim, out_dim, dropout),
ResidualBlock(out_dim, out_dim, dropout, cpu_cache=cpu_cache),
AttentionBlock(out_dim),
ResidualBlock(out_dim, out_dim, dropout),
ResidualBlock(out_dim, out_dim, dropout, cpu_cache=cpu_cache),
)
# # output blocks
@@ -730,7 +764,8 @@ class Decoder3d(nn.Module):
attn_scales=[],
temperal_upsample=[False, True, True],
dropout=0.0,
pruning_rate=0.0):
pruning_rate=0.0,
cpu_cache=False):
super().__init__()
self.dim = dim
self.z_dim = z_dim
@@ -738,6 +773,7 @@ class Decoder3d(nn.Module):
self.num_res_blocks = num_res_blocks
self.attn_scales = attn_scales
self.temperal_upsample = temperal_upsample
self.cpu_cache = cpu_cache
# dimensions
dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
@@ -748,9 +784,9 @@ class Decoder3d(nn.Module):
self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
# middle blocks
self.middle = nn.Sequential(ResidualBlock(dims[0], dims[0], dropout),
self.middle = nn.Sequential(ResidualBlock(dims[0], dims[0], dropout, cpu_cache=cpu_cache),
AttentionBlock(dims[0]),
ResidualBlock(dims[0], dims[0], dropout))
ResidualBlock(dims[0], dims[0], dropout, cpu_cache=cpu_cache))
# upsample blocks
upsamples = []
@@ -759,7 +795,7 @@ class Decoder3d(nn.Module):
if i == 1 or i == 2 or i == 3:
in_dim = in_dim // 2
for _ in range(num_res_blocks + 1):
upsamples.append(ResidualBlock(in_dim, out_dim, dropout))
upsamples.append(ResidualBlock(in_dim, out_dim, dropout, cpu_cache=cpu_cache))
if scale in attn_scales:
upsamples.append(AttentionBlock(out_dim))
in_dim = out_dim
@@ -838,7 +874,8 @@ class Decoder3d_38(nn.Module):
attn_scales=[],
temperal_upsample=[False, True, True],
dropout=0.0,
pruning_rate=0.0):
pruning_rate=0.0,
cpu_cache=False):
super().__init__()
self.dim = dim
self.z_dim = z_dim
@@ -855,9 +892,9 @@ class Decoder3d_38(nn.Module):
self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
# middle blocks
self.middle = nn.Sequential(ResidualBlock(dims[0], dims[0], dropout),
self.middle = nn.Sequential(ResidualBlock(dims[0], dims[0], dropout, cpu_cache=cpu_cache),
AttentionBlock(dims[0]),
ResidualBlock(dims[0], dims[0], dropout))
ResidualBlock(dims[0], dims[0], dropout, cpu_cache=cpu_cache))
# upsample blocks
upsamples = []
@@ -951,7 +988,9 @@ class VideoVAE_(nn.Module):
dropout=0.0,
mean=None,
inv_std=None,
pruning_rate=0.0):
pruning_rate=0.0,
cpu_cache=False,
verbose=False):
super().__init__()
self.dim = dim
self.z_dim = z_dim
@@ -962,14 +1001,15 @@ class VideoVAE_(nn.Module):
self.temperal_upsample = temperal_downsample[::-1]
self.mean = mean
self.inv_std = inv_std
self.verbose = verbose
# modules
self.encoder = Encoder3d(dim, z_dim * 2, dim_mult, num_res_blocks,
attn_scales, self.temperal_downsample, dropout, pruning_rate)
attn_scales, self.temperal_downsample, dropout, pruning_rate, cpu_cache=cpu_cache)
self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1)
self.conv2 = CausalConv3d(z_dim, z_dim, 1)
self.decoder = Decoder3d(dim, z_dim, dim_mult, num_res_blocks,
attn_scales, self.temperal_upsample, dropout, pruning_rate)
attn_scales, self.temperal_upsample, dropout, pruning_rate, cpu_cache=cpu_cache)
def forward(self, x):
mu, log_var = self.encode(x)
@@ -1016,10 +1056,15 @@ class VideoVAE_(nn.Module):
def encode(self, x, pbar=True, sample=False):
t = x.shape[2]
iter_ = 1 + (t - 1) // 4
input_shape = x.shape
if pbar:
pbar = ProgressBar(iter_)
try:
torch.cuda.reset_peak_memory_stats(device)
except:
pass
for i in range(iter_):
for i in tqdm(range(iter_), desc="WanVAE encoding frames", disable=not pbar):
self._enc_conv_idx = [0]
if i == 0:
out = self.encoder(x[:, :, :1, :, :],
@@ -1042,15 +1087,22 @@ class VideoVAE_(nn.Module):
std = torch.exp(0.5 * log_var.clamp(-30.0, 20.0))
eps = torch.randn_like(std)
return mu + std * eps
if self.verbose:
try:
log.info(f"WanVAE encoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE encode")
torch.cuda.reset_peak_memory_stats(device)
except:
pass
return mu
#modification originally by @raindrop313 https://github.com/raindrop313/ComfyUI-WanVideoStartEndFrames
def decode_2(self, z):
# z: [b,c,t,h,w]
z = z / self.inv_std.to(z) + self.mean.to(z)
iter_ = z.shape[2]
z_head=z[:,:,:-1,:,:]
z_tail=z[:,:,-1,:,:].unsqueeze(2)
@@ -1065,12 +1117,12 @@ class VideoVAE_(nn.Module):
out_ = self.decoder(x[:, :, -1, :, :].unsqueeze(2),
feat_cache=None,
feat_idx=self._conv_idx)
out = torch.cat([out, out_], 2) # may add tensor offload
out = torch.cat([out, out_], 2)
else:
out_ = self.decoder(x[:, :, i:i + 1, :, :],
feat_cache=self._feat_map,
feat_idx=self._conv_idx)
out = torch.cat([out, out_], 2) # may add tensor offload
out = torch.cat([out, out_], 2)
self.clear_cache()
return out
@@ -1079,11 +1131,16 @@ class VideoVAE_(nn.Module):
def decode(self, z, pbar=True):
# z: [b,c,t,h,w]
z = z / self.inv_std.to(z) + self.mean.to(z)
input_shape = z.shape
iter_ = z.shape[2]
if pbar:
pbar = ProgressBar(iter_)
try:
torch.cuda.reset_peak_memory_stats(device)
except:
pass
x = self.conv2(z)
for i in range(iter_):
for i in tqdm(range(iter_), desc="WanVAE decoding frames", disable=not pbar):
self._conv_idx = [0]
if i == 0:
out = self.decoder(x[:, :, i:i + 1, :, :],
@@ -1093,13 +1150,20 @@ class VideoVAE_(nn.Module):
out_ = self.decoder(x[:, :, i:i + 1, :, :],
feat_cache=self._feat_map,
feat_idx=self._conv_idx)
out = torch.cat([out, out_], 2) # may add tensor offload
out = torch.cat([out, out_], 2)
if pbar:
pbar.update(1)
if pbar:
pbar.update_absolute(0)
self.clear_cache()
if self.verbose:
try:
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device)
except:
pass
return out
def reparameterize(self, mu, log_var):
@@ -1126,11 +1190,12 @@ class VideoVAE_(nn.Module):
class WanVideoVAE(nn.Module):
def __init__(self, z_dim=16, dtype=torch.float32, pruning_rate=0.0):
def __init__(self, z_dim=16, dtype=torch.float32, pruning_rate=0.0, cpu_cache=False, verbose=False):
super().__init__()
self.dtype = dtype
self.cpu_cache = cpu_cache
self.verbose = verbose
mean = [
-0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508,
0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921
@@ -1144,7 +1209,7 @@ class WanVideoVAE(nn.Module):
self.z_dim = z_dim
# init model
self.model = VideoVAE_(z_dim=z_dim, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate).eval().requires_grad_(False)
self.model = VideoVAE_(z_dim=z_dim, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=self.cpu_cache, verbose=self.verbose).eval().requires_grad_(False)
self.upsampling_factor = 8
@@ -1170,7 +1235,7 @@ class WanVideoVAE(nn.Module):
return mask
def tiled_decode(self, hidden_states, device, tile_size, tile_stride, pbar=True):
def tiled_decode(self, hidden_states, device, tile_size, tile_stride, end_=False, pbar=True):
_, _, T, H, W = hidden_states.shape
size_h, size_w = tile_size
stride_h, stride_w = tile_stride
@@ -1187,14 +1252,20 @@ class WanVideoVAE(nn.Module):
data_device = "cpu"
computation_device = device
out_T = T * 4 - 3
weight = torch.zeros((1, 1, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
values = torch.zeros((1, 3, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
weight, values = None, None
if pbar:
pbar = ProgressBar(len(tasks))
for h, h_, w, w_ in tqdm(tasks, desc="VAE decoding"):
hidden_states_batch = hidden_states[:, :, :, h:h_, w:w_].to(computation_device)
hidden_states_batch = self.model.decode(hidden_states_batch).to(data_device)
if end_:
hidden_states_batch = self.model.decode_2(hidden_states_batch).to(data_device)
else:
hidden_states_batch = self.model.decode(hidden_states_batch).to(data_device)
if weight is None:
weight = torch.zeros((1, 1, hidden_states_batch.shape[2], H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
if values is None:
values = torch.zeros((1, 3, hidden_states_batch.shape[2], H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
mask = self.build_mask(
hidden_states_batch,
@@ -1227,7 +1298,7 @@ class WanVideoVAE(nn.Module):
def tiled_encode(self, video, device, tile_size, tile_stride, end_=False, pbar=True):
_, _, T, H, W = video.shape
if tile_size is None and tile_stride is None:
size_h, size_w = H //2, W // 2
stride_h, stride_w = size_h // 2, size_w // 2
@@ -1338,7 +1409,7 @@ class WanVideoVAE(nn.Module):
for hidden_state in hidden_states:
hidden_state = hidden_state.unsqueeze(0)
if tiled:
video = self.tiled_decode(hidden_state, device, tile_size, tile_stride, pbar=pbar)
video = self.tiled_decode(hidden_state, device, tile_size, tile_stride, end_=end_, pbar=pbar)
else:
if end_:
video = self.double_decode(hidden_state, device)
@@ -1363,7 +1434,9 @@ class VideoVAE38_(VideoVAE_):
dtype=torch.bfloat16,
mean=None,
inv_std=None,
pruning_rate=0.0):
pruning_rate=0.0,
cpu_cache=False,
verbose=False):
super(VideoVAE_, self).__init__()
self.dim = dim
self.z_dim = z_dim
@@ -1375,24 +1448,30 @@ class VideoVAE38_(VideoVAE_):
self.dtype = dtype
self.mean = mean
self.inv_std = inv_std
self.cpu_cache = cpu_cache
self.verbose = verbose
# modules
self.encoder = Encoder3d_38(dim, z_dim * 2, dim_mult, num_res_blocks,
attn_scales, self.temperal_downsample, dropout, pruning_rate)
attn_scales, self.temperal_downsample, dropout, pruning_rate, cpu_cache=cpu_cache)
self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1)
self.conv2 = CausalConv3d(z_dim, z_dim, 1)
self.decoder = Decoder3d_38(dec_dim, z_dim, dim_mult, num_res_blocks,
attn_scales, self.temperal_upsample, dropout, pruning_rate)
attn_scales, self.temperal_upsample, dropout, pruning_rate, cpu_cache=cpu_cache)
def encode(self, x, pbar=True, sample=False):
input_shape = x.shape
self.clear_cache()
try:
torch.cuda.reset_peak_memory_stats(device)
except:
pass
x = patchify(x, patch_size=2)
t = x.shape[2]
iter_ = 1 + (t - 1) // 4
if pbar:
pbar = ProgressBar(iter_)
for i in range(iter_):
for i in tqdm(range(iter_), desc="WanVAE encoding frames", disable=not pbar):
self._enc_conv_idx = [0]
if i == 0:
out = self.encoder(x[:, :, :1, :, :],
@@ -1408,18 +1487,30 @@ class VideoVAE38_(VideoVAE_):
mu = self.conv1(out).chunk(2, dim=1)[0]
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
self.clear_cache()
if self.verbose:
try:
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device)
except:
pass
return mu
def decode(self, z, pbar=True):
self.clear_cache()
input_shape = z.shape
try:
torch.cuda.reset_peak_memory_stats(device)
except:
pass
z = z / self.inv_std.to(z) + self.mean.to(z)
iter_ = z.shape[2]
if pbar:
pbar = ProgressBar(iter_)
x = self.conv2(z)
for i in range(iter_):
for i in tqdm(range(iter_), desc="WanVAE decoding frames", disable=not pbar):
self._conv_idx = [0]
if i == 0:
out = self.decoder(x[:, :, i:i + 1, :, :],
@@ -1435,12 +1526,19 @@ class VideoVAE38_(VideoVAE_):
pbar.update(1)
out = unpatchify(out, patch_size=2)
self.clear_cache()
if self.verbose:
try:
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device)
except:
pass
return out
class WanVideoVAE38(WanVideoVAE):
def __init__(self, z_dim=48, dim=160, dtype=torch.bfloat16, pruning_rate=0.0):
def __init__(self, z_dim=48, dim=160, dtype=torch.bfloat16, pruning_rate=0.0, cpu_cache=False, verbose=False):
super(WanVideoVAE, self).__init__()
mean = [
@@ -1463,7 +1561,9 @@ class WanVideoVAE38(WanVideoVAE):
self.inv_std = (1.0 / torch.tensor(std)).view(1, z_dim, 1, 1, 1)
self.dtype = dtype
self.z_dim = z_dim
self.cpu_cache = cpu_cache
self.verbose = verbose
# init model
self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate).eval().requires_grad_(False)
self.upsampling_factor = 16
self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=cpu_cache, verbose=verbose).eval().requires_grad_(False)
self.upsampling_factor = 16