142 Commits
Author SHA1 Message Date
kijai af696b02f3 update 2025-06-15 19:01:09 +03:00
kijai 307437e495 Update nodes.py 2025-06-15 02:15:57 +03:00
kijai 985268bb69 update 2025-06-14 17:55:26 +03:00
kijai e5bbb95804 Update nodes.py 2025-06-12 16:39:10 +03:00
kijai 8410a6977b init 2025-06-12 16:31:48 +03:00
kijai ced1ddaa1a Update basic_flowmatch.py 2025-06-12 16:29:53 +03:00
kijai 5917f51837 Small adjustments and options to dwpose detection 2025-06-12 16:29:31 +03:00
kijai 6eddec54a6 fix 2025-06-09 23:18:01 +03:00
kijai 9b380ec3c0 Allow loading state dicts from .pt sub dicts 2025-06-09 22:33:07 +03:00
kijai 39412cf422 Node to use ATI on native workflows 2025-06-09 04:24:37 +03:00
kijai 820ac3008e Update nodes.py 2025-06-07 21:48:50 +03:00
kijai d3a09a1a6b bump version 2025-06-07 11:57:38 +03:00
kijai 87a980f6bf This can still be append 2025-06-06 12:10:42 +03:00
kijai a290f75b5e Update model.py 2025-06-05 23:41:17 +03:00
kijai 9c27705dac Update nodes.py 2025-06-05 23:04:40 +03:00
kijai efb87445d5 Update nodes.py 2025-06-05 19:28:56 +03:00
kijai 6139017535 Update wanvideo_ATI_testing_01.json 2025-06-05 18:44:37 +03:00
kijai 45dca5c41c Fix ATI start/end percent controls 2025-06-04 10:02:20 +03:00
kijai bdd3828596 Update nodes.py 2025-06-03 20:47:27 +03:00
kijai 821881429b Update nodes.py 2025-06-03 10:03:58 +03:00
kijai bd6c853051 Better track parsing 2025-06-03 08:11:40 +03:00
kijai f4a1157f71 Fix typo 2025-06-03 08:04:08 +03:00
kijai e2f42b5773 realisdance 2025-06-03 08:02:59 +03:00
kijai 9e8978731b Update wanvideo_ATI_testing_01.json 2025-05-31 20:46:44 +03:00
kijai 7d63a771d8 Better ATI visualization 2025-05-31 20:44:33 +03:00
kijai da3c7def22 Update wanvideo_ATI_testing_01.json 2025-05-31 19:42:05 +03:00
kijai abf61ad833 More efficient and compatible tensor reshape for ATI motion_patch 2025-05-31 19:30:06 +03:00
kijai 774b055452 Create wanvideo_ATI_testing_01.json 2025-05-31 02:07:59 +03:00
kijai 8bc74daad9 cleanup 2025-05-31 01:58:21 +03:00
kijai cd2884d88a Initial ATI support
https://github.com/bytedance/ATI
2025-05-31 01:17:18 +03:00
kijai c4a8f6d835 Clearer error when trying to use uni3c controlnet with incompatible model 2025-05-30 18:11:36 +03:00
kijai 5c676d3a49 Allow further filtering of LoRA layers to apply 2025-05-30 16:24:44 +03:00
kijai ef40577b70 fixes 2025-05-30 12:29:38 +03:00
kijai fb73022a06 controlnet device fix 2025-05-29 18:42:46 +03:00
kijai 0323a9d2c7 Fix TAEW decoding 2025-05-29 18:37:13 +03:00
kijai 129f368380 Fix uni3c on fp16 and add strength parameters 2025-05-29 13:25:52 +03:00
kijai 87ae18e203 Error on trying to use VACE with I2V models. 2025-05-29 00:14:35 +03:00
kijai 07c7fc6c2a Don't show Phantom references in live preview 2025-05-27 19:14:22 +03:00
kijai 81e7022ab5 Force steps to match possible cfg schedule length 2025-05-27 18:06:36 +03:00
kijai 934dc127a8 Update wanvideo_phantom_subject2vid_example_01.json 2025-05-27 17:32:05 +03:00
kijai ea01ae9ed9 Allow scheduling phantom cfg too 2025-05-27 15:27:38 +03:00
kijai 0f7de6bac2 Fix 14B controlnet loading 2025-05-26 23:42:26 +03:00
kijai 23e3368de3 Initial Dilated Controlnet support
https://huggingface.co/TheDenk/wan2.1-t2v-1.3b-controlnet-hed-v1
https://huggingface.co/TheDenk/wan2.1-t2v-14b-controlnet-hed-v1
2025-05-26 21:54:48 +03:00
kijai ece2917a41 Fix compile 2025-05-26 15:49:03 +03:00
kijai ef1ed29178 partial Uni3C implementation 2025-05-26 12:35:36 +03:00
Jukka Seppänen d9ca90c1f0 Merge pull request #545 from anastasiuspernat/main
Fix chained VACE embeddings ignored with WanVideo Context Options
2025-05-26 00:46:53 +03:00
Anastasiy Safari 75de496d6a Merge branch 'kijai:main' into main 2025-05-25 14:39:41 -07:00
kijai fe7a3d5b46 Forgot to remove these 2025-05-23 18:44:59 +03:00
kijai 370837233a Fix more VACEs than 2....
silly mistake, more than 2 never worked
2025-05-21 19:15:30 +03:00
kijai 7875031efe Support list input for VACE strength
Allows finer control of VACE strength over timesteps
2025-05-21 19:04:25 +03:00
kijai fa7015315a Possibly reduce TeaCache VRAM use a bit
Uses 500-700MB less at 1280x720
2025-05-20 21:07:15 +03:00
kijai d547482ad6 Update nodes.py 2025-05-18 21:09:36 +03:00
kijai fd562e9730 Update nodes.py 2025-05-16 13:17:11 +03:00
kijai 7bba1e8e17 typo 2025-05-16 11:03:26 +03:00
kijai 1c3060f59f remove torchao quants, fix long unianimate crash 2025-05-15 20:25:43 +03:00
kijai 4494bda603 bump version 2025-05-15 12:47:23 +03:00
kijai ec9f310486 update VACE example 2025-05-15 11:09:32 +03:00
kijai 6aec1d97a8 better error for VACE module loading 2025-05-14 18:23:37 +03:00
kijai 555c67a8cd Support 14B VACE 2025-05-14 17:33:01 +03:00
kijai 3cb4902ef1 Update nodes.py 2025-05-14 11:55:16 +03:00
kijai 1e6a780729 Update nodes.py 2025-05-14 11:40:37 +03:00
kijai 5aed18d658 remove print 2025-05-14 11:27:28 +03:00
kijai 1701375214 Fix TAEW preview 2025-05-14 11:26:54 +03:00
kijai abdac5ec3a Update model.py 2025-05-14 10:31:44 +03:00
kijai bfed2afc2e Add flex_attention for CausVid 2025-05-13 21:35:34 +03:00
Anastasiy Safari ab870a198f Fix chained VACE embeddings ignored with WanVideo Context Options 2025-05-10 13:35:43 -07:00
kijai f2bc29b931 Fix TeaCache when not using coefficients 2025-05-05 21:58:15 +03:00
kijai 5f7164788e Don't use matplotlib for dwpose 2025-05-05 10:42:55 +03:00
kijai 03edd8cbf2 Context windows with camera control 2025-05-01 18:58:44 +03:00
kijai 77f78ceded Support Fun Control-Camera 2025-04-29 20:41:27 +03:00
kijai 4a41fb0aaf cleanup experimental context window stuff
Didn't really work anyway
2025-04-29 17:36:07 +03:00
kijai b78505014c Fix Phantom and FantasyTalking when not using TeaCache 2025-04-29 10:31:52 +03:00
kijai e8074288df Update wanvideo_I2V_FantasyTalking_example_01.json 2025-04-29 09:58:12 +03:00
kijai e3afc7fc75 Add helper node for cfg scheduling and allow scheduling audio_cfg as well 2025-04-29 09:00:01 +03:00
kijai 1b3592f6b4 Update nodes.py 2025-04-29 07:34:36 +03:00
kijai ed068e1fe3 Fix lowvram lora load when not using quantization 2025-04-29 02:21:11 +03:00
kijai fe25d21fba Fix audio cfg
silly mistake
2025-04-29 02:07:05 +03:00
kijai 2d0e0fe698 Update wanvideo_I2V_FantasyTalking_example_01.json 2025-04-29 01:46:53 +03:00
kijai ca0974d5f0 cleanup 2025-04-29 01:46:39 +03:00
kijai a9ec21bb8f Update wanvideo_I2V_FantasyTalking_example_01.json 2025-04-29 00:19:59 +03:00
kijai a7fa683c2c Update wanvideo_I2V_FantasyTalking_example_01.json 2025-04-29 00:07:25 +03:00
kijai 1f6fd1e5a8 Don't download redundant wave2vec model files.
.bin and .h5 can be safely deleted
2025-04-28 23:24:01 +03:00
kijai e334f78124 Update nodes.py 2025-04-28 21:50:57 +03:00
kijai c8a084d27a Update nodes.py 2025-04-28 21:37:48 +03:00
kijai 8117c6f033 Context windows with FantasyTalking 2025-04-28 20:24:41 +03:00
kijai d96787b312 Update wanvideo_I2V_FantasyTalking_example_01.json 2025-04-28 19:27:49 +03:00
kijai 4e562b5d92 Add FantasyTalking example 2025-04-28 19:14:03 +03:00
kijai 5683d8306a Support FantasyTalking 2025-04-28 19:10:32 +03:00
kijai df95c85283 Fix sageattn when using fp32 2025-04-27 17:31:06 +03:00
kijai 4fc4159c88 fix up default rope dtype 2025-04-27 17:22:02 +03:00
kijai fc7ab666a2 support loading Fun camera model
Input not functional yet
2025-04-25 17:02:10 +03:00
kijai caebbafab8 version after merge 2025-04-25 16:21:19 +03:00
kijai 03eeef32a3 Merge branch 'dev' 2025-04-25 16:20:59 +03:00
kijai c09b981d1e version 2025-04-25 16:20:50 +03:00
kijai e75e6d6933 Update nodes.py 2025-04-25 16:19:52 +03:00
kijai 031d8ac817 disable enhance always with DF 2025-04-25 15:47:25 +03:00
kijai 33b8f1c082 Support Fun control 1.1
Supports the new reference image input to the 1.1 control model
2025-04-25 15:44:18 +03:00
kijai 863829d083 dtype fixes and maybe allow fp8_fast to work
Unsure when this has changed but seems that _scaled_mm now works with same dtype? This fixes the quality degradation of fp8_fast in initial tests.
2025-04-25 03:12:10 +03:00
kijai c7b2635f01 Update model.py 2025-04-23 20:37:48 +03:00
kijai 96a2172e13 Phantom + VACE testing 2025-04-23 20:34:28 +03:00
kijai b3f3b0cd93 Update model.py 2025-04-23 20:11:08 +03:00
kijai 4602b9c885 Fix phantom start percent 2025-04-23 10:30:33 +03:00
kijai e0e5fcf713 Add start/end percent for Phantom and other small fixes 2025-04-23 01:00:41 +03:00
kijai 2b0bf44994 Update wanvideo_phantom_subject2vid_example_01.json 2025-04-22 22:42:58 +03:00
kijai 24876cffc1 Update nodes.py 2025-04-22 22:34:45 +03:00
kijai 1f53574387 Create wanvideo_phantom_subject2vid_example_01.json 2025-04-22 21:15:54 +03:00
kijai 6fadcbd957 Support Phantom, refactor model dtypes, reduce DF model memory use 2025-04-22 21:13:27 +03:00
kijai a623f87dca Update wanvideo_skyreels_diffusion_forcing_extension_example_01.json 2025-04-22 16:05:26 +03:00
kijai 6099ad393b Update wanvideo_skyreels_diffusion_forcing_extension_example_01.json 2025-04-22 10:24:26 +03:00
kijai 949c887e0c Fix FLF2V when using vram management node 2025-04-22 09:44:46 +03:00
kijai 5109a74839 Update wanvideo_skyreels_diffusion_forcing_extension_example_01.json 2025-04-22 08:13:26 +03:00
kijai 6ba7bc811e Update nodes.py 2025-04-22 02:05:08 +03:00
kijai f93e54d21b Allow unianimate with DF 2025-04-22 01:01:22 +03:00
kijai 3ae0bc1fec Create wanvideo_skyreels_diffusion_forcing_extension_example_01.json 2025-04-22 00:29:57 +03:00
kijai e3ea2bf392 Allow VACE with DF 2025-04-21 20:01:20 +03:00
kijai 884069121f Fix latent preview fallback if disabled 2025-04-21 19:11:16 +03:00
kijai f371c0a736 Allow TeaCache for DF
Guess it works
2025-04-21 18:30:40 +03:00
kijai 65f5505fca Initial Skyreels DiffusionForcing model support
On it's own sampler at least for now while I figure out how it all works.
2025-04-21 18:13:51 +03:00
kijai e5a326c981 Use penultimate_hidden_states instead when using native clip vision
The new Skyreels 1.3B I2V did not work otherwise, this seems to make it match the wrapper clip vision output better
2025-04-21 01:33:51 +03:00
kijai 7d005201a2 Allow loading UniAnimate LoRA with low vram load as well 2025-04-20 18:05:43 +03:00
kijai 604f0e2714 Handle non-detected pose frames 2025-04-20 13:11:27 +03:00
kijai 04cc8fa80e Update nodes.py 2025-04-20 12:53:30 +03:00
kijai 18aa47cc74 Update nodes.py 2025-04-19 22:41:50 +03:00
kijai 20088d0fbf Update nodes.py 2025-04-19 22:38:19 +03:00
kijai 3692f3580a Update nodes.py 2025-04-19 21:15:24 +03:00
kijai 52d6f00770 Update nodes.py 2025-04-19 20:20:42 +03:00
kijai b5b4e44512 Make dwpose keypoint size adjustable 2025-04-19 20:08:30 +03:00
kijai 7b6c34e26d Add options what to draw on dwpose detector 2025-04-19 19:47:09 +03:00
kijai b9b8ec7bf8 Apply Fresca for cfg 1.0 as well 2025-04-19 17:50:17 +03:00
kijai 2cfae9fa8f Make UniAnimate ref pose optional 2025-04-19 17:41:04 +03:00
kijai 90c9415d31 Add lcm scheduler, pad/truncate UniAnim poses if pose count doesn't match latent count 2025-04-19 17:24:18 +03:00
kijai 421e375f13 Strength controls for Unianimate 2025-04-18 23:06:49 +03:00
kijai 19044adc78 Update nodes.py 2025-04-18 20:45:16 +03:00
kijai 3813c615d3 Make ref pose optional 2025-04-18 20:07:39 +03:00
kijai 6d521a9f4e Update nodes.py 2025-04-18 19:45:04 +03:00
kijai 7462356743 Support UniAnimate-DiT
https://github.com/ali-vilab/UniAnimate-DiT
2025-04-18 19:43:40 +03:00
kijai 00de4f5e0e Detect (a guess for now) TeaCache coefficients for SkyreelsV2 2025-04-18 11:47:19 +03:00
kijai cc8450b1a7 ReCamMaster custom orbit cam generation node 2025-04-17 02:48:54 +03:00
kijai 8257cd1f8a Fix TeaCache end step 2025-04-17 01:18:42 +03:00
kijai 761d188dba Update nodes.py 2025-04-16 23:07:16 +03:00
kijai 425fcfcf2c Create wanvideo_FLF2V_720P_example_01.json 2025-04-16 20:05:55 +03:00
kijai 7c81d7ce10 Better support for the official FirstLastFrame2Video -model (FLF2V)
https://huggingface.co/Kijai/WanVideo_comfy/blob/main/Wan2_1-FLF2V-14B-720P_fp16.safetensors
https://huggingface.co/Kijai/WanVideo_comfy/blob/main/Wan2_1-FLF2V-14B-720P_fp8_e4m3fn.safetensors
2025-04-16 19:40:04 +03:00
48 changed files with 27040 additions and 3119 deletions
+2 -1
View File
@@ -9,4 +9,5 @@ logs/
.idea
tools/
.vscode/
convert_*
convert_*
*.pt
+42
View File
@@ -0,0 +1,42 @@
# Copyright (c) 2024-2025 Bytedance Ltd. and/or its affiliates
#
# 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.
from typing import Dict, List, Optional, Tuple, Union
import numpy as np
import torch
def process_tracks(tracks_np: np.ndarray, frame_size: Tuple[int, int], quant_multi: int = 8, **kwargs):
# tracks: shape [t, h, w, 3] => samples align with 24 fps, model trained with 16 fps.
# frame_size: tuple (W, H)
tracks = torch.from_numpy(tracks_np).float()
if tracks.shape[1] == 121:
tracks = torch.permute(tracks, (1, 0, 2, 3))
tracks, visibles = tracks[..., :2], tracks[..., 2:3]
short_edge = min(*frame_size)
tracks = tracks - torch.tensor([*frame_size]).type_as(tracks) / 2
tracks = tracks / short_edge * 2
visibles = visibles * 2 - 1
trange = torch.linspace(-1, 1, tracks.shape[0]).view(-1, 1, 1, 1).expand(*visibles.shape)
out_ = torch.cat([trange, tracks, visibles], dim=-1).view(121, -1, 4)
out_0 = out_[:1]
out_l = out_[1:] # 121 => 120 | 1
out_l = torch.repeat_interleave(out_l, 2, dim=0)[1::3] # 120 => 240 => 80
return torch.cat([out_0, out_l], dim=0)
+142
View File
@@ -0,0 +1,142 @@
# Copyright (c) 2024-2025 Bytedance Ltd. and/or its affiliates
#
# 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.
from typing import List, Optional, Tuple, Union
import torch
# Refer to https://github.com/Angtian/VoGE/blob/main/VoGE/Utils.py
def ind_sel(target: torch.Tensor, ind: torch.Tensor, dim: int = 1):
"""
:param target: [... (can be k or 1), n > M, ...]
:param ind: [... (k), M]
:param dim: dim to apply index on
:return: sel_target [... (k), M, ...]
"""
assert (
len(ind.shape) > dim
), "Index must have the target dim, but get dim: %d, ind shape: %s" % (dim, str(ind.shape))
target = target.expand(
*tuple(
[ind.shape[k] if target.shape[k] == 1 else -1 for k in range(dim)]
+ [
-1,
]
* (len(target.shape) - dim)
)
)
ind_pad = ind
if len(target.shape) > dim + 1:
for _ in range(len(target.shape) - (dim + 1)):
ind_pad = ind_pad.unsqueeze(-1)
ind_pad = ind_pad.expand(*(-1,) * (dim + 1), *target.shape[(dim + 1) : :])
return torch.gather(target, dim=dim, index=ind_pad)
def merge_final(vert_attr: torch.Tensor, weight: torch.Tensor, vert_assign: torch.Tensor):
"""
:param vert_attr: [n, d] or [b, n, d] color or feature of each vertex
:param weight: [b(optional), w, h, M] weight of selected vertices
:param vert_assign: [b(optional), w, h, M] selective index
:return:
"""
target_dim = len(vert_assign.shape) - 1
if len(vert_attr.shape) == 2:
assert vert_attr.shape[0] > vert_assign.max()
# [n, d] ind: [b(optional), w, h, M]-> [b(optional), w, h, M, d]
# sel_attr = ind_sel(
# vert_attr[(None,) * target_dim], vert_assign.type(torch.long), dim=target_dim
# )
new_shape = [1] * target_dim + list(vert_attr.shape)
tensor = vert_attr.reshape(new_shape)
sel_attr = ind_sel(tensor, vert_assign.type(torch.long), dim=target_dim)
else:
assert vert_attr.shape[1] > vert_assign.max()
#sel_attr = ind_sel(
# vert_attr[:, *(None,) * (target_dim - 1)], vert_assign.type(torch.long), dim=target_dim
#)
new_shape = [vert_attr.shape[0]] + [1] * (target_dim - 1) + list(vert_attr.shape[1:])
tensor = vert_attr.reshape(new_shape)
sel_attr = ind_sel(tensor, vert_assign.type(torch.long), dim=target_dim)
# [b(optional), w, h, M]
final_attr = torch.sum(sel_attr * weight.unsqueeze(-1), dim=-2)
return final_attr
def patch_motion(
tracks: torch.FloatTensor, # (B, T, N, 4)
vid: torch.FloatTensor, # (C, T, H, W)
temperature: float = 220.0,
vae_divide: tuple = (4, 16),
topk: int = 2,
):
with torch.no_grad():
_, T, H, W = vid.shape
N = tracks.shape[2]
_, tracks, visible = torch.split(
tracks, [1, 2, 1], dim=-1
) # (B, T, N, 2) | (B, T, N, 1)
tracks_n = tracks / torch.tensor([W / min(H, W), H / min(H, W)], device=tracks.device)
tracks_n = tracks_n.clamp(-1, 1)
visible = visible.clamp(0, 1)
xx = torch.linspace(-W / min(H, W), W / min(H, W), W)
yy = torch.linspace(-H / min(H, W), H / min(H, W), H)
grid = torch.stack(torch.meshgrid(yy, xx, indexing="ij")[::-1], dim=-1).to(
tracks.device
)
tracks_pad = tracks[:, 1:]
visible_pad = visible[:, 1:]
visible_align = visible_pad.view(T - 1, 4, *visible_pad.shape[2:]).sum(1)
tracks_align = (tracks_pad * visible_pad).view(T - 1, 4, *tracks_pad.shape[2:]).sum(
1
) / (visible_align + 1e-5)
dist_ = (
(tracks_align[:, None, None] - grid[None, :, :, None]).pow(2).sum(-1)
) # T, H, W, N
weight = torch.exp(-dist_ * temperature) * visible_align.clamp(0, 1).view(
T - 1, 1, 1, N
)
vert_weight, vert_index = torch.topk(
weight, k=min(topk, weight.shape[-1]), dim=-1
)
grid_mode = "bilinear"
point_feature = torch.nn.functional.grid_sample(
vid[vae_divide[0]:].permute(1, 0, 2, 3)[:1],
tracks_n[:, :1].type(vid.dtype),
mode=grid_mode,
padding_mode="zeros",
align_corners=False,
)
point_feature = point_feature.squeeze(0).squeeze(1).permute(1, 0) # N, C=16
out_feature = merge_final(point_feature, vert_weight, vert_index).permute(3, 0, 1, 2) # T - 1, H, W, C => C, T - 1, H, W
out_weight = vert_weight.sum(-1) # T - 1, H, W
# out feature -> already soft weighted
mix_feature = out_feature + vid[vae_divide[0]:, 1:] * (1 - out_weight.clamp(0, 1))
out_feature_full = torch.cat([vid[vae_divide[0]:, :1], mix_feature], dim=1) # C, T, H, W
out_mask_full = torch.cat([torch.ones_like(out_weight[:1]), out_weight], dim=0) # T, H, W
return torch.cat([out_mask_full[None].expand(vae_divide[0], -1, -1, -1), out_feature_full], dim=0)
+329
View File
@@ -0,0 +1,329 @@
import json
from .motion import process_tracks
import numpy as np
from typing import List, Tuple
import torch
FIXED_LENGTH = 121
def pad_pts(tr):
"""Convert list of {x,y} to (FIXED_LENGTH,1,3) array, padding/truncating."""
pts = np.array([[p['x'], p['y'], 1] for p in tr], dtype=np.float32)
n = pts.shape[0]
if n < FIXED_LENGTH:
pad = np.zeros((FIXED_LENGTH - n, 3), dtype=np.float32)
pts = np.vstack((pts, pad))
else:
pts = pts[:FIXED_LENGTH]
return pts.reshape(FIXED_LENGTH, 1, 3)
def age_to_bgr(ratio: float) -> Tuple[int,int,int]:
"""
Map ratio∈[0,1] through: 0→blue, 1/3→green, 2/3→yellow, 1→red.
Returns (B,G,R) for OpenCV.
"""
if ratio <= 1/3:
# blue→green
t = ratio / (1/3)
b = int(255 * (1 - t))
g = int(255 * t)
r = 0
elif ratio <= 2/3:
# green→yellow
t = (ratio - 1/3) / (1/3)
b = 0
g = 255
r = int(255 * t)
else:
# yellow→red
t = (ratio - 2/3) / (1/3)
b = 0
g = int(255 * (1 - t))
r = 255
return (r, g, b)
def paint_point_track(
frames: np.ndarray,
point_tracks: np.ndarray,
visibles: np.ndarray,
min_radius: int = 1,
max_radius: int = 6,
max_retain: int = 50
) -> np.ndarray:
"""
Draws every past point of each track on each frame, with radius and color
interpolated by the point's age (old→small to new→large).
Args:
frames: [F, H, W, 3] uint8 RGB
point_tracks:[N, F, 2] float32 – (x,y) in pixel coords
visibles: [N, F] bool – visibility mask
min_radius: radius for the very first point (oldest)
max_radius: radius for the current point (newest)
Returns:
video: [F, H, W, 3] uint8 RGB
"""
import cv2
num_points, num_frames = point_tracks.shape[:2]
H, W = frames.shape[1:3]
video = frames.copy()
for t in range(num_frames):
# start from the original frame
frame = video[t].copy()
for i in range(num_points):
# draw every past step τ = 0..t
for τ in range(t + 1):
if not visibles[i, τ]:
continue
if t - τ > max_retain:
continue
# sub-pixel offset + clamp
x, y = point_tracks[i, τ] + 0.5
xi = int(np.clip(x, 0, W - 1))
yi = int(np.clip(y, 0, H - 1))
# age‐ratio in [0,1]
if num_frames > 1:
ratio = 1 - float(t - τ) / max_retain
else:
ratio = 1.0
# interpolated radius
radius = int(round(min_radius + (max_radius - min_radius) * ratio))
# OpenCV draws in BGR order:
color_rgb = age_to_bgr(ratio)
# filled circle
cv2.circle(frame, (xi, yi), radius, color_rgb, thickness=-1)
video[t] = frame
return video
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
print(f"Warning: Unexpected track format: {type(tracks_data[0])}")
except json.JSONDecodeError as e:
print(f"Error parsing tracks JSON: {e}")
tracks_data = []
return tracks_data
class WanVideoATITracks:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": ("WANVIDEOMODEL", ),
"tracks": ("STRING",),
"width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the image to encode"}),
"height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the image to encode"}),
"temperature": ("FLOAT", {"default": 220.0, "min": 0.0, "max": 1000.0, "step": 0.1}),
"topk": ("INT", {"default": 2, "min": 1, "max": 10, "step": 1}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the steps to apply ATI"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the steps to apply ATI"}),
},
}
RETURN_TYPES = ("WANVIDEOMODEL",)
RETURN_NAMES = ("model",)
FUNCTION = "patchmodel"
CATEGORY = "WanVideoWrapper"
def patchmodel(self, model, tracks, width, height, temperature, topk, start_percent, end_percent):
tracks_data = parse_json_tracks(tracks)
arrs = []
for track in tracks_data:
pts = pad_pts(track)
arrs.append(pts)
tracks_np = np.stack(arrs, axis=0)
processed_tracks = process_tracks(tracks_np, (width, height))
patcher = model.clone()
patcher.model_options["transformer_options"]["ati_tracks"] = processed_tracks.unsqueeze(0)
patcher.model_options["transformer_options"]["ati_temperature"] = temperature
patcher.model_options["transformer_options"]["ati_topk"] = topk
patcher.model_options["transformer_options"]["ati_start_percent"] = start_percent
patcher.model_options["transformer_options"]["ati_end_percent"] = end_percent
return (patcher,)
class WanVideoATITracksVisualize:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"images": ("IMAGE",),
"tracks": ("STRING",),
"min_radius": ("INT", {"default": 1, "min": 0, "max": 100, "step": 1, "tooltip": "radius for the very first point (oldest)"}),
"max_radius": ("INT", {"default": 6, "min": 0, "max": 100, "step": 1, "tooltip": "radius for the current point (newest)"}),
"max_retain": ("INT", {"default": 50, "min": 0, "max": 100, "step": 1, "tooltip": "Maximum number of points to retain"}),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("images",)
FUNCTION = "patchmodel"
CATEGORY = "WanVideoWrapper"
def patchmodel(self, images, tracks, min_radius, max_radius, max_retain):
tracks_data = parse_json_tracks(tracks)
arrs = []
for track in tracks_data:
pts = pad_pts(track)
arrs.append(pts)
tracks_np = np.stack(arrs, axis=0)
track = np.repeat(tracks_np, 2, axis=1)[:, ::3]
points = track[:, :, 0, :2].astype(np.float32)
visibles = track[:, :, 0, 2].astype(np.float32)
if images.shape[0] < points.shape[1]:
repeat_count = (points.shape[1] + images.shape[0] - 1) // images.shape[0]
images = images.repeat(repeat_count, 1, 1, 1)
images = images[:points.shape[1]]
elif images.shape[0] > points.shape[1]:
images = images[:points.shape[1]]
video_viz = paint_point_track(images.cpu().numpy(), points, visibles, min_radius, max_radius, max_retain)
video_viz = torch.from_numpy(video_viz).float()
return (video_viz,)
from comfy import utils
import types
from .motion_patch import patch_motion
class WanConcatCondPatch:
def __init__(self, tracks, temperature, topk):
self.tracks = tracks
self.temperature = temperature
self.topk = topk
def __get__(self, obj, objtype=None):
# Create bound method with stored parameters
def wrapped_concat_cond(self_module, *args, **kwargs):
return modified_concat_cond(self_module, self.tracks, self.temperature, self.topk, *args, **kwargs)
return types.MethodType(wrapped_concat_cond, obj)
def modified_concat_cond(self, tracks, temperature, topk, **kwargs):
noise = kwargs.get("noise", None)
extra_channels = self.diffusion_model.patch_embedding.weight.shape[1] - noise.shape[1]
if extra_channels == 0:
return None
image = kwargs.get("concat_latent_image", None)
device = kwargs["device"]
if image is None:
shape_image = list(noise.shape)
shape_image[1] = extra_channels
image = torch.zeros(shape_image, dtype=noise.dtype, layout=noise.layout, device=noise.device)
else:
image = utils.common_upscale(image.to(device), noise.shape[-1], noise.shape[-2], "bilinear", "center")
for i in range(0, image.shape[1], 16):
image[:, i: i + 16] = self.process_latent_in(image[:, i: i + 16])
image = utils.resize_to_batch_size(image, noise.shape[0])
if not self.image_to_video or extra_channels == image.shape[1]:
return image
if image.shape[1] > (extra_channels - 4):
image = image[:, :(extra_channels - 4)]
mask = kwargs.get("concat_mask", kwargs.get("denoise_mask", None))
if mask is None:
mask = torch.zeros_like(noise)[:, :4]
else:
if mask.shape[1] != 4:
mask = torch.mean(mask, dim=1, keepdim=True)
mask = 1.0 - mask
mask = utils.common_upscale(mask.to(device), noise.shape[-1], noise.shape[-2], "bilinear", "center")
if mask.shape[-3] < noise.shape[-3]:
mask = torch.nn.functional.pad(mask, (0, 0, 0, 0, 0, noise.shape[-3] - mask.shape[-3]), mode='constant', value=0)
if mask.shape[1] == 1:
mask = mask.repeat(1, 4, 1, 1, 1)
mask = utils.resize_to_batch_size(mask, noise.shape[0])
image_cond = torch.cat((mask, image), dim=1)
image_cond_ati = patch_motion(tracks.to(image_cond.device, image_cond.dtype), image_cond[0],
temperature=temperature, topk=topk)
return image_cond_ati.unsqueeze(0)
class WanVideoATI_comfy:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": ("MODEL", ),
"width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the image to encode"}),
"height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the image to encode"}),
"tracks": ("STRING",),
"temperature": ("FLOAT", {"default": 220.0, "min": 0.0, "max": 1000.0, "step": 0.1}),
"topk": ("INT", {"default": 2, "min": 1, "max": 10, "step": 1}),
},
}
RETURN_TYPES = ("MODEL",)
RETURN_NAMES = ("model", )
FUNCTION = "patchcond"
CATEGORY = "WanVideoWrapper"
def patchcond(self, model, tracks, width, height, temperature, topk):
tracks_data = parse_json_tracks(tracks)
arrs = []
for track in tracks_data:
pts = pad_pts(track)
arrs.append(pts)
tracks_np = np.stack(arrs, axis=0)
processed_tracks = process_tracks(tracks_np, (width, height))
model_clone = model.clone()
model_clone.add_object_patch(
"concat_cond",
WanConcatCondPatch(
processed_tracks.unsqueeze(0), temperature, topk
).__get__(model.model, model.model.__class__)
)
return (model_clone,)
NODE_CLASS_MAPPINGS = {
"WanVideoATITracks": WanVideoATITracks,
"WanVideoATITracksVisualize": WanVideoATITracksVisualize,
"WanVideoATI_comfy": WanVideoATI_comfy,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoATITracks": "WanVideo ATI Tracks",
"WanVideoATITracksVisualize": "WanVideo ATI Tracks Visualize",
"WanVideoATI_comfy": "WanVideo ATI Comfy",
}
+29 -1
View File
@@ -1,7 +1,35 @@
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
from .recammaster.nodes import NODE_CLASS_MAPPINGS as RECAM_MASTER_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS
from .unianimate.nodes import NODE_CLASS_MAPPINGS as UNIANIMATE_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS
from .skyreels.nodes import NODE_CLASS_MAPPINGS as SKYREELS_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as SKYREELS_NODE_DISPLAY_NAME_MAPPINGS
from .fantasytalking.nodes import NODE_CLASS_MAPPINGS as FANTASYTALKING_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FANTASYTALKING_NODE_DISPLAY_NAME_MAPPINGS
from .fun_camera.nodes import NODE_CLASS_MAPPINGS as FUN_CAMERA_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FUN_CAMERA_NODE_DISPLAY_NAME_MAPPINGS
from .uni3c.nodes import NODE_CLASS_MAPPINGS as UNI3C_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNI3C_NODE_DISPLAY_NAME_MAPPINGS
from .controlnet.nodes import NODE_CLASS_MAPPINGS as CONTROLNET_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as CONTROLNET_NODE_DISPLAY_NAME_MAPPINGS
from .ATI.nodes import NODE_CLASS_MAPPINGS as ATI_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as ATI_NODE_DISPLAY_NAME_MAPPINGS
from .causvid.nodes import NODE_CLASS_MAPPINGS as CAUSVID_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as CAUSVID_NODE_DISPLAY_NAME_MAPPINGS
NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(FANTASYTALKING_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(FUN_CAMERA_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(UNI3C_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(CONTROLNET_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(ATI_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(CAUSVID_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(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(CAUSVID_NODE_DISPLAY_NAME_MAPPINGS)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+612
View File
@@ -0,0 +1,612 @@
import os
import torch
import gc
from ..utils import log, print_memory, fourier_filter
import math
from tqdm import tqdm
from ..wanvideo.modules.model import rope_params
from ..wanvideo.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from ..wanvideo.utils.scheduling_flow_match_lcm import FlowMatchLCMScheduler
from ..wanvideo.utils.basic_flowmatch import FlowMatchScheduler
from ..nodes import optimized_scale
from einops import rearrange
from ..enhance_a_video.globals import disable_enhance
import comfy.model_management as mm
from comfy.utils import load_torch_file, ProgressBar, common_upscale
from comfy.clip_vision import clip_preprocess, ClipVisionModel
from comfy.cli_args import args, LatentPreviewMethod
script_directory = os.path.dirname(os.path.abspath(__file__))
#region Sampler
class WanVideoCausVidSampler:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("WANVIDEOMODEL",),
"text_embeds": ("WANVIDEOTEXTEMBEDS", ),
"image_embeds": ("WANVIDIMAGE_EMBEDS", ),
"steps": ("INT", {"default": 30, "min": 1}),
"shift": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 1000.0, "step": 0.01}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"force_offload": ("BOOLEAN", {"default": True, "tooltip": "Moves the model to the offload device after sampling"}),
"scheduler": ([
"flowmatch_causvid", "flowmatch_causvid_14b", "flowmatch_causvid_self_forcing",
#"unipc", "unipc/beta", "euler", "euler/beta", "lcm", "lcm/beta"
],
{
"default": 'flowmatch_causvid'
}),
"kv_cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}),
},
"optional": {
"samples": ("LATENT", {"tooltip": "init Latents to use for video2video process"} ),
"prefix_samples": ("LATENT", {"tooltip": "prefix latents"} ),
"denoise_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"rope_function": (["default", "comfy"], {"default": "default", "tooltip": "Comfy's RoPE implementation doesn't use complex numbers and can thus be compiled, that should be a lot faster when using torch.compile"}),
"experimental_args": ("EXPERIMENTALARGS", ),
}
}
RETURN_TYPES = ("LATENT", )
RETURN_NAMES = ("samples",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def _initialize_kv_cache(self, batch_size, dtype, device, num_blocks=30, num_heads=12):
"""
Initialize a Per-GPU KV cache for the Wan model.
"""
kv_cache1 = []
for _ in range(num_blocks):
kv_cache1.append({
"k": torch.zeros([batch_size, self.cache_window_size, num_heads, 128], dtype=dtype, device=device),
"v": torch.zeros([batch_size, self.cache_window_size, num_heads, 128], dtype=dtype, device=device),
"global_end_index": torch.tensor([0], dtype=torch.long, device=device),
"local_end_index": torch.tensor([0], dtype=torch.long, device=device)
})
self.kv_cache1 = kv_cache1 # always store the clean cache
def _initialize_crossattn_cache(self, batch_size, dtype, device, num_blocks=30, num_heads=12):
"""
Initialize a Per-GPU cross-attention cache for the Wan model.
"""
crossattn_cache = []
for _ in range(num_blocks):
crossattn_cache.append({
"k": torch.zeros([batch_size, 512, num_heads, 128], dtype=dtype, device=device),
"v": torch.zeros([batch_size, 512, num_heads, 128], dtype=dtype, device=device),
"is_init": False
})
self.crossattn_cache = crossattn_cache
def _shift_kv_cache(self):
"""
Shift the KV cache left by shift_blocks * num_frame_per_block * frame_seq_length.
This is called when kv_start exceeds window_size.
The first block is preserved, and shifting starts from the second block.
"""
shift_length = self.shift_blocks * self.num_frame_per_block * self.frame_seq_length
for block in self.kv_cache1:
block["k"] = torch.roll(block["k"], shifts=-shift_length, dims=1)
block["v"] = torch.roll(block["v"], shifts=-shift_length, dims=1)
# Clear the shifted-out part (except the first block)
block["k"][:, -shift_length:] = 0
block["v"][:, -shift_length:] = 0
# Update kv_start
self.kv_start -= shift_length
return shift_length
def process(self, model, text_embeds, image_embeds, shift, steps, seed, scheduler, kv_cache_device,
force_offload=True, samples=None, prefix_samples=None, denoise_strength=1.0, rope_function="default",
experimental_args=None):
#assert not (context_options and teacache_args), "Context options cannot currently be used together with teacache."
patcher = model
model = model.model
transformer = model.diffusion_model
dtype = model["dtype"]
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
if kv_cache_device == "main_device":
cache_device = mm.get_torch_device()
else:
cache_device = mm.unet_offload_device()
steps = int(steps/denoise_strength)
timesteps = None
if 'unipc' in scheduler:
sample_scheduler = FlowUniPCMultistepScheduler(shift=shift)
sample_scheduler.set_timesteps(steps, device=device, shift=shift, use_beta_sigmas=('beta' in scheduler))
elif 'euler' in scheduler:
sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta'))
sample_scheduler.set_timesteps(steps, device=device)
elif 'lcm' in scheduler:
sample_scheduler = FlowMatchLCMScheduler(shift=shift, use_beta_sigmas=(scheduler == 'lcm/beta'))
sample_scheduler.set_timesteps(steps, device=device)
elif 'flowmatch_causvid' in scheduler:
sample_scheduler = FlowMatchScheduler(
shift=shift, sigma_min=0.0, extra_one_step=True
)
sample_scheduler.set_timesteps(1000, training=True)
denoising_step_list = torch.tensor([1000, 757, 522], dtype=torch.long)
if "warp" in scheduler or "self_forcing" in scheduler:
denoising_step_list = torch.tensor([1000, 750, 500, 250] , dtype=torch.long)
timesteps = torch.cat((sample_scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32)))
denoising_step_list = timesteps[1000 - denoising_step_list]
elif "14b" in scheduler:
denoising_step_list = torch.tensor([1000, 934, 862, 756, 603, 410, 250, 140, 74], dtype=torch.long)
# sample_scheduler = FlowMatchScheduler(num_inference_steps=steps, shift=shift, sigma_min=0, extra_one_step=True)
# sample_scheduler.timesteps = torch.tensor(denoising_step_list).to(device)
# sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)])
#print(sample_scheduler.sigmas)
timesteps = denoising_step_list
#timesteps = torch.tensor(denoising_list).to(device)
print("timesteps", timesteps)
if denoise_strength < 1.0:
steps = int(steps * denoise_strength)
timesteps = timesteps[-(steps + 1):]
seed_g = torch.Generator(device=torch.device("cpu"))
seed_g.manual_seed(seed)
clip_fea, clip_fea_neg = None, None
vace_data, vace_context, vace_scale = None, None, None
image_cond = image_embeds.get("image_embeds", None)
target_shape = image_embeds.get("target_shape", None)
if target_shape is None:
raise ValueError("Empty image embeds must be provided for T2V (Text to Video")
has_ref = image_embeds.get("has_ref", False)
vace_context = image_embeds.get("vace_context", None)
vace_scale = image_embeds.get("vace_scale", None)
vace_start_percent = image_embeds.get("vace_start_percent", 0.0)
vace_end_percent = image_embeds.get("vace_end_percent", 1.0)
vace_seqlen = image_embeds.get("vace_seq_len", None)
vace_additional_embeds = image_embeds.get("additional_vace_inputs", [])
if vace_context is not None:
vace_data = [
{"context": vace_context,
"scale": vace_scale,
"start": vace_start_percent,
"end": vace_end_percent,
"seq_len": vace_seqlen
}
]
if len(vace_additional_embeds) > 0:
for i in range(len(vace_additional_embeds)):
if vace_additional_embeds[i].get("has_ref", False):
has_ref = True
vace_data.append({
"context": vace_additional_embeds[i]["vace_context"],
"scale": vace_additional_embeds[i]["vace_scale"],
"start": vace_additional_embeds[i]["vace_start_percent"],
"end": vace_additional_embeds[i]["vace_end_percent"],
"seq_len": vace_additional_embeds[i]["vace_seq_len"]
})
noise = torch.randn(
target_shape[0],
target_shape[1] + 1 if has_ref else target_shape[1],
target_shape[2],
target_shape[3],
dtype=torch.float32,
device=torch.device("cpu"),
generator=seed_g)
noise = noise.to(device, dtype)
latent_video_length = noise.shape[1]
if samples is not None:
input_samples = samples["samples"].squeeze(0).to(noise)
if input_samples.shape[1] != noise.shape[1]:
input_samples = torch.cat([input_samples[:, :1].repeat(1, noise.shape[1] - input_samples.shape[1], 1, 1), input_samples], dim=1)
original_image = input_samples.to(device)
if denoise_strength < 1.0:
latent_timestep = timesteps[:1].to(noise)
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
mask = samples.get("mask", None)
if mask is not None:
if mask.shape[2] != noise.shape[1]:
mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - mask.shape[2], noise.shape[2], noise.shape[3]), mask], dim=2)
init_latents = noise.to(device)
fps_embeds = None
if hasattr(transformer, "fps_embedding"):
fps = round(fps, 2)
log.info(f"Model has fps embedding, using {fps} fps")
fps_embeds = [fps]
fps_embeds = [0 if i == 16 else 1 for i in fps_embeds]
prefix_video = prefix_samples["samples"].to(noise) if prefix_samples is not None else None
prefix_video_latent_length = prefix_video.shape[2] if prefix_video is not None else 0
if prefix_video is not None:
log.info(f"Prefix video of length: {prefix_video_latent_length}")
init_latents[:, :prefix_video_latent_length] = prefix_video[0]
disable_enhance() #not sure if this can work, disabling for now to avoid errors if it's enabled by another sampler
freqs = None
transformer.rope_embedder.k = None
transformer.rope_embedder.num_frames = None
if rope_function=="comfy":
transformer.rope_embedder.k = 0
transformer.rope_embedder.num_frames = latent_video_length
else:
d = transformer.dim // transformer.num_heads
freqs = torch.cat([
rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=0),
rope_params(1024, 2 * (d // 6)),
rope_params(1024, 2 * (d // 6))
],
dim=1)
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1])
log.info(f"Seq len: {seq_len}")
seq_len = latent_video_length * 1560
log.info(f"Seq len: {seq_len}")
if args.preview_method in [LatentPreviewMethod.Auto, LatentPreviewMethod.Latent2RGB]: #default for latent2rgb
from latent_preview import prepare_callback
else:
from ..latent_preview import prepare_callback #custom for tiny VAE previews
#blockswap init
transformer_options = patcher.model_options.get("transformer_options", None)
if transformer_options is not None:
block_swap_args = transformer_options.get("block_swap_args", None)
if block_swap_args is not None:
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", True)
for name, param in transformer.named_parameters():
if "block" not in name:
param.data = param.data.to(device)
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking)
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking)
transformer.block_swap(
block_swap_args["blocks_to_swap"] - 1 ,
block_swap_args["offload_txt_emb"],
block_swap_args["offload_img_emb"],
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
)
elif model["auto_cpu_offload"]:
for module in transformer.modules():
if hasattr(module, "offload"):
module.offload()
if hasattr(module, "onload"):
module.onload()
elif model["manual_offloading"]:
transformer.to(device)
use_fresca = False
if experimental_args is not None:
video_attention_split_steps = experimental_args.get("video_attention_split_steps", [])
if video_attention_split_steps:
transformer.video_attention_split_steps = [int(x.strip()) for x in video_attention_split_steps.split(",")]
else:
transformer.video_attention_split_steps = []
use_zero_init = experimental_args.get("use_zero_init", True)
use_cfg_zero_star = experimental_args.get("cfg_zero_star", False)
zero_star_steps = experimental_args.get("zero_star_steps", 0)
use_fresca = experimental_args.get("use_fresca", False)
if use_fresca:
fresca_scale_low = experimental_args.get("fresca_scale_low", 1.0)
fresca_scale_high = experimental_args.get("fresca_scale_high", 1.25)
fresca_freq_cutoff = experimental_args.get("fresca_freq_cutoff", 20)
#region model pred
def model_pred(z, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
vace_data=None, unianim_data=None, teacache_state=None, kv_cache=None, crossattn_cache=None, current_kv_cache_start=0, kv_start=0, kv_end=0):
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])):
nonlocal patcher
current_step_percentage = idx / len(timesteps)
control_lora_enabled = False
image_cond_input = image_cond
base_params = {
'seq_len': seq_len,
'device': device,
'freqs': freqs,
't': timestep,
'current_step': idx,
'control_lora_enabled': control_lora_enabled,
'vace_data': vace_data,
'unianim_data': unianim_data,
'kv_cache': kv_cache,
'crossattn_cache': crossattn_cache,
'current_kv_cache_start': current_kv_cache_start,
"kv_start": kv_start,
"kv_end": kv_end
}
#cond
noise_pred_cond, teacache_state_cond = transformer(
[z], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None,
clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage,
pred_id=teacache_state[0] if teacache_state else None,
**base_params
)
noise_pred_cond = noise_pred_cond[0].to(intermediate_device)
if use_fresca:
noise_pred_cond = fourier_filter(
noise_pred_cond,
scale_low=fresca_scale_low,
scale_high=fresca_scale_high,
freq_cutoff=fresca_freq_cutoff,
)
return noise_pred_cond, [teacache_state_cond]
def convert_flow_pred_to_x0(flow_pred: torch.Tensor, xt: torch.Tensor, timestep: torch.Tensor) -> torch.Tensor:
"""
Convert flow matching's prediction to x0 prediction.
flow_pred: the prediction with shape [B, C, H, W]
xt: the input noisy data with shape [B, C, H, W]
timestep: the timestep with shape [B]
pred = noise - x0
x_t = (1-sigma_t) * x0 + sigma_t * noise
we have x0 = x_t - sigma_t * pred
see derivations https://chatgpt.com/share/67bf8589-3d04-8008-bc6e-4cf1a24e2d0e
"""
# use higher precision for calculations
original_dtype = flow_pred.dtype
flow_pred, xt, sigmas, timesteps = map(
lambda x: x.double().to(flow_pred.device), [flow_pred, xt,
sample_scheduler.sigmas,
sample_scheduler.timesteps]
)
timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
x0_pred = xt - sigma_t * flow_pred
return x0_pred.to(original_dtype)
log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {init_latents.shape[3]*8}x{init_latents.shape[2]*8} with {steps} steps")
intermediate_device = device
#clear memory before sampling
mm.unload_all_models()
mm.soft_empty_cache()
gc.collect()
try:
torch.cuda.reset_peak_memory_stats(device)
except:
pass
#main loop
self.num_frame_per_block = 3
num_frames = noise.shape[1]
assert num_frames % self.num_frame_per_block == 0
num_blocks = num_frames // self.num_frame_per_block
print("num_blocks: ", num_blocks)
context_noise = 0
self.frame_seq_length = 1560
print("frame_seq_length: ", self.frame_seq_length)
self.cache_window_size = self.frame_seq_length * num_frames
self.shift_blocks = 1
self.kv_start = 0
self.kv_end = 0
output_latents = torch.zeros(
(target_shape[0],
target_shape[1],
target_shape[2],
target_shape[3]), device=device, dtype=dtype)
print("output_latents shape: ", output_latents.shape)
# Step 1: Initialize KV cache to all zeros
self._initialize_kv_cache(
batch_size=1,
dtype=noise.dtype,
device=cache_device,
num_blocks=transformer.num_layers,
num_heads=transformer.num_heads,
)
self._initialize_crossattn_cache(
batch_size=1,
dtype=noise.dtype,
device=cache_device,
num_blocks=transformer.num_layers,
num_heads=transformer.num_heads,
)
# Step 2: Cache context feature
current_kv_cache_start_frame = 0
num_input_frames = 0
# Step 3: Temporal denoising loop
all_num_frames = [self.num_frame_per_block] * num_blocks
print("all_num_frames", all_num_frames)
pbar = ProgressBar(num_blocks)
callback = prepare_callback(patcher, num_blocks)
for i,current_num_frames in enumerate(all_num_frames):
print("current_kv_cache_start_frame: ", current_kv_cache_start_frame)
#noisy_input = noise[:, current_kv_cache_start_frame - num_input_frames:current_kv_cache_start_frame + current_num_frames - num_input_frames]
noisy_input = noise[:, i * self.num_frame_per_block:(i + 1) * self.num_frame_per_block]
print("noisy_input shape: ", noisy_input.shape)
kv_end = self.kv_start + self.num_frame_per_block * self.frame_seq_length
print("kv_end: ", kv_end)
# Spatial denoising loop
for step_index, current_timestep in enumerate(timesteps):
print(f"current_timestep: {current_timestep}")
# set current timestep
timestep = torch.ones(
[1, current_num_frames],
device=noise.device,
dtype=torch.int64) * current_timestep
if step_index < len(timesteps) - 1:
flow_pred, self.teacache_state = model_pred(
noisy_input.to(dtype),
text_embeds["prompt_embeds"],
text_embeds["negative_prompt_embeds"],
timestep, step_index, image_cond, clip_fea,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_kv_cache_start=current_kv_cache_start_frame * self.frame_seq_length,
kv_start = self.kv_start,
kv_end = kv_end
)
#print("noise_pred shape: ", noise_pred.shape)
denoised_pred = convert_flow_pred_to_x0(
flow_pred=flow_pred.transpose(0, 1),
xt=noisy_input.transpose(0, 1),
timestep=timestep.flatten(0, 1)
)
next_timestep = timesteps[step_index + 1]
print("step_index: ", step_index, "next_timestep: ", next_timestep)
noisy_input = sample_scheduler.add_noise(
denoised_pred,
torch.randn_like(denoised_pred),
next_timestep * torch.ones(
[current_num_frames], device=noise.device, dtype=torch.long)
)
noisy_input = noisy_input.transpose(0, 1)
else:
# for getting real output
flow_pred, self.teacache_state = model_pred(
noisy_input.to(dtype),
text_embeds["prompt_embeds"],
text_embeds["negative_prompt_embeds"],
timestep, step_index, image_cond, clip_fea,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_kv_cache_start=current_kv_cache_start_frame * self.frame_seq_length,
kv_start = self.kv_start,
kv_end = kv_end
)
denoised_pred = convert_flow_pred_to_x0(
flow_pred=flow_pred.transpose(0, 1),
xt=noisy_input.transpose(0, 1),
timestep=timestep.flatten(0, 1)
)
denoised_pred = denoised_pred.transpose(0, 1)
# Step 3.2: record the model's output
#print("denoised_pred shape before output: ", denoised_pred.shape)
#output_latents[:, current_kv_cache_start_frame:current_kv_cache_start_frame + current_num_frames] = denoised_pred
output_latents[:, i * self.num_frame_per_block:(i + 1) * self.num_frame_per_block] = denoised_pred
# Step 3.3: rerun with timestep zero to update KV cache using clean context
print("cleaning KV cache")
context_timestep = torch.ones_like(timestep) * context_noise
model_pred(
denoised_pred.to(dtype),
text_embeds["prompt_embeds"],
text_embeds["negative_prompt_embeds"],
context_timestep, step_index, image_cond, clip_fea,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_kv_cache_start=current_kv_cache_start_frame * self.frame_seq_length,
kv_start = self.kv_start,
kv_end = kv_end
)
# Update positions for next block
r_shift_length = self.num_frame_per_block * self.frame_seq_length
self.kv_start += r_shift_length
kv_end += r_shift_length
#self.rope_start += r_shift_length
# Check if we need to shift the cache
if kv_end > self.cache_window_size:
print("Shifting KV cache")
kv_end -= self._shift_kv_cache()
# Step 3.4: update the start and end frame indices
current_kv_cache_start_frame += current_num_frames
if callback is not None:
#callback_latent = output_latents[:, :current_kv_cache_start_frame].float().detach().permute(1,0,2,3)
callback_latent = denoised_pred.float().detach().permute(1,0,2,3)
callback(i, callback_latent, None, num_blocks)
else:
pbar.update(1)
# reset cross attn cache
for block_index in range(transformer.num_layers):
self.crossattn_cache[block_index]["is_init"] = False
# reset kv cache
for block_index in range(len(self.kv_cache1)):
self.kv_cache1[block_index]["global_end_index"] = torch.tensor(
[0], dtype=torch.long, device=noise.device)
self.kv_cache1[block_index]["local_end_index"] = torch.tensor(
[0], dtype=torch.long, device=noise.device)
self.kv_cache1 = None
self.crossattn_cache = None
if force_offload:
if model["manual_offloading"]:
transformer.to(offload_device)
mm.soft_empty_cache()
gc.collect()
try:
print_memory(device)
torch.cuda.reset_peak_memory_stats(device)
except:
pass
return ({
"samples": output_latents.unsqueeze(0).cpu(),
}, )
NODE_CLASS_MAPPINGS = {
"WanVideoCausVidSampler": WanVideoCausVidSampler,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoCausVidSampler": "WanVideo CausVid Sampler",
}
+170
View File
@@ -0,0 +1,170 @@
import torch
from ..utils import log
import comfy.model_management as mm
from comfy.utils import load_torch_file
from tqdm import tqdm
import gc
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
import folder_paths
class WanVideoControlnetLoader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": (folder_paths.get_filename_list("controlnet"), {"tooltip": "These models are loaded from the 'ComfyUI/models/controlnet' -folder",}),
"base_precision": (["fp32", "bf16", "fp16"], {"default": "bf16"}),
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2', 'fp8_e4m3fn_fast_no_ffn'], {"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"}),
},
}
RETURN_TYPES = ("WANVIDEOCONTROLNET",)
RETURN_NAMES = ("controlnet", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Loads ControlNet model from 'https://huggingface.co/collections/TheDenk/wan21-controlnets-68302b430411dafc0d74d2fc'"
def loadmodel(self, model, base_precision, load_device, quantization):
device = mm.get_torch_device()
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)
num_layers = 8 if "blocks.7.scale_shift_table" in sd else 6
out_proj_dim = 5120 if num_layers == 6 else 1536
if not "control_encoder.0.0.weight" in sd:
raise ValueError("Invalid ControlNet model")
controlnet_cfg = {
"added_kv_proj_dim": None,
"attention_head_dim": 128,
"cross_attn_norm": None,
"downscale_coef": 8,
"eps": 1e-06,
"ffn_dim": 8960,
"freq_dim": 256,
"image_dim": None,
"in_channels": 3,
"num_attention_heads": 12,
"num_layers": num_layers,
"out_proj_dim": out_proj_dim,
"patch_size": [
1,
2,
2
],
"qk_norm": "rms_norm_across_heads",
"rope_max_seq_len": 1024,
"text_dim": 4096,
"vae_channels": 16
}
from .wan_controlnet import WanControlnet
with init_empty_weights():
controlnet = WanControlnet(**controlnet_cfg)
controlnet.eval()
if quantization == "disabled":
for k, v in sd.items():
if isinstance(v, torch.Tensor):
if v.dtype == torch.float8_e4m3fn:
quantization = "fp8_e4m3fn"
break
elif v.dtype == torch.float8_e5m2:
quantization = "fp8_e5m2"
break
if "fp8_e4m3fn" in quantization:
dtype = torch.float8_e4m3fn
elif quantization == "fp8_e5m2":
dtype = torch.float8_e5m2
else:
dtype = base_dtype
params_to_keep = {"norm", "head", "time_in", "vector_in", "controlnet_patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter"}
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}",
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 load_device == "offload_device" and controlnet.device != offload_device:
log.info(f"Moving controlnet model from {controlnet.device} to {offload_device}")
controlnet.to(offload_device)
gc.collect()
mm.soft_empty_cache()
return (controlnet,)
class WanVideoControlnetApply:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("WANVIDEOMODEL", ),
"controlnet": ("WANVIDEOCONTROLNET", ),
"control_images": ("IMAGE", ),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001, "tooltip": "controlnet strength"}),
"control_stride": ("INT", {"default": 3, "min": 1, "max": 8, "step": 1, "tooltip": "controlnet stride"}),
"control_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the steps to apply controlnet"}),
"control_end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the steps to apply controlnet"}),
}
}
RETURN_TYPES = ("WANVIDEOMODEL",)
RETURN_NAMES = ("model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, model, controlnet, control_images, strength, control_stride, control_start_percent, control_end_percent):
patcher = model.clone()
if 'transformer_options' not in patcher.model_options:
patcher.model_options['transformer_options'] = {}
control_input = control_images.permute(3, 0, 1, 2).unsqueeze(0).contiguous()
control_input = control_input * 2.0 - 1.0
controlnet = {
"controlnet": controlnet,
"control_latents": control_input,
"controlnet_strength": strength,
"control_stride": control_stride,
"controlnet_start": control_start_percent,
"controlnet_end": control_end_percent
}
patcher.model_options["transformer_options"]["controlnet"] = controlnet
return (patcher,)
NODE_CLASS_MAPPINGS = {
"WanVideoControlnetLoader": WanVideoControlnetLoader,
"WanVideoControlnet": WanVideoControlnetApply,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoControlnetLoader": "WanVideo Controlnet Loader",
"WanVideoControlnet": "WanVideo Controlnet Apply",
}
+267
View File
@@ -0,0 +1,267 @@
# source https://github.com/TheDenk/wan2.1-dilated-controlnet/blob/main/wan_controlnet.py
from typing import Any, Dict, Optional, Tuple, Union
import torch
import torch.nn as nn
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscale_lora_layers
from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.transformers.transformer_wan import (
WanTimeTextImageEmbedding,
WanRotaryPosEmbed,
WanTransformerBlock
)
def zero_module(module):
for p in module.parameters():
nn.init.zeros_(p)
return module
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin):
r"""
A Controlnet Transformer model for video-like data used in the Wan model.
Args:
patch_size (`Tuple[int]`, defaults to `(1, 2, 2)`):
3D patch dimensions for video embedding (t_patch, h_patch, w_patch).
num_attention_heads (`int`, defaults to `40`):
Fixed length for text embeddings.
attention_head_dim (`int`, defaults to `128`):
The number of channels in each head.
vae_channels (`int`, defaults to `16`):
The number of channels in the vae input.
in_channels (`int`, defaults to `16`):
The number of channels in the controlnet input.
text_dim (`int`, defaults to `512`):
Input dimension for text embeddings.
freq_dim (`int`, defaults to `256`):
Dimension for sinusoidal time embeddings.
ffn_dim (`int`, defaults to `13824`):
Intermediate dimension in feed-forward network.
num_layers (`int`, defaults to `40`):
The number of layers of transformer blocks to use.
window_size (`Tuple[int]`, defaults to `(-1, -1)`):
Window size for local attention (-1 indicates global attention).
cross_attn_norm (`bool`, defaults to `True`):
Enable cross-attention normalization.
qk_norm (`bool`, defaults to `True`):
Enable query/key normalization.
eps (`float`, defaults to `1e-6`):
Epsilon value for normalization layers.
add_img_emb (`bool`, defaults to `False`):
Whether to use img_emb.
added_kv_proj_dim (`int`, *optional*, defaults to `None`):
The number of channels to use for the added key and value projections. If `None`, no projection is used.
downscale_coef (`int`, *optional*, defaults to `8`):
Coeficient for downscale controlnet input video.
out_proj_dim (`int`, *optional*, defaults to `128 * 12`):
Output projection dimention for last linear layers.
"""
_supports_gradient_checkpointing = True
_skip_layerwise_casting_patterns = ["patch_embedding", "condition_embedder", "norm"]
_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,
patch_size: Tuple[int] = (1, 2, 2),
num_attention_heads: int = 40,
attention_head_dim: int = 128,
in_channels: int = 3,
vae_channels: int = 16,
text_dim: int = 4096,
freq_dim: int = 256,
ffn_dim: int = 13824,
num_layers: int = 20,
cross_attn_norm: bool = True,
qk_norm: Optional[str] = "rms_norm_across_heads",
eps: float = 1e-6,
image_dim: Optional[int] = None,
added_kv_proj_dim: Optional[int] = None,
rope_max_seq_len: int = 1024,
downscale_coef: int = 8,
out_proj_dim: int = 128 * 12,
) -> None:
super().__init__()
start_channels = in_channels * (downscale_coef ** 2)
input_channels = [start_channels, start_channels // 2, start_channels // 4]
self.control_encoder = nn.ModuleList([
## Spatial compression with time awareness
nn.Sequential(
nn.Conv3d(
in_channels,
input_channels[0],
kernel_size=(3, downscale_coef + 1, downscale_coef + 1),
stride=(1, downscale_coef, downscale_coef),
padding=(1, downscale_coef // 2, downscale_coef // 2)
),
nn.GELU(approximate="tanh"),
nn.GroupNorm(2, input_channels[0]),
),
## Temporal compression with spatial awareness
nn.Sequential(
nn.Conv3d(input_channels[0], input_channels[1], kernel_size=3, stride=(2, 1, 1), padding=1),
nn.GELU(approximate="tanh"),
nn.GroupNorm(2, input_channels[1]),
),
## Temporal compression with spatial awareness
nn.Sequential(
nn.Conv3d(input_channels[1], input_channels[2], kernel_size=3, stride=(2, 1, 1), padding=1),
nn.GELU(approximate="tanh"),
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)
# 2. Condition embeddings
# image_embedding_dim=1280 for I2V model
self.condition_embedder = WanTimeTextImageEmbedding(
dim=inner_dim,
time_freq_dim=freq_dim,
time_proj_dim=inner_dim * 6,
text_embed_dim=text_dim,
image_embed_dim=image_dim,
)
# 3. Transformer blocks
self.blocks = nn.ModuleList(
[
WanTransformerBlock(
inner_dim, ffn_dim, num_attention_heads, qk_norm, cross_attn_norm, eps, added_kv_proj_dim
)
for _ in range(num_layers)
]
)
# 4 Controlnet modules
self.controlnet_blocks = nn.ModuleList([])
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,
timestep: torch.LongTensor,
encoder_hidden_states: torch.Tensor,
controlnet_states: torch.Tensor,
encoder_hidden_states_image: Optional[torch.Tensor] = None,
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
if attention_kwargs is not None:
attention_kwargs = attention_kwargs.copy()
lora_scale = attention_kwargs.pop("scale", 1.0)
else:
lora_scale = 1.0
if USE_PEFT_BACKEND:
# weight the lora layers by setting `lora_scale` for each PEFT layer
scale_lora_layers(self, lora_scale)
else:
if attention_kwargs is not None and attention_kwargs.get("scale", None) is not None:
logger.warning(
"Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective."
)
rotary_emb = self.rope(hidden_states)
# 0. Controlnet encoder
for control_encoder_block in self.control_encoder:
controlnet_states = control_encoder_block(controlnet_states)
hidden_states = torch.cat([hidden_states, controlnet_states], dim=1)
## 1. Patch embedding and stack
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image
)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
# 2. Transformer blocks
controlnet_hidden_states = ()
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block, controlnet_block in zip(self.blocks, self.controlnet_blocks):
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states, timestep_proj, rotary_emb
)
controlnet_hidden_states += (controlnet_block(hidden_states),)
else:
for block, controlnet_block in zip(self.blocks, self.controlnet_blocks):
hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, rotary_emb)
controlnet_hidden_states += (controlnet_block(hidden_states),)
if USE_PEFT_BACKEND:
# remove `lora_scale` from each PEFT layer
unscale_lora_layers(self, lora_scale)
if not return_dict:
return (controlnet_hidden_states,)
return Transformer2DModelOutput(sample=controlnet_hidden_states)
if __name__ == "__main__":
parameters = {
"added_kv_proj_dim": None,
"attention_head_dim": 128,
"cross_attn_norm": True,
"eps": 1e-06,
"ffn_dim": 8960,
"freq_dim": 256,
"image_dim": None,
"in_channels": 3,
"num_attention_heads": 12,
"num_layers": 2,
"patch_size": [1, 2, 2],
"qk_norm": "rms_norm_across_heads",
"rope_max_seq_len": 1024,
"text_dim": 4096,
"downscale_coef": 8,
"out_proj_dim": 12 * 128,
}
controlnet = WanControlnet(**parameters)
hidden_states = torch.rand(1, 16, 21, 60, 90)
timestep = torch.randint(low=0, high=1000, size=(1,), dtype=torch.long)
encoder_hidden_states = torch.rand(1, 512, 4096)
controlnet_states = torch.rand(1, 3, 81, 480, 720)
controlnet_hidden_states = controlnet(
hidden_states=hidden_states,
timestep=timestep,
encoder_hidden_states=encoder_hidden_states,
controlnet_states=controlnet_states,
return_dict=False
)
print("Output states count", len(controlnet_hidden_states[0]))
for out_hidden_states in controlnet_hidden_states[0]:
print(out_hidden_states.shape)
+1 -1
View File
@@ -75,7 +75,7 @@ def enable_vram_management_recursively(model: torch.nn.Module, module_map: dict,
for name, module in model.named_children():
for source_module, target_module in module_map.items():
if isinstance(module, source_module):
if "rope_embedder" in name or "patch_embedding" in name:
if "rope_embedder" in name or "patch_embedding" in name or "emb_pos" in name:
continue
num_param = sum(p.numel() for p in module.parameters())
Binary file not shown.

After

Width:  |  Height:  |  Size: 192 KiB

Binary file not shown.
@@ -1,8 +1,8 @@
{
"id": "c6e410bc-5e2c-460b-ae81-c91b6094fbb1",
"revision": 0,
"last_node_id": 204,
"last_link_id": 336,
"last_node_id": 206,
"last_link_id": 341,
"nodes": [
{
"id": 42,
@@ -101,7 +101,7 @@
200
],
"flags": {},
"order": 21,
"order": 22,
"mode": 2,
"inputs": [
{
@@ -258,7 +258,7 @@
174
],
"flags": {},
"order": 34,
"order": 39,
"mode": 0,
"inputs": [
{
@@ -371,7 +371,7 @@
200
],
"flags": {},
"order": 22,
"order": 23,
"mode": 2,
"inputs": [
{
@@ -413,7 +413,7 @@
86
],
"flags": {},
"order": 24,
"order": 25,
"mode": 0,
"inputs": [
{
@@ -647,9 +647,7 @@
"ver": "0.3.27",
"Node name for S&R": "PreviewImage"
},
"widgets_values": [
""
]
"widgets_values": []
},
{
"id": 125,
@@ -663,15 +661,12 @@
190.28567504882812
],
"flags": {},
"order": 40,
"order": 41,
"mode": 0,
"inputs": [
{
"name": "text",
"type": "STRING",
"widget": {
"name": "text"
},
"link": 215
}
],
@@ -689,7 +684,7 @@
"Node name for S&R": "ShowText|pysssss"
},
"widgets_values": [
"",
"A man in a suit and tie walking down a hallway. He has a friendly expression and is looking directly at the camera. The hallway has beige walls adorned with framed black and white photographs. There is a door on the left side of the hallway and a poster on the wall. The lighting is soft and natural. The image is high quality and has a watermark in the bottom right corner.",
"A man in a suit and tie walking down a hallway. He has a friendly expression and is looking directly at the camera. The hallway has beige walls adorned with framed black and white photographs. There is a door on the left side of the hallway and a poster on the wall. The lighting is soft and natural. The image is high quality and has a watermark in the bottom right corner."
]
},
@@ -745,7 +740,7 @@
261.5306701660156
],
"flags": {},
"order": 41,
"order": 42,
"mode": 0,
"inputs": [
{
@@ -842,7 +837,7 @@
"flags": {
"collapsed": true
},
"order": 23,
"order": 24,
"mode": 0,
"inputs": [
{
@@ -914,7 +909,7 @@
266
],
"flags": {},
"order": 28,
"order": 29,
"mode": 0,
"inputs": [
{
@@ -922,28 +917,22 @@
"type": "IMAGE",
"link": 244
},
{
"name": "get_image_size",
"shape": 7,
"type": "IMAGE",
"link": null
},
{
"name": "width_input",
"shape": 7,
"type": "INT",
"widget": {
"name": "width_input"
},
"link": null
},
{
"name": "height_input",
"shape": 7,
"type": "INT",
"widget": {
"name": "height_input"
},
"link": null
},
{
"name": "get_image_size",
"shape": 7,
"type": "IMAGE",
"link": null
}
],
@@ -978,8 +967,6 @@
"lanczos",
false,
16,
0,
0,
"center"
]
},
@@ -1025,50 +1012,6 @@
"color": "#2a363b",
"bgcolor": "#3f5159"
},
{
"id": 74,
"type": "WidgetToString",
"pos": [
2128.5166015625,
-432.6599426269531
],
"size": [
315,
154
],
"flags": {},
"order": 29,
"mode": 0,
"inputs": [
{
"name": "any_input",
"shape": 7,
"type": "*",
"link": 102
}
],
"outputs": [
{
"name": "STRING",
"type": "STRING",
"links": [
212
]
}
],
"properties": {
"cnr_id": "comfyui-kjnodes",
"ver": "d57154c3a808b8a3f232ed293eaa2d000867c884",
"Node name for S&R": "WidgetToString"
},
"widgets_values": [
0,
"camera_type",
false,
"",
2
]
},
{
"id": 58,
"type": "WanVideoEncode",
@@ -1139,7 +1082,7 @@
"flags": {
"collapsed": true
},
"order": 42,
"order": 43,
"mode": 0,
"inputs": [
{
@@ -1177,7 +1120,7 @@
555.8994140625
],
"flags": {},
"order": 35,
"order": 40,
"mode": 0,
"inputs": [
{
@@ -1192,9 +1135,7 @@
"ver": "0.3.27",
"Node name for S&R": "PreviewImage"
},
"widgets_values": [
""
]
"widgets_values": []
},
{
"id": 128,
@@ -1292,7 +1233,7 @@
274
],
"flags": {},
"order": 39,
"order": 44,
"mode": 0,
"inputs": [
{
@@ -1304,9 +1245,6 @@
"name": "caption",
"shape": 7,
"type": "STRING",
"widget": {
"name": "caption"
},
"link": null
},
{
@@ -1341,8 +1279,7 @@
"black",
"FreeMonoBoldOblique.otf",
"input",
"up",
""
"up"
]
},
{
@@ -1353,7 +1290,7 @@
-1155.6121826171875
],
"size": [
390.5999755859375,
421.6000061035156,
202
],
"flags": {},
@@ -1383,82 +1320,6 @@
128
]
},
{
"id": 138,
"type": "ReCamMasterPoseVisualizer",
"pos": [
1597.2598876953125,
177.93458557128906
],
"size": [
349.6756591796875,
130
],
"flags": {},
"order": 31,
"mode": 0,
"inputs": [
{
"name": "camera_poses",
"type": "CAMERAPOSES",
"link": 242
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
243
]
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "2dc25c150ec4288e0b8689fc49fe2c9c8ab01999",
"Node name for S&R": "ReCamMasterPoseVisualizer"
},
"widgets_values": [
0.10000000000000002,
0.20000000000000004,
0.4000000000000001,
0.5000000000000001
]
},
{
"id": 157,
"type": "GetNode",
"pos": [
1394.278076171875,
-105.60762786865234
],
"size": [
210,
60
],
"flags": {
"collapsed": true
},
"order": 14,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
274
]
}
],
"title": "Get_InputLatents",
"properties": {},
"widgets_values": [
"InputLatents"
],
"color": "#323",
"bgcolor": "#535"
},
{
"id": 127,
"type": "WanVideoExperimentalArgs",
@@ -1471,7 +1332,7 @@
130
],
"flags": {},
"order": 15,
"order": 14,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1509,7 +1370,7 @@
"flags": {
"collapsed": true
},
"order": 16,
"order": 15,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1543,7 +1404,7 @@
"flags": {
"collapsed": true
},
"order": 17,
"order": 16,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1575,7 +1436,7 @@
46
],
"flags": {},
"order": 27,
"order": 28,
"mode": 2,
"inputs": [
{
@@ -1617,7 +1478,7 @@
"flags": {
"collapsed": true
},
"order": 44,
"order": 45,
"mode": 0,
"inputs": [
{
@@ -1657,7 +1518,7 @@
"flags": {
"collapsed": true
},
"order": 18,
"order": 17,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1689,7 +1550,7 @@
688.150634765625
],
"flags": {},
"order": 30,
"order": 34,
"mode": 0,
"inputs": [
{
@@ -1754,6 +1615,12 @@
"shape": 7,
"type": "EXPERIMENTALARGS",
"link": 334
},
{
"name": "sigmas",
"shape": 7,
"type": "SIGMAS",
"link": null
}
],
"outputs": [
@@ -1781,8 +1648,7 @@
0,
1,
false,
"comfy",
""
"comfy"
]
},
{
@@ -1797,7 +1663,7 @@
178
],
"flags": {},
"order": 19,
"order": 18,
"mode": 0,
"inputs": [],
"outputs": [
@@ -1832,10 +1698,10 @@
],
"size": [
908.9017944335938,
912.1107788085938
334
],
"flags": {},
"order": 43,
"order": 46,
"mode": 0,
"inputs": [
{
@@ -1900,53 +1766,6 @@
}
}
},
{
"id": 56,
"type": "WanVideoReCamMasterCameraEmbed",
"pos": [
1379.0372314453125,
-36.21258544921875
],
"size": [
356.0601806640625,
79.98188781738281
],
"flags": {},
"order": 25,
"mode": 0,
"inputs": [
{
"name": "latents",
"type": "LATENT",
"link": 274
}
],
"outputs": [
{
"name": "camera_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"links": [
102,
272
]
},
{
"name": "camera_poses",
"type": "CAMERAPOSES",
"links": [
242
]
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "11e9166d0b00fe3b1e6ebb0a3d1db50ce7a56d58",
"Node name for S&R": "WanVideoReCamMasterCameraEmbed"
},
"widgets_values": [
"arc_right"
]
},
{
"id": 22,
"type": "WanVideoModelLoader",
@@ -1959,7 +1778,7 @@
234
],
"flags": {},
"order": 20,
"order": 19,
"mode": 0,
"inputs": [
{
@@ -2017,6 +1836,248 @@
],
"color": "#223",
"bgcolor": "#335"
},
{
"id": 138,
"type": "ReCamMasterPoseVisualizer",
"pos": [
1758.8209228515625,
168.25579833984375
],
"size": [
349.6756591796875,
130
],
"flags": {},
"order": 35,
"mode": 0,
"inputs": [
{
"name": "camera_poses",
"type": "CAMERAPOSES",
"link": 242
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
243
]
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "2dc25c150ec4288e0b8689fc49fe2c9c8ab01999",
"Node name for S&R": "ReCamMasterPoseVisualizer"
},
"widgets_values": [
0.10000000000000002,
0.20000000000000004,
0.4000000000000001,
0.5000000000000001
]
},
{
"id": 157,
"type": "GetNode",
"pos": [
1339.9281005859375,
-96.67338562011719
],
"size": [
210,
60
],
"flags": {
"collapsed": true
},
"order": 20,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
274,
338
]
}
],
"title": "Get_InputLatents",
"properties": {},
"widgets_values": [
"InputLatents"
],
"color": "#323",
"bgcolor": "#535"
},
{
"id": 56,
"type": "WanVideoReCamMasterCameraEmbed",
"pos": [
1338.8333740234375,
-36.212589263916016
],
"size": [
356.0601806640625,
79.98188781738281
],
"flags": {},
"order": 30,
"mode": 0,
"inputs": [
{
"name": "camera_poses",
"type": "CAMERAPOSES",
"link": 340
},
{
"name": "latents",
"type": "LATENT",
"link": 274
}
],
"outputs": [
{
"name": "camera_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"links": [
272
]
},
{
"name": "camera_poses",
"type": "CAMERAPOSES",
"links": [
242
]
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "11e9166d0b00fe3b1e6ebb0a3d1db50ce7a56d58",
"Node name for S&R": "WanVideoReCamMasterCameraEmbed"
},
"widgets_values": []
},
{
"id": 206,
"type": "WanVideoReCamMasterGenerateOrbitCamera",
"pos": [
1305.5355224609375,
185.5100555419922
],
"size": [
384.8144836425781,
82
],
"flags": {},
"order": 21,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "camera_poses",
"type": "CAMERAPOSES",
"links": []
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "8257cd1f8abaa6504248b946f31c5173c0228b3d",
"Node name for S&R": "WanVideoReCamMasterGenerateOrbitCamera"
},
"widgets_values": [
81,
90
]
},
{
"id": 74,
"type": "WidgetToString",
"pos": [
2128.5166015625,
-432.6599426269531
],
"size": [
315,
154
],
"flags": {},
"order": 31,
"mode": 0,
"inputs": [
{
"name": "any_input",
"shape": 7,
"type": "*",
"link": 341
}
],
"outputs": [
{
"name": "STRING",
"type": "STRING",
"links": [
212
]
}
],
"properties": {
"cnr_id": "comfyui-kjnodes",
"ver": "d57154c3a808b8a3f232ed293eaa2d000867c884",
"Node name for S&R": "WidgetToString"
},
"widgets_values": [
0,
"camera_type",
false,
"",
2
]
},
{
"id": 205,
"type": "WanVideoReCamMasterDefaultCamera",
"pos": [
1317.4481201171875,
-241.10047912597656
],
"size": [
388.8835754394531,
58
],
"flags": {},
"order": 27,
"mode": 0,
"inputs": [
{
"name": "latents",
"type": "LATENT",
"link": 338
}
],
"outputs": [
{
"name": "camera_poses",
"type": "CAMERAPOSES",
"links": [
340,
341
]
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "8257cd1f8abaa6504248b946f31c5173c0228b3d",
"Node name for S&R": "WanVideoReCamMasterDefaultCamera"
},
"widgets_values": [
"pan_right"
]
}
],
"links": [
@@ -2060,14 +2121,6 @@
1,
"CONDITIONING"
],
[
102,
56,
0,
74,
0,
"*"
],
[
210,
28,
@@ -2257,7 +2310,7 @@
157,
0,
56,
0,
1,
"LATENT"
],
[
@@ -2291,6 +2344,30 @@
155,
6,
"TEACACHEARGS"
],
[
338,
157,
0,
205,
0,
"LATENT"
],
[
340,
205,
0,
56,
0,
"CAMERAPOSES"
],
[
341,
205,
0,
74,
0,
"*"
]
],
"groups": [
@@ -2337,13 +2414,12 @@
"config": {},
"extra": {
"ds": {
"scale": 0.7400249944258357,
"scale": 0.611590904484162,
"offset": [
1311.6629502036258,
1366.2150672288358
1176.5764579377562,
1095.9393193240473
]
},
"linkExtensions": [],
"node_versions": {
"ComfyUI-WanVideoWrapper": "5a2383621a05825d0d0437781afcb8552d9590fd",
"comfy-core": "0.3.26",
@@ -2352,7 +2428,8 @@
"VHS_latentpreview": true,
"VHS_latentpreviewrate": 0,
"VHS_MetadataImage": true,
"VHS_KeepIntermediate": true
"VHS_KeepIntermediate": true,
"frontendVersion": "1.16.7"
},
"version": 0.4
}
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+130
View File
@@ -0,0 +1,130 @@
import os
import torch
import torch.nn as nn
import torch.nn.functional as F
from safetensors import safe_open
class AudioProjModel(nn.Module):
def __init__(self, audio_in_dim=1024, cross_attention_dim=1024):
super().__init__()
self.cross_attention_dim = cross_attention_dim
self.proj = torch.nn.Linear(audio_in_dim, cross_attention_dim, bias=False)
self.norm = torch.nn.LayerNorm(cross_attention_dim)
def forward(self, audio_embeds):
context_tokens = self.proj(audio_embeds)
context_tokens = self.norm(context_tokens)
return context_tokens # [B,L,C]
class FantasyTalkingAudioConditionModel(nn.Module):
def __init__(self, audio_in_dim: int, audio_proj_dim: int):
super().__init__()
self.audio_in_dim = audio_in_dim
self.audio_proj_dim = audio_proj_dim
# audio proj model
self.proj_model = self.init_proj(self.audio_proj_dim)
def init_proj(self, cross_attention_dim=5120):
proj_model = AudioProjModel(
audio_in_dim=self.audio_in_dim, cross_attention_dim=cross_attention_dim
)
return proj_model
def get_proj_fea(self, audio_fea=None):
return self.proj_model(audio_fea) if audio_fea is not None else None
def split_audio_sequence(self, audio_proj_length, num_frames=81):
"""
Map the audio feature sequence to corresponding latent frame slices.
Args:
audio_proj_length (int): The total length of the audio feature sequence
(e.g., 173 in audio_proj[1, 173, 768]).
num_frames (int): The number of video frames in the training data (default: 81).
Returns:
list: A list of [start_idx, end_idx] pairs. Each pair represents the index range
(within the audio feature sequence) corresponding to a latent frame.
"""
# Average number of tokens per original video frame
tokens_per_frame = audio_proj_length / num_frames
# Each latent frame covers 4 video frames, and we want the center
tokens_per_latent_frame = tokens_per_frame * 4
half_tokens = int(tokens_per_latent_frame / 2)
pos_indices = []
for i in range(int((num_frames - 1) / 4) + 1):
if i == 0:
pos_indices.append(0)
else:
start_token = tokens_per_frame * ((i - 1) * 4 + 1)
end_token = tokens_per_frame * (i * 4 + 1)
center_token = int((start_token + end_token) / 2) - 1
pos_indices.append(center_token)
# Build index ranges centered around each position
pos_idx_ranges = [[idx - half_tokens, idx + half_tokens] for idx in pos_indices]
# Adjust the first range to avoid negative start index
pos_idx_ranges[0] = [
-(half_tokens * 2 - pos_idx_ranges[1][0]),
pos_idx_ranges[1][0],
]
return pos_idx_ranges
def split_tensor_with_padding(self, input_tensor, pos_idx_ranges, expand_length=0):
"""
Split the input tensor into subsequences based on index ranges, and apply right-side zero-padding
if the range exceeds the input boundaries.
Args:
input_tensor (Tensor): Input audio tensor of shape [1, L, 768].
pos_idx_ranges (list): A list of index ranges, e.g. [[-7, 1], [1, 9], ..., [165, 173]].
expand_length (int): Number of tokens to expand on both sides of each subsequence.
Returns:
sub_sequences (Tensor): A tensor of shape [1, F, L, 768], where L is the length after padding.
Each element is a padded subsequence.
k_lens (Tensor): A tensor of shape [F], representing the actual (unpadded) length of each subsequence.
Useful for ignoring padding tokens in attention masks.
"""
pos_idx_ranges = [
[idx[0] - expand_length, idx[1] + expand_length] for idx in pos_idx_ranges
]
sub_sequences = []
seq_len = input_tensor.size(1) # 173
max_valid_idx = seq_len - 1 # 172
k_lens_list = []
for start, end in pos_idx_ranges:
# Calculate the fill amount
pad_front = max(-start, 0)
pad_back = max(end - max_valid_idx, 0)
# Calculate the start and end indices of the valid part
valid_start = max(start, 0)
valid_end = min(end, max_valid_idx)
# Extract the valid part
if valid_start <= valid_end:
valid_part = input_tensor[:, valid_start : valid_end + 1, :]
else:
valid_part = input_tensor.new_zeros((1, 0, input_tensor.size(2)))
# In the sequence dimension (the 1st dimension) perform padding
padded_subseq = F.pad(
valid_part,
(0, 0, 0, pad_back + pad_front, 0, 0),
mode="constant",
value=0,
)
k_lens_list.append(padded_subseq.size(-2) - pad_back - pad_front)
sub_sequences.append(padded_subseq)
return torch.stack(sub_sequences, dim=1), torch.tensor(
k_lens_list, dtype=torch.long
)
+192
View File
@@ -0,0 +1,192 @@
import os
import torch
import gc
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
import folder_paths
script_directory = os.path.dirname(os.path.abspath(__file__))
class DownloadAndLoadWav2VecModel:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": (["facebook/wav2vec2-base-960h"],),
"base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}),
"load_device": (["main_device", "offload_device"], {"default": "main_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
},
}
RETURN_TYPES = ("WAV2VECMODEL",)
RETURN_NAMES = ("wav2vec_model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, model, base_precision, load_device):
from transformers import Wav2Vec2Model, Wav2Vec2Processor
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]
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
if load_device == "offload_device":
transfomer_load_device = offload_device
else:
transfomer_load_device = device
model_path = os.path.join(folder_paths.models_dir, "transformers", model)
if not os.path.exists(model_path):
log.info(f"Downloading Qwen model to: {model_path}")
from huggingface_hub import snapshot_download
snapshot_download(
repo_id=model,
ignore_patterns=["*.bin", "*.h5"],
local_dir=model_path,
local_dir_use_symlinks=False,
)
wav2vec_processor = Wav2Vec2Processor.from_pretrained(model_path)
wav2vec = Wav2Vec2Model.from_pretrained(model_path).to(base_dtype).to(transfomer_load_device).eval()
wav2vec_processor_model = {
"processor": wav2vec_processor,
"model": wav2vec,
"dtype": base_dtype,}
return (wav2vec_processor_model,)
class FantasyTalkingModelLoader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}),
"base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}),
},
}
RETURN_TYPES = ("FANTASYTALKINGMODEL",)
RETURN_NAMES = ("model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, model, base_precision):
from .model import FantasyTalkingAudioConditionModel
device = mm.get_torch_device()
offload_device = mm.unet_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("diffusion_models", model)
sd = load_torch_file(model_path, device=offload_device, safe_load=True)
with init_empty_weights():
fantasytalking_proj_model = FantasyTalkingAudioConditionModel(audio_in_dim=768, audio_proj_dim=2048)
#fantasytalking_proj_model.load_state_dict(sd, strict=False)
for name, param in fantasytalking_proj_model.named_parameters():
set_module_tensor_to_device(fantasytalking_proj_model, name, device=offload_device, dtype=base_dtype, value=sd[name])
fantasytalking = {
"proj_model": fantasytalking_proj_model,
"sd": sd,
}
return (fantasytalking,)
class FantasyTalkingWav2VecEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"wav2vec_model": ("WAV2VECMODEL",),
"fantasytalking_model": ("FANTASYTALKINGMODEL",),
"audio": ("AUDIO",),
"num_frames": ("INT", {"default": 81, "min": 1, "max": 1000, "step": 1}),
"fps": ("FLOAT", {"default": 23.0, "min": 1.0, "max": 60.0, "step": 0.1}),
"audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1, "tooltip": "Strength of the audio conditioning"}),
"audio_cfg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1, "tooltip": "When not 1.0, an extra model pass without audio conditioning is done: slower inference but more motion is allowed"}),
},
}
RETURN_TYPES = ("FANTASYTALKING_EMBEDS", )
RETURN_NAMES = ("fantasytalking_embeds",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, wav2vec_model, fantasytalking_model, fps, num_frames, audio_scale, audio_cfg_scale, audio):
import torchaudio
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
dtype = wav2vec_model["dtype"]
wav2vec = wav2vec_model["model"]
wav2vec_processor = wav2vec_model["processor"]
audio_proj_model = fantasytalking_model["proj_model"]
sr = 16000
audio_input = audio["waveform"]
sample_rate = audio["sample_rate"]
if sample_rate != sr:
audio_input = torchaudio.functional.resample(audio_input, sample_rate, sr)
audio_input = audio_input[0][0]
start_time = 0
end_time = num_frames / fps
start_sample = int(start_time * sr)
end_sample = int(end_time * sr)
try:
audio_segment = audio_input[start_sample:end_sample]
except:
audio_segment = audio_input
print("audio_segment.shape", audio_segment.shape)
input_values = wav2vec_processor(
audio_segment.numpy(), sampling_rate=sr, return_tensors="pt"
).input_values.to(dtype).to(device)
audio_features = wav2vec(input_values).last_hidden_state
audio_proj_model.proj_model.to(device)
audio_proj_fea = audio_proj_model.get_proj_fea(audio_features)
pos_idx_ranges = audio_proj_model.split_audio_sequence(
audio_proj_fea.size(1), num_frames=num_frames
)
audio_proj_split, audio_context_lens = audio_proj_model.split_tensor_with_padding(
audio_proj_fea, pos_idx_ranges, expand_length=4
) # [b,21,9+8,768]
audio_proj_model.proj_model.to(offload_device)
mm.soft_empty_cache()
out = {
"audio_proj": audio_proj_split,
"audio_context_lens": audio_context_lens,
"audio_scale": audio_scale,
"audio_cfg_scale": audio_cfg_scale
}
return (out,)
NODE_CLASS_MAPPINGS = {
"DownloadAndLoadWav2VecModel": DownloadAndLoadWav2VecModel,
"FantasyTalkingModelLoader": FantasyTalkingModelLoader,
"FantasyTalkingWav2VecEmbeds": FantasyTalkingWav2VecEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadWav2VecModel": "(Down)load Wav2Vec Model",
"FantasyTalkingModelLoader": "FantasyTalking Model Loader",
"FantasyTalkingWav2VecEmbeds": "FantasyTalking Wav2Vec Embeds",
}
+2 -2
View File
@@ -7,8 +7,8 @@ def fp8_linear_forward(cls, original_dtype, input):
weight_dtype = cls.weight.dtype
if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
if len(input.shape) == 3:
target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn
inn = input.reshape(-1, input.shape[2]).to(target_dtype)
#target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn
inn = input.reshape(-1, input.shape[2]).to(weight_dtype)
w = cls.weight.t()
scale = torch.ones((1), device=input.device, dtype=torch.float32)
+186
View File
@@ -0,0 +1,186 @@
import numpy as np
import os
import torch
from einops import rearrange
script_directory = os.path.dirname(os.path.abspath(__file__))
class Camera(object):
"""Copied from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
"""
def __init__(self, entry):
fx, fy, cx, cy = entry[1:5]
self.fx = fx
self.fy = fy
self.cx = cx
self.cy = cy
w2c_mat = np.array(entry[7:]).reshape(3, 4)
w2c_mat_4x4 = np.eye(4)
w2c_mat_4x4[:3, :] = w2c_mat
self.w2c_mat = w2c_mat_4x4
self.c2w_mat = np.linalg.inv(w2c_mat_4x4)
def custom_meshgrid(*args):
"""Copied from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
"""
# ref: https://pytorch.org/docs/stable/generated/torch.meshgrid.html?highlight=meshgrid#torch.meshgrid
return torch.meshgrid(*args)
def get_relative_pose(cam_params):
"""Copied from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
"""
abs_w2cs = [cam_param.w2c_mat for cam_param in cam_params]
abs_c2ws = [cam_param.c2w_mat for cam_param in cam_params]
cam_to_origin = 0
target_cam_c2w = np.array([
[1, 0, 0, 0],
[0, 1, 0, -cam_to_origin],
[0, 0, 1, 0],
[0, 0, 0, 1]
])
abs2rel = target_cam_c2w @ abs_w2cs[0]
ret_poses = [target_cam_c2w, ] + [abs2rel @ abs_c2w for abs_c2w in abs_c2ws[1:]]
ret_poses = np.array(ret_poses, dtype=np.float32)
return ret_poses
def ray_condition(K, c2w, H, W, device):
"""Copied from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
"""
# c2w: B, V, 4, 4
# K: B, V, 4
B = K.shape[0]
j, i = custom_meshgrid(
torch.linspace(0, H - 1, H, device=device, dtype=c2w.dtype),
torch.linspace(0, W - 1, W, device=device, dtype=c2w.dtype),
)
i = i.reshape([1, 1, H * W]).expand([B, 1, H * W]) + 0.5 # [B, HxW]
j = j.reshape([1, 1, H * W]).expand([B, 1, H * W]) + 0.5 # [B, HxW]
fx, fy, cx, cy = K.chunk(4, dim=-1) # B,V, 1
zs = torch.ones_like(i) # [B, HxW]
xs = (i - cx) / fx * zs
ys = (j - cy) / fy * zs
zs = zs.expand_as(ys)
directions = torch.stack((xs, ys, zs), dim=-1) # B, V, HW, 3
directions = directions / directions.norm(dim=-1, keepdim=True) # B, V, HW, 3
rays_d = directions @ c2w[..., :3, :3].transpose(-1, -2) # B, V, 3, HW
rays_o = c2w[..., :3, 3] # B, V, 3
rays_o = rays_o[:, :, None].expand_as(rays_d) # B, V, 3, HW
# c2w @ dirctions
rays_dxo = torch.cross(rays_o, rays_d)
plucker = torch.cat([rays_dxo, rays_d], dim=-1)
plucker = plucker.reshape(B, c2w.shape[1], H, W, 6) # B, V, H, W, 6
# plucker = plucker.permute(0, 1, 4, 2, 3)
return plucker
def process_poses(poses, width=672, height=384, original_pose_width=1280, original_pose_height=720, device='cpu', return_poses=False):
"""Modified from https://github.com/hehao13/CameraCtrl/blob/main/inference.py
"""
cam_params = [[float(x) for x in pose] for pose in poses]
if return_poses:
return cam_params
else:
cam_params = [Camera(cam_param) for cam_param in cam_params]
sample_wh_ratio = width / height
pose_wh_ratio = original_pose_width / original_pose_height # Assuming placeholder ratios, change as needed
if pose_wh_ratio > sample_wh_ratio:
resized_ori_w = height * pose_wh_ratio
for cam_param in cam_params:
cam_param.fx = resized_ori_w * cam_param.fx / width
else:
resized_ori_h = width / pose_wh_ratio
for cam_param in cam_params:
cam_param.fy = resized_ori_h * cam_param.fy / height
intrinsic = np.asarray([[cam_param.fx * width,
cam_param.fy * height,
cam_param.cx * width,
cam_param.cy * height]
for cam_param in cam_params], dtype=np.float32)
K = torch.as_tensor(intrinsic)[None] # [1, 1, 4]
c2ws = get_relative_pose(cam_params) # Assuming this function is defined elsewhere
c2ws = torch.as_tensor(c2ws)[None] # [1, n_frame, 4, 4]
plucker_embedding = ray_condition(K, c2ws, height, width, device=device)[0].permute(0, 3, 1, 2).contiguous() # V, 6, H, W
plucker_embedding = plucker_embedding[None]
plucker_embedding = rearrange(plucker_embedding, "b f c h w -> b f h w c")[0]
return plucker_embedding
class WanVideoFunCameraEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"poses": ("CAMERACTRL_POSES", ),
"width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the image to encode"}),
"height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the image to encode"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Strength of the camera motion"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the steps to apply camera motion"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the steps to apply camera motion"}),
},
# "optional": {
# "fun_ref_image": ("LATENT", {"tooltip": "Reference latent for the Fun 1.1 -model"}),
# }
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, poses, width, height, strength, start_percent, end_percent, fun_ref_image=None):
num_frames = len(poses)
control_camera_video = process_poses(poses, width, height)
control_camera_video = control_camera_video.permute([3, 0, 1, 2]).unsqueeze(0)
print("control_camera_video.shape", control_camera_video.shape)
# Rearrange dimensions
# Concatenate and transpose dimensions
control_camera_latents = torch.concat(
[
torch.repeat_interleave(control_camera_video[:, :, 0:1], repeats=4, dim=2),
control_camera_video[:, :, 1:]
], dim=2
).transpose(1, 2)
# Reshape, transpose, and view into desired shape
b, f, c, h, w = control_camera_latents.shape
control_camera_latents = control_camera_latents.contiguous().view(b, f // 4, 4, c, h, w).transpose(2, 3)
control_camera_latents = control_camera_latents.contiguous().view(b, f // 4, c * 4, h, w).transpose(1, 2)
print("control_camera_latents.shape", control_camera_latents.shape)
vae_stride = (4, 8, 8)
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,
"control_embeds": {
"control_camera_latents": control_camera_latents * strength,
"control_camera_start_percent": start_percent,
"control_camera_end_percent": end_percent,
"fun_ref_image": fun_ref_image["samples"][:,:, 0] if fun_ref_image is not None else None,
}
}
return (embeds,)
NODE_CLASS_MAPPINGS = {
"WanVideoFunCameraEmbeds": WanVideoFunCameraEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoFunCameraEmbeds": "WanVideo FunCamera Embeds",
}
+10 -8
View File
@@ -84,7 +84,7 @@ def get_previewer(device, latent_format):
taew_sd = comfy.utils.load_torch_file(taehv_path)
taesd = TAEHV(taew_sd).to(device)
previewer = TAESDPreviewerImpl(taesd)
previewer = WrappedPreviewer(previewer, rate=16)
previewer = WrappedPreviewer(previewer, rate=3)
if previewer is None:
if latent_format.latent_rgb_factors is not None:
@@ -164,17 +164,19 @@ class WrappedPreviewer(LatentPreviewer):
self.c_index = (self.c_index + num_previews) % num_images
return None
def process_previews(self, image_tensor, ind, leng):
max_size = 256
max_size = 512
image_tensor = self.decode_latent_to_preview(image_tensor)
if image_tensor.size(1) > max_size or image_tensor.size(2) > max_size:
image_tensor = image_tensor.movedim(-1,0)
if image_tensor.size(2) < image_tensor.size(3):
height = (max_size * image_tensor.size(2)) // image_tensor.size(3)
image_tensor = F.interpolate(image_tensor, (height,max_size), mode='bilinear')
image_tensor = F.interpolate(image_tensor, (height,max_size), mode='bicubic')
else:
width = (max_size * image_tensor.size(3)) // image_tensor.size(2)
image_tensor = F.interpolate(image_tensor, (max_size, width), mode='bilinear')
image_tensor = F.interpolate(image_tensor, (max_size, width), mode='bicubic')
image_tensor = image_tensor.movedim(0,-1)
#image_tensor = image_tensor.repeat_interleave(2, dim=0)
previews_ubyte = (image_tensor.clamp(0, 1)
.mul(0xFF) # to 0..255
).to(device="cpu", dtype=torch.uint8)
@@ -189,10 +191,10 @@ class WrappedPreviewer(LatentPreviewer):
#NOTE: send sync already uses call_soon_threadsafe
serv.send_sync(server.BinaryEventTypes.PREVIEW_IMAGE,
message.getvalue(), serv.client_id)
if self.rate == 16:
ind = (ind + 1) % ((leng-1) * 4 - 1)
else:
ind = (ind + 1) % leng
#if self.rate == 16:
ind = (ind + 1) % ((leng-1) * 4 - 1)
#else:
# ind = (ind + 1) % leng
# Send SwarmUI preview if detected
if self.swarmui_env:
+852 -345
View File
File diff suppressed because it is too large Load Diff
+2 -2
View File
@@ -1,7 +1,7 @@
[project]
name = "ComfyUI-WanVideoWrapper"
description = "ComfyUI diffusers wrapper nodes for WanVideo"
version = "1.1.5"
description = "ComfyUI wrapper nodes for WanVideo"
version = "1.1.9"
license = {file = "LICENSE"}
dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.32.0", "ftfy"]
+105 -18
View File
@@ -10,9 +10,16 @@ class Camera(object):
c2w_mat = np.array(c2w).reshape(4, 4)
self.c2w_mat = c2w_mat
self.w2c_mat = np.linalg.inv(c2w_mat)
class WanVideoReCamMasterCameraEmbed:
def parse_matrix(matrix_str):
rows = matrix_str.strip().split('] [')
matrix = []
for row in rows:
row = row.replace('[', '').replace(']', '')
matrix.append(list(map(float, row.split())))
return np.array(matrix)
class WanVideoReCamMasterDefaultCamera:
@classmethod
def INPUT_TYPES(s):
return {"required": {
@@ -32,18 +39,16 @@ class WanVideoReCamMasterCameraEmbed:
},
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "CAMERAPOSES",)
RETURN_NAMES = ("camera_embeds", "camera_poses",)
RETURN_TYPES = ("CAMERAPOSES",)
RETURN_NAMES = ("camera_poses",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "https://github.com/KwaiVGI/ReCamMaster"
def process(self, camera_type, latents):
# load camera
import json
from einops import rearrange
camera_data_path = os.path.join(script_directory, "camera_extrinsics.json")
camera_data_path = os.path.join(script_directory, "recam_extrinsics.json")
with open(camera_data_path, 'r') as file:
cam_data = json.load(file)
@@ -65,10 +70,96 @@ class WanVideoReCamMasterCameraEmbed:
}
cam_idx = list(range(num_frames))[::4]
traj = [self.parse_matrix(cam_data[f"frame{idx}"][f"cam{int(camera_type_map[camera_type]):02d}"]) for idx in cam_idx]
traj = [parse_matrix(cam_data[f"frame{idx}"][f"cam{int(camera_type_map[camera_type]):02d}"]) for idx in cam_idx]
traj = np.stack(traj).transpose(0, 2, 1)
return (traj,)
class WanVideoReCamMasterGenerateOrbitCamera:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"num_frames": ("INT", {"default": 81, "min": 1, "max": 1000, "step": 1, "tooltip": "Number of frames to generate"}),
"degrees": ("INT", {"default": 90, "min": -180, "max": 180, "step": 1, "tooltip": "Degrees to orbit"}),
},
}
RETURN_TYPES = ("CAMERAPOSES",)
RETURN_NAMES = ("camera_poses",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "https://github.com/KwaiVGI/ReCamMaster"
def process(self, degrees, num_frames):
def generate_orbit(num_frames=num_frames, degrees=degrees):
camera_data = []
center = np.array([3390, 1380, 240]) # Center point of orbit
for i in range(num_frames):
# Calculate angle from 0 to specified degrees
angle = i * degrees / (num_frames - 1)
angle_rad = np.radians(angle)
# Calculate position - circular path around center
x = center[0] - np.cos(angle_rad)
y = center[1] - np.sin(angle_rad)
z = center[2]
# Calculate direction from camera to center point
camera_pos = np.array([x, y, z])
dir_to_center = center - camera_pos
# Calculate the angle needed to face the center
look_angle = np.arctan2(dir_to_center[1], dir_to_center[0])
# Rotation matrix for facing the center (corrected)
cos_look = np.cos(look_angle)
sin_look = np.sin(look_angle)
# Create transformation matrix directly
transform = np.array([
[cos_look, -sin_look, 0, x],
[sin_look, cos_look, 0, y],
[0, 0, 1, z],
[0, 0, 0, 1]
])
camera_data.append(transform)
return camera_data
# Generate orbit data
camera_transforms = generate_orbit(num_frames=num_frames, degrees=degrees)
traj = camera_transforms[::4]
traj = np.stack(traj)
return (traj,)
class WanVideoReCamMasterCameraEmbed:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"camera_poses": ("CAMERAPOSES",),
"latents": ("LATENT", {"tooltip": "source video"}),
},
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "CAMERAPOSES",)
RETURN_NAMES = ("camera_embeds", "camera_poses",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "https://github.com/KwaiVGI/ReCamMaster"
def process(self, camera_poses, latents):
from einops import rearrange
samples = latents["samples"].squeeze(0)
C, T, H, W = samples.shape
num_frames = (T - 1) * 4 + 1
c2ws = []
for c2w in traj:
for c2w in camera_poses:
c2w = c2w[:, [1, 2, 0, 3]]
c2w[:3, 1] *= -1.
c2w[:3, 3] /= 100
@@ -93,15 +184,7 @@ class WanVideoReCamMasterCameraEmbed:
}
}
return (embeds, traj,)
def parse_matrix(self, matrix_str):
rows = matrix_str.strip().split('] [')
matrix = []
for row in rows:
row = row.replace('[', '').replace(']', '')
matrix.append(list(map(float, row.split())))
return np.array(matrix)
return (embeds, camera_poses,)
def get_relative_pose(self, cam_params):
abs_w2cs = [cam_param.w2c_mat for cam_param in cam_params]
@@ -274,8 +357,12 @@ or a .txt file with RealEstate camera intrinsics and coordinates, in a 3D plot.
NODE_CLASS_MAPPINGS = {
"WanVideoReCamMasterCameraEmbed": WanVideoReCamMasterCameraEmbed,
"ReCamMasterPoseVisualizer": ReCamMasterPoseVisualizer,
"WanVideoReCamMasterGenerateOrbitCamera": WanVideoReCamMasterGenerateOrbitCamera,
"WanVideoReCamMasterDefaultCamera": WanVideoReCamMasterDefaultCamera,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoReCamMasterCameraEmbed": "WanVideo ReCamMaster Camera Embed",
"ReCamMasterPoseVisualizer": "ReCamMaster Pose Visualizer",
"WanVideoReCamMasterGenerateOrbitCamera": "WanVideo ReCamMaster Generate Orbit Camera",
"WanVideoReCamMasterDefaultCamera": "WanVideo ReCamMaster Default Camera",
}
+595
View File
@@ -0,0 +1,595 @@
import os
import torch
import gc
from ..utils import log, print_memory, fourier_filter
import math
from tqdm import tqdm
from ..wanvideo.modules.model import rope_params
from ..wanvideo.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from ..wanvideo.utils.scheduling_flow_match_lcm import FlowMatchLCMScheduler
from ..nodes import optimized_scale
from einops import rearrange
from ..enhance_a_video.globals import disable_enhance
import comfy.model_management as mm
from comfy.utils import load_torch_file, ProgressBar, common_upscale
from comfy.clip_vision import clip_preprocess, ClipVisionModel
from comfy.cli_args import args, LatentPreviewMethod
script_directory = os.path.dirname(os.path.abspath(__file__))
def generate_timestep_matrix(
num_frames,
step_template,
base_num_frames,
ar_step=5,
num_pre_ready=0,
casual_block_size=1,
shrink_interval_with_mask=False,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, list[tuple]]:
step_matrix, step_index = [], []
update_mask, valid_interval = [], []
num_iterations = len(step_template) + 1
num_frames_block = num_frames // casual_block_size
base_num_frames_block = base_num_frames // casual_block_size
if base_num_frames_block < num_frames_block:
infer_step_num = len(step_template)
gen_block = base_num_frames_block
min_ar_step = infer_step_num / gen_block
assert ar_step >= min_ar_step, f"ar_step should be at least {math.ceil(min_ar_step)} in your setting"
# print(num_frames, step_template, base_num_frames, ar_step, num_pre_ready, casual_block_size, num_frames_block, base_num_frames_block)
step_template = torch.cat(
[
torch.tensor([999], dtype=torch.int64, device=step_template.device),
step_template.long(),
torch.tensor([0], dtype=torch.int64, device=step_template.device),
]
) # to handle the counter in row works starting from 1
pre_row = torch.zeros(num_frames_block, dtype=torch.long)
if num_pre_ready > 0:
pre_row[: num_pre_ready // casual_block_size] = num_iterations
while torch.all(pre_row >= (num_iterations - 1)) == False:
new_row = torch.zeros(num_frames_block, dtype=torch.long)
for i in range(num_frames_block):
if i == 0 or pre_row[i - 1] >= (
num_iterations - 1
): # the first frame or the last frame is completely denoised
new_row[i] = pre_row[i] + 1
else:
new_row[i] = new_row[i - 1] - ar_step
new_row = new_row.clamp(0, num_iterations)
update_mask.append(
(new_row != pre_row) & (new_row != num_iterations)
) # False: no need to update, True: need to update
step_index.append(new_row)
step_matrix.append(step_template[new_row])
pre_row = new_row
# for long video we split into several sequences, base_num_frames is set to the model max length (for training)
terminal_flag = base_num_frames_block
if shrink_interval_with_mask:
idx_sequence = torch.arange(num_frames_block, dtype=torch.int64)
update_mask = update_mask[0]
update_mask_idx = idx_sequence[update_mask]
last_update_idx = update_mask_idx[-1].item()
terminal_flag = last_update_idx + 1
# for i in range(0, len(update_mask)):
for curr_mask in update_mask:
if terminal_flag < num_frames_block and curr_mask[terminal_flag]:
terminal_flag += 1
valid_interval.append((max(terminal_flag - base_num_frames_block, 0), terminal_flag))
step_update_mask = torch.stack(update_mask, dim=0)
step_index = torch.stack(step_index, dim=0)
step_matrix = torch.stack(step_matrix, dim=0)
if casual_block_size > 1:
step_update_mask = step_update_mask.unsqueeze(-1).repeat(1, 1, casual_block_size).flatten(1).contiguous()
step_index = step_index.unsqueeze(-1).repeat(1, 1, casual_block_size).flatten(1).contiguous()
step_matrix = step_matrix.unsqueeze(-1).repeat(1, 1, casual_block_size).flatten(1).contiguous()
valid_interval = [(s * casual_block_size, e * casual_block_size) for s, e in valid_interval]
return step_matrix, step_index, step_update_mask, valid_interval
#region Sampler
class WanVideoDiffusionForcingSampler:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("WANVIDEOMODEL",),
"text_embeds": ("WANVIDEOTEXTEMBEDS", ),
"image_embeds": ("WANVIDIMAGE_EMBEDS", ),
"addnoise_condition": ("INT", {"default": 10, "min": 0, "max": 1000, "tooltip": "Improves consistency in long video generation"}),
"fps": ("FLOAT", {"default": 24.0, "min": 1.0, "max": 120.0, "step": 0.01}),
"steps": ("INT", {"default": 30, "min": 1}),
"cfg": ("FLOAT", {"default": 6.0, "min": 0.0, "max": 30.0, "step": 0.01}),
"shift": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 1000.0, "step": 0.01}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"force_offload": ("BOOLEAN", {"default": True, "tooltip": "Moves the model to the offload device after sampling"}),
"scheduler": (["unipc", "unipc/beta", "euler", "euler/beta", "lcm", "lcm/beta"],
{
"default": 'unipc'
}),
},
"optional": {
"samples": ("LATENT", {"tooltip": "init Latents to use for video2video process"} ),
"prefix_samples": ("LATENT", {"tooltip": "prefix latents"} ),
"denoise_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"teacache_args": ("TEACACHEARGS", ),
"slg_args": ("SLGARGS", ),
"rope_function": (["default", "comfy"], {"default": "comfy", "tooltip": "Comfy's RoPE implementation doesn't use complex numbers and can thus be compiled, that should be a lot faster when using torch.compile"}),
"experimental_args": ("EXPERIMENTALARGS", ),
"unianimate_poses": ("UNIANIMATE_POSE", ),
}
}
RETURN_TYPES = ("LATENT", )
RETURN_NAMES = ("samples",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, model, text_embeds, image_embeds, shift, fps, steps, addnoise_condition, cfg, seed, scheduler,
force_offload=True, samples=None, prefix_samples=None, denoise_strength=1.0, slg_args=None, rope_function="default", teacache_args=None,
experimental_args=None, unianimate_poses=None):
#assert not (context_options and teacache_args), "Context options cannot currently be used together with teacache."
patcher = model
model = model.model
transformer = model.diffusion_model
dtype = model["dtype"]
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
steps = int(steps/denoise_strength)
timesteps = None
if 'unipc' in scheduler:
sample_scheduler = FlowUniPCMultistepScheduler(shift=shift)
sample_scheduler.set_timesteps(steps, device=device, shift=shift, use_beta_sigmas=('beta' in scheduler))
elif 'euler' in scheduler:
sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta'))
sample_scheduler.set_timesteps(steps, device=device)
elif 'lcm' in scheduler:
sample_scheduler = FlowMatchLCMScheduler(shift=shift, use_beta_sigmas=(scheduler == 'lcm/beta'))
sample_scheduler.set_timesteps(steps, device=device)
init_timesteps = sample_scheduler.timesteps
if denoise_strength < 1.0:
steps = int(steps * denoise_strength)
timesteps = timesteps[-(steps + 1):]
seed_g = torch.Generator(device=torch.device("cpu"))
seed_g.manual_seed(seed)
clip_fea, clip_fea_neg = None, None
vace_data, vace_context, vace_scale = None, None, None
image_cond = image_embeds.get("image_embeds", None)
target_shape = image_embeds.get("target_shape", None)
if target_shape is None:
raise ValueError("Empty image embeds must be provided for T2V (Text to Video")
has_ref = image_embeds.get("has_ref", False)
vace_context = image_embeds.get("vace_context", None)
vace_scale = image_embeds.get("vace_scale", None)
if not isinstance(vace_scale, list):
vace_scale = [vace_scale] * (steps+1)
vace_start_percent = image_embeds.get("vace_start_percent", 0.0)
vace_end_percent = image_embeds.get("vace_end_percent", 1.0)
vace_seqlen = image_embeds.get("vace_seq_len", None)
vace_additional_embeds = image_embeds.get("additional_vace_inputs", [])
if vace_context is not None:
vace_data = [
{"context": vace_context,
"scale": vace_scale,
"start": vace_start_percent,
"end": vace_end_percent,
"seq_len": vace_seqlen
}
]
if len(vace_additional_embeds) > 0:
for i in range(len(vace_additional_embeds)):
if vace_additional_embeds[i].get("has_ref", False):
has_ref = True
vace_scale = vace_additional_embeds[i]["vace_scale"]
if not isinstance(vace_scale, list):
vace_scale = [vace_scale] * (steps+1)
vace_data.append({
"context": vace_additional_embeds[i]["vace_context"],
"scale": vace_scale,
"start": vace_additional_embeds[i]["vace_start_percent"],
"end": vace_additional_embeds[i]["vace_end_percent"],
"seq_len": vace_additional_embeds[i]["vace_seq_len"]
})
noise = torch.randn(
target_shape[0],
target_shape[1] + 1 if has_ref else target_shape[1],
target_shape[2],
target_shape[3],
dtype=torch.float32,
device=torch.device("cpu"),
generator=seed_g)
latent_video_length = noise.shape[1]
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1])
if samples is not None:
input_samples = samples["samples"].squeeze(0).to(noise)
if input_samples.shape[1] != noise.shape[1]:
input_samples = torch.cat([input_samples[:, :1].repeat(1, noise.shape[1] - input_samples.shape[1], 1, 1), input_samples], dim=1)
original_image = input_samples.to(device)
if denoise_strength < 1.0:
latent_timestep = timesteps[:1].to(noise)
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
mask = samples.get("mask", None)
if mask is not None:
if mask.shape[2] != noise.shape[1]:
mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - mask.shape[2], noise.shape[2], noise.shape[3]), mask], dim=2)
latents = noise.to(device)
fps_embeds = None
if hasattr(transformer, "fps_embedding"):
fps = round(fps, 2)
log.info(f"Model has fps embedding, using {fps} fps")
fps_embeds = [fps]
fps_embeds = [0 if i == 16 else 1 for i in fps_embeds]
prefix_video = prefix_samples["samples"].to(noise) if prefix_samples is not None else None
prefix_video_latent_length = prefix_video.shape[2] if prefix_video is not None else 0
if prefix_video is not None:
log.info(f"Prefix video of length: {prefix_video_latent_length}")
latents[:, :prefix_video_latent_length] = prefix_video[0]
#base_num_frames = (base_num_frames - 1) // 4 + 1 if base_num_frames is not None else latent_video_length
base_num_frames=latent_video_length
ar_step = 0
causal_block_size = 1
step_matrix, _, step_update_mask, valid_interval = generate_timestep_matrix(
latent_video_length, init_timesteps, base_num_frames, ar_step, prefix_video_latent_length, causal_block_size
)
sample_schedulers = []
for _ in range(latent_video_length):
if 'unipc' in scheduler:
sample_scheduler = FlowUniPCMultistepScheduler(shift=shift)
sample_scheduler.set_timesteps(steps, device=device, shift=shift, use_beta_sigmas=('beta' in scheduler))
elif 'euler' in scheduler:
sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift)
sample_scheduler.set_timesteps(steps, device=device)
elif 'lcm' in scheduler:
sample_scheduler = FlowMatchLCMScheduler(shift=shift, use_beta_sigmas=(scheduler == 'lcm/beta'))
sample_scheduler.set_timesteps(steps, device=device)
sample_schedulers.append(sample_scheduler)
sample_schedulers_counter = [0] * latent_video_length
unianim_data = None
if unianimate_poses is not None:
transformer.dwpose_embedding.to(device)
transformer.randomref_embedding_pose.to(device)
dwpose_data = unianimate_poses["pose"]
dwpose_data = transformer.dwpose_embedding(
(torch.cat([dwpose_data[:,:,:1].repeat(1,1,3,1,1), dwpose_data], dim=2)
).to(device)).to(model["dtype"])
log.info(f"UniAnimate pose embed shape: {dwpose_data.shape}")
if dwpose_data.shape[2] > latent_video_length:
log.warning(f"UniAnimate pose embed length {dwpose_data.shape[2]} is longer than the video length {latent_video_length}, truncating")
dwpose_data = dwpose_data[:,:, :latent_video_length]
elif dwpose_data.shape[2] < latent_video_length:
log.warning(f"UniAnimate pose embed length {dwpose_data.shape[2]} is shorter than the video length {latent_video_length}, padding with last pose")
pad_len = latent_video_length - dwpose_data.shape[2]
pad = dwpose_data[:,:,:1].repeat(1,1,pad_len,1,1)
dwpose_data = torch.cat([dwpose_data, pad], dim=2)
dwpose_data_flat = rearrange(dwpose_data, 'b c f h w -> b (f h w) c').contiguous()
random_ref_dwpose_data = None
if image_cond is not None:
random_ref_dwpose = unianimate_poses.get("ref", None)
if random_ref_dwpose is not None:
random_ref_dwpose_data = transformer.randomref_embedding_pose(
random_ref_dwpose.to(device)
).unsqueeze(2).to(model["dtype"]) # [1, 20, 104, 60]
unianim_data = {
"dwpose": dwpose_data_flat,
"random_ref": random_ref_dwpose_data.squeeze(0) if random_ref_dwpose_data is not None else None,
"strength": unianimate_poses["strength"],
"start_percent": unianimate_poses["start_percent"],
"end_percent": unianimate_poses["end_percent"]
}
disable_enhance() #not sure if this can work, disabling for now to avoid errors if it's enabled by another sampler
freqs = None
transformer.rope_embedder.k = None
transformer.rope_embedder.num_frames = None
if rope_function=="comfy":
transformer.rope_embedder.k = 0
transformer.rope_embedder.num_frames = latent_video_length
else:
d = transformer.dim // transformer.num_heads
freqs = torch.cat([
rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=0),
rope_params(1024, 2 * (d // 6)),
rope_params(1024, 2 * (d // 6))
],
dim=1)
if not isinstance(cfg, list):
cfg = [cfg] * (steps +1)
log.info(f"Seq len: {seq_len}")
pbar = ProgressBar(steps)
if args.preview_method in [LatentPreviewMethod.Auto, LatentPreviewMethod.Latent2RGB]: #default for latent2rgb
from latent_preview import prepare_callback
else:
from ..latent_preview import prepare_callback #custom for tiny VAE previews
callback = prepare_callback(patcher, steps)
#blockswap init
transformer_options = patcher.model_options.get("transformer_options", None)
if transformer_options is not None:
block_swap_args = transformer_options.get("block_swap_args", None)
if block_swap_args is not None:
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", True)
for name, param in transformer.named_parameters():
if "block" not in name:
param.data = param.data.to(device)
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking)
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking)
transformer.block_swap(
block_swap_args["blocks_to_swap"] - 1 ,
block_swap_args["offload_txt_emb"],
block_swap_args["offload_img_emb"],
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
)
elif model["auto_cpu_offload"]:
for module in transformer.modules():
if hasattr(module, "offload"):
module.offload()
if hasattr(module, "onload"):
module.onload()
elif model["manual_offloading"]:
transformer.to(device)
# Initialize TeaCache if enabled
if teacache_args is not None:
transformer.enable_teacache = True
transformer.rel_l1_thresh = teacache_args["rel_l1_thresh"]
transformer.teacache_start_step = teacache_args["start_step"]
transformer.teacache_cache_device = teacache_args["cache_device"]
log.info(f"TeaCache: Using cache device: {transformer.teacache_state.cache_device}")
transformer.teacache_end_step = len(init_timesteps)-1 if teacache_args["end_step"] == -1 else teacache_args["end_step"]
transformer.teacache_use_coefficients = teacache_args["use_coefficients"]
transformer.teacache_mode = teacache_args["mode"]
transformer.teacache_state.clear_all()
else:
transformer.enable_teacache = False
if slg_args is not None:
transformer.slg_blocks = slg_args["blocks"]
transformer.slg_start_percent = slg_args["start_percent"]
transformer.slg_end_percent = slg_args["end_percent"]
else:
transformer.slg_blocks = None
self.teacache_state = [None, None]
self.teacache_state_source = [None, None]
self.teacache_states_context = []
use_cfg_zero_star, use_fresca = False, False
if experimental_args is not None:
video_attention_split_steps = experimental_args.get("video_attention_split_steps", [])
if video_attention_split_steps:
transformer.video_attention_split_steps = [int(x.strip()) for x in video_attention_split_steps.split(",")]
else:
transformer.video_attention_split_steps = []
use_zero_init = experimental_args.get("use_zero_init", True)
use_cfg_zero_star = experimental_args.get("cfg_zero_star", False)
zero_star_steps = experimental_args.get("zero_star_steps", 0)
use_fresca = experimental_args.get("use_fresca", False)
if use_fresca:
fresca_scale_low = experimental_args.get("fresca_scale_low", 1.0)
fresca_scale_high = experimental_args.get("fresca_scale_high", 1.25)
fresca_freq_cutoff = experimental_args.get("fresca_freq_cutoff", 20)
#region model pred
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
vace_data=None, unianim_data=None, teacache_state=None):
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])):
if use_cfg_zero_star and (idx <= zero_star_steps) and use_zero_init:
return latent_model_input*0, None
nonlocal patcher
current_step_percentage = idx / len(init_timesteps)
control_lora_enabled = False
image_cond_input = image_cond
base_params = {
'seq_len': seq_len,
'device': device,
'freqs': freqs,
't': timestep,
'current_step': idx,
'control_lora_enabled': control_lora_enabled,
'vace_data': vace_data,
'unianim_data': unianim_data,
'fps_embeds': fps_embeds,
}
batch_size = 1
if not math.isclose(cfg_scale, 1.0) and len(positive_embeds) > 1:
negative_embeds = negative_embeds * len(positive_embeds)
#cond
noise_pred_cond, teacache_state_cond = transformer(
[z], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None,
clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage,
pred_id=teacache_state[0] if teacache_state else None,
**base_params
)
noise_pred_cond = noise_pred_cond[0].to(intermediate_device)
if math.isclose(cfg_scale, 1.0):
if use_fresca:
noise_pred_cond = fourier_filter(
noise_pred_cond,
scale_low=fresca_scale_low,
scale_high=fresca_scale_high,
freq_cutoff=fresca_freq_cutoff,
)
return noise_pred_cond, [teacache_state_cond]
#uncond
noise_pred_uncond, teacache_state_uncond = transformer(
[z], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea,
y=[image_cond_input] if image_cond_input is not None else None,
is_uncond=True, current_step_percentage=current_step_percentage,
pred_id=teacache_state[1] if teacache_state else None,
**base_params
)
noise_pred_uncond = noise_pred_uncond[0].to(intermediate_device)
#cfg
#https://github.com/WeichenFan/CFG-Zero-star/
if use_cfg_zero_star:
alpha = optimized_scale(
noise_pred_cond.view(batch_size, -1),
noise_pred_uncond.view(batch_size, -1)
).view(batch_size, 1, 1, 1)
else:
alpha = 1.0
#https://github.com/WikiChao/FreSca
if use_fresca:
filtered_cond = fourier_filter(
noise_pred_cond - noise_pred_uncond,
scale_low=fresca_scale_low,
scale_high=fresca_scale_high,
freq_cutoff=fresca_freq_cutoff,
)
noise_pred = noise_pred_uncond * alpha + cfg_scale * filtered_cond * alpha
else:
noise_pred = noise_pred_uncond * alpha + cfg_scale * (noise_pred_cond - noise_pred_uncond * alpha)
return noise_pred, [teacache_state_cond, teacache_state_uncond]
log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latents.shape[3]*8}x{latents.shape[2]*8} with {steps} steps")
intermediate_device = device
#clear memory before sampling
mm.unload_all_models()
mm.soft_empty_cache()
gc.collect()
try:
torch.cuda.reset_peak_memory_stats(device)
except:
pass
#region main loop start
for i, timestep_i in enumerate(tqdm(step_matrix)):
update_mask_i = step_update_mask[i]
valid_interval_i = valid_interval[i]
valid_interval_start, valid_interval_end = valid_interval_i
timestep = timestep_i[None, valid_interval_start:valid_interval_end].clone()
latent_model_input = latents[:, valid_interval_start:valid_interval_end, :, :].clone()
if addnoise_condition > 0 and valid_interval_start < prefix_video_latent_length:
noise_factor = 0.001 * addnoise_condition
timestep_for_noised_condition = addnoise_condition
latent_model_input[:, valid_interval_start:prefix_video_latent_length] = (
latent_model_input[:, valid_interval_start:prefix_video_latent_length] * (1.0 - noise_factor)
+ torch.randn_like(latent_model_input[:, valid_interval_start:prefix_video_latent_length])
* noise_factor
)
timestep[:, valid_interval_start:prefix_video_latent_length] = timestep_for_noised_condition
#print("timestep", timestep)
noise_pred, self.teacache_state = predict_with_cfg(
latent_model_input.to(dtype),
cfg[i],
text_embeds["prompt_embeds"],
text_embeds["negative_prompt_embeds"],
timestep, i, image_cond, clip_fea, unianim_data=unianim_data, vace_data=vace_data,
teacache_state=self.teacache_state)
for idx in range(valid_interval_start, valid_interval_end):
if update_mask_i[idx].item():
latents[:, idx] = sample_schedulers[idx].step(
noise_pred[:, idx - valid_interval_start],
timestep_i[idx],
latents[:, idx],
return_dict=False,
generator=seed_g,
)[0]
sample_schedulers_counter[idx] += 1
x0 = latents.unsqueeze(0)
if callback is not None:
callback_latent = (latent_model_input - noise_pred.to(timestep_i[idx].device) * timestep_i[idx] / 1000).detach().permute(1,0,2,3)
callback(i, callback_latent, None, steps)
else:
pbar.update(1)
if teacache_args is not None:
states = transformer.teacache_state.states
state_names = {
0: "conditional",
1: "unconditional"
}
for pred_id, state in states.items():
name = state_names.get(pred_id, f"prediction_{pred_id}")
if 'skipped_steps' in state:
log.info(f"TeaCache skipped: {len(state['skipped_steps'])} {name} steps: {state['skipped_steps']}")
transformer.teacache_state.clear_all()
if force_offload:
if model["manual_offloading"]:
transformer.to(offload_device)
mm.soft_empty_cache()
gc.collect()
try:
print_memory(device)
torch.cuda.reset_peak_memory_stats(device)
except:
pass
return ({
"samples": x0.cpu(),
}, )
NODE_CLASS_MAPPINGS = {
"WanVideoDiffusionForcingSampler": WanVideoDiffusionForcingSampler,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoDiffusionForcingSampler": "WanVideo Diffusion Forcing Sampler",
}
+89
View File
@@ -0,0 +1,89 @@
import einops
import torch
import torch.nn.functional as F
@torch.amp.autocast("cuda", enabled=False)
def batch_sample_rays(intrinsic, extrinsic, image_h=None, image_w=None):
''' get rays
Args:
intrinsic: [BF, 3, 3],
extrinsic: [BF, 4, 4],
h, w: int
# normalize: let the first camera R=I
Returns:
rays_o, rays_d: [BF, N, 3]
'''
# FIXME: PPU does not support inverse in GPU
device = intrinsic.device
B = intrinsic.shape[0]
c2w = torch.inverse(extrinsic)[:, :3, :4].to(device) # [BF,3,4]
x = torch.arange(image_w, device=device).float() - 0.5
y = torch.arange(image_h, device=device).float() + 0.5
points = torch.stack(torch.meshgrid(x, y, indexing='ij'), -1)
points = einops.repeat(points, 'w h c -> b (h w) c', b=B)
points = torch.cat([points, torch.ones_like(points)[:, :, 0:1]], dim=-1)
directions = points @ intrinsic.inverse().to(device).transpose(-1, -2) * 1 # depth is 1
rays_d = F.normalize(directions @ c2w[:, :3, :3].transpose(-1, -2), dim=-1) # [BF,N,3]
rays_o = c2w[..., :3, 3] # [BF, 3]
rays_o = rays_o[:, None, :].expand_as(rays_d) # [BF, N, 3]
return rays_o, rays_d
@torch.amp.autocast("cuda", enabled=False)
def embed_rays(rays_o, rays_d, nframe):
if len(rays_o.shape) == 4: # [b,f,n,3]
rays_o = einops.rearrange(rays_o, "b f n c -> (b f) n c")
rays_d = einops.rearrange(rays_d, "b f n c -> (b f) n c")
cross_od = torch.cross(rays_o, rays_d, dim=-1)
cam_emb = torch.cat([rays_d, cross_od], dim=-1)
cam_emb = einops.rearrange(cam_emb, "(b f) n c -> b f n c", f=nframe)
return cam_emb
@torch.amp.autocast("cuda", enabled=False)
def camera_center_normalization(w2c, nframe, camera_scale=2.0):
# copy from SEVA, w2c: [BF, 4, 4]
# ensure the first view is eye matrix
c2w_view0 = w2c[::nframe].inverse() # [B,4,4]
c2w_view0 = c2w_view0.repeat_interleave(nframe, dim=0) # [BF,4,4]
w2c = c2w_view0 @ w2c
# camera centering
c2w = torch.linalg.inv(w2c)
camera_dist_2med = torch.norm(c2w[:, :3, 3] - c2w[:, :3, 3].median(0, keepdim=True).values, dim=-1)
valid_mask = camera_dist_2med <= torch.clamp(torch.quantile(camera_dist_2med, 0.97) * 10, max=1e6)
c2w[:, :3, 3] -= c2w[valid_mask, :3, 3].mean(0, keepdim=True)
w2c = torch.linalg.inv(c2w)
# camera normalization
camera_dists = c2w[:, :3, 3].clone()
translation_scaling_factor = (
camera_scale
if torch.isclose(
torch.norm(camera_dists[0]),
torch.zeros(1, dtype=camera_dists.dtype, device=camera_dists.device),
atol=1e-5,
).any()
else (camera_scale / torch.norm(camera_dists[0]))
)
w2c[:, :3, 3] *= translation_scaling_factor
c2w[:, :3, 3] *= translation_scaling_factor
return w2c
def get_camera_embedding(intrinsic, extrinsic, f, h, w, normalize=True):
if normalize:
extrinsic = camera_center_normalization(extrinsic, nframe=f)
rays_o, rays_d = batch_sample_rays(intrinsic, extrinsic, image_h=h, image_w=w)
camera_embedding = embed_rays(rays_o, rays_d, nframe=f)
camera_embedding = einops.rearrange(camera_embedding, "b f (h w) c -> b c f h w", h=h, w=w)
return camera_embedding
+263
View File
@@ -0,0 +1,263 @@
import torch
import torch.nn as nn
from diffusers.models import ModelMixin
from typing import Optional
import torch.nn.functional as F
from diffusers.models.attention_processor import Attention
from diffusers.models.transformers.transformer_wan import WanRotaryPosEmbed
from einops import rearrange
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
def __call__(
self,
attn: Attention,
hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
rotary_emb: Optional[torch.Tensor] = None,
**kwargs
) -> torch.Tensor:
query = attn.to_q(hidden_states)
key = attn.to_k(hidden_states)
value = attn.to_v(hidden_states)
if attn.norm_q is not None:
query = attn.norm_q(query)
if attn.norm_k is not None:
key = attn.norm_k(key)
query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2)
key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2)
value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2) # [b,head,l,c]
if rotary_emb is not None:
def apply_rotary_emb(hidden_states: torch.Tensor, freqs: torch.Tensor):
x_rotated = torch.view_as_complex(hidden_states.to(torch.float64).unflatten(3, (-1, 2)))
x_out = torch.view_as_real(x_rotated * freqs).flatten(3, 4)
return x_out.type_as(hidden_states)
query = apply_rotary_emb(query, rotary_emb)
key = apply_rotary_emb(key, rotary_emb)
if self.attention_mode == 'sdpa':
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
elif self.attention_mode == 'sageattn':
hidden_states = sageattn_func(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
hidden_states = hidden_states.transpose(1, 2).flatten(2, 3)
hidden_states = hidden_states.type_as(query)
hidden_states = attn.to_out[0](hidden_states)
hidden_states = attn.to_out[1](hidden_states)
return hidden_states
class SimpleCogVideoXLayerNormZero(nn.Module):
def __init__(
self,
conditioning_dim: int,
embedding_dim: int,
elementwise_affine: bool = True,
eps: float = 1e-5,
bias: bool = True,
) -> None:
super().__init__()
self.silu = nn.SiLU()
self.linear = nn.Linear(conditioning_dim, 3 * embedding_dim, bias=bias)
self.norm = nn.LayerNorm(embedding_dim, eps=eps, elementwise_affine=elementwise_affine)
def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor):
shift, scale, gate = self.linear(self.silu(temb)).chunk(3, dim=1)
hidden_states = self.norm(hidden_states) * (1 + scale)[:, None, :] + shift[:, None, :]
return hidden_states, gate[:, None, :]
class SingleAttentionBlock(nn.Module):
def __init__(
self,
dim,
ffn_dim,
num_heads,
time_embed_dim=512,
qk_norm="rms_norm_across_heads",
eps=1e-6,
attention_mode="sdpa",
):
super().__init__()
self.dim = dim
self.ffn_dim = ffn_dim
self.num_heads = num_heads
self.qk_norm = qk_norm
self.eps = eps
# layers
self.norm1 = SimpleCogVideoXLayerNormZero(
time_embed_dim, dim, elementwise_affine=True, eps=1e-5, bias=True
)
self.self_attn = Attention(
query_dim=dim,
heads=num_heads,
kv_heads=num_heads,
dim_head=dim // num_heads,
qk_norm=qk_norm,
eps=eps,
bias=True,
cross_attention_dim=None,
out_bias=True,
processor=SimpleAttnProcessor2_0(attention_mode),
)
self.norm2 = SimpleCogVideoXLayerNormZero(
time_embed_dim, dim, elementwise_affine=True, eps=1e-5, bias=True
)
self.ffn = nn.Sequential(
nn.Linear(dim, ffn_dim),
nn.GELU(approximate='tanh'),
nn.Linear(ffn_dim, dim)
)
def forward(
self,
hidden_states,
temb,
rotary_emb,
):
# norm & modulate
norm_hidden_states, gate_msa = self.norm1(hidden_states, temb)
# attention
attn_hidden_states = self.self_attn(hidden_states=norm_hidden_states,
rotary_emb=rotary_emb)
hidden_states = hidden_states + gate_msa * attn_hidden_states
# norm & modulate
norm_hidden_states, gate_ff = self.norm2(hidden_states, temb)
# feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = hidden_states + gate_ff * ff_output
return hidden_states
class MaskCamEmbed(nn.Module):
def __init__(self, controlnet_cfg) -> None:
super().__init__()
# padding bug fixed
if controlnet_cfg.get("interp", False):
self.mask_padding = [0, 0, 0, 0, 3, 3] # 左右上下前后, I2V-interp,首尾帧
else:
self.mask_padding = [0, 0, 0, 0, 3, 0] # 左右上下前后, I2V
add_channels = controlnet_cfg.get("add_channels", 1)
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)))
def forward(self, add_inputs: torch.Tensor):
# render_mask.shape [b,c,f,h,w]
warp_add_pad = F.pad(add_inputs, self.mask_padding, mode="constant", value=0)
add_embeds = self.mask_proj(warp_add_pad) # [B,C,F,H,W]
add_embeds = self.mask_zero_proj(add_embeds)
add_embeds = rearrange(add_embeds, "b c f h w -> b (f h w) c")
return add_embeds
class WanControlNet(ModelMixin):
def __init__(self, controlnet_cfg):
super().__init__()
self.rope_max_seq_len = 1024
self.patch_size = (1, 2, 2)
self.in_channels = controlnet_cfg["in_channels"]
self.dim = controlnet_cfg["dim"]
self.num_heads = controlnet_cfg["num_heads"]
if controlnet_cfg["conv_out_dim"] != controlnet_cfg["dim"]:
self.proj_in = nn.Linear(controlnet_cfg["conv_out_dim"], controlnet_cfg["dim"])
else:
self.proj_in = nn.Identity()
self.controlnet_blocks = nn.ModuleList(
[
SingleAttentionBlock(
dim=self.dim,
ffn_dim=controlnet_cfg["ffn_dim"],
num_heads=self.num_heads,
time_embed_dim=controlnet_cfg["time_embed_dim"],
qk_norm="rms_norm_across_heads",
attention_mode=controlnet_cfg["attention_mode"],
)
for _ in range(controlnet_cfg["num_layers"])
]
)
self.proj_out = nn.ModuleList(
[
zero_module(nn.Linear(self.dim, 5120))
for _ in range(controlnet_cfg["num_layers"])
]
)
self.gradient_checkpointing = False
self.controlnet_rope = WanRotaryPosEmbed(self.dim // self.num_heads,
self.patch_size, self.rope_max_seq_len)
self.controlnet_patch_embedding = nn.Conv3d(
self.in_channels,
controlnet_cfg["conv_out_dim"],
kernel_size=self.patch_size,
stride=self.patch_size,
dtype=torch.float32
)
self.controlnet_mask_embedding = MaskCamEmbed(controlnet_cfg)
def forward(self, render_latent, render_mask, camera_embedding, temb, device):
controlnet_rotary_emb = self.controlnet_rope(render_latent)
controlnet_inputs = self.controlnet_patch_embedding(render_latent.to(torch.float32)).to(render_latent.dtype)
controlnet_inputs = controlnet_inputs.to(render_latent.dtype)
controlnet_inputs = controlnet_inputs.flatten(2).transpose(1, 2)
# additional inputs (mask, camera embedding)
add_inputs = None
if camera_embedding is not None and render_mask is not None:
add_inputs = torch.cat([render_mask, camera_embedding], dim=1)
elif render_mask is not None:
add_inputs = render_mask
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 = []
for i, block in enumerate(self.controlnet_blocks):
hidden_states = block(
hidden_states=hidden_states,
temb=temb,
rotary_emb=controlnet_rotary_emb
)
controlnet_states.append(self.proj_out[i](hidden_states).to(device))
return controlnet_states
+241
View File
@@ -0,0 +1,241 @@
import torch
from ..utils import log
import comfy.model_management as mm
from comfy.utils import ProgressBar, load_torch_file
from tqdm import tqdm
import gc
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
def INPUT_TYPES(s):
return {
"required": {
"model": (folder_paths.get_filename_list("controlnet"), {"tooltip": "These models are loaded from the 'ComfyUI/models/controlnet' -folder",}),
"base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}),
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2', 'fp8_e4m3fn_fast_no_ffn'], {"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"}),
"attention_mode": ([
"sdpa",
"sageattn",
], {"default": "sdpa"}),
},
"optional": {
"compile_args": ("WANCOMPILEARGS", ),
#"block_swap_args": ("BLOCKSWAPARGS", ),
}
}
RETURN_TYPES = ("WANVIDEOCONTROLNET",)
RETURN_NAMES = ("controlnet", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, model, base_precision, load_device, quantization, attention_mode, compile_args=None):
device = mm.get_torch_device()
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]
controlnet_cfg = {
"in_channels": in_channels,
"conv_out_dim": 5120,
"time_embed_dim": 5120,
"dim": 1024,
"ffn_dim": ffn_dim,
"num_heads": 16,
"num_layers": 20,
"add_channels": 7,
"mid_channels": 256,
"attention_mode": attention_mode
}
from .controlnet import WanControlNet
with init_empty_weights():
controlnet = WanControlNet(controlnet_cfg)
controlnet.eval()
if quantization == "disabled":
for k, v in sd.items():
if isinstance(v, torch.Tensor):
if v.dtype == torch.float8_e4m3fn:
quantization = "fp8_e4m3fn"
break
elif v.dtype == torch.float8_e5m2:
quantization = "fp8_e5m2"
break
if "fp8_e4m3fn" in quantization:
dtype = torch.float8_e4m3fn
elif quantization == "fp8_e5m2":
dtype = torch.float8_e5m2
else:
dtype = base_dtype
params_to_keep = {"norm", "head", "time_in", "vector_in", "controlnet_patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter"}
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}",
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:
torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"]
try:
if hasattr(torch, '_dynamo') and hasattr(torch._dynamo, 'config'):
torch._dynamo.config.recompile_limit = compile_args["dynamo_recompile_limit"]
except Exception as e:
log.warning(f"Could not set recompile_limit: {e}")
if compile_args["compile_transformer_blocks_only"]:
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"])
if load_device == "offload_device" and controlnet.device != offload_device:
log.info(f"Moving controlnet model from {controlnet.device} to {offload_device}")
controlnet.to(offload_device)
gc.collect()
mm.soft_empty_cache()
return (controlnet,)
class WanVideoUni3C_embeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"controlnet": ("WANVIDEOCONTROLNET",),
"render_latent": ("LATENT",),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the steps to apply the controlnet"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the steps to apply the controlnet"}),
},
"optional": {
"render_mask": ("MASK", {"tooltip": "NOT IMPLEMENTED!"}),
},
}
RETURN_TYPES = ("UNI3C_EMBEDS", )
RETURN_NAMES = ("uni3c_embeds",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, controlnet, render_latent, strength, start_percent, end_percent, render_mask=None):
device = mm.get_torch_device()
latent_mask = None
latents = render_latent["samples"]
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}")
# # load camera
# cam_info = json.load(open(f"{render_path}/cam_info.json"))
# w2cs = torch.tensor(np.array(cam_info["extrinsic"]), dtype=torch.float32, device=device)
# intrinsic = torch.tensor(np.array(cam_info["intrinsic"]), dtype=torch.float32, device=device)
# intrinsic[0, :] = intrinsic[0, :] / cam_info["width"] * width
# intrinsic[1, :] = intrinsic[1, :] / cam_info["height"] * height
# intrinsic = intrinsic[None].repeat(nframe, 1, 1)
# from .utils import build_cameras, set_initial_camera, traj_map
# focal_length = 1.0
# start_elevation = 5.0
# depth_avg = 0.5
# traj_type = "orbit"
# cam_traj, x_offset, y_offset, z_offset, d_theta, d_phi, d_r = traj_map(traj_type)
# focallength_px = focal_length * width
# K = torch.tensor([[focallength_px, 0, width / 2],
# [0, focallength_px, height / 2],
# [0, 0, 1]], dtype=torch.float32)
# 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,
# c2w_0=c2w_0,
# intrinsic=intrinsic,
# nframe=nframe,
# focal_length=focal_length,
# d_theta=d_theta,
# d_phi=d_phi,
# d_r=d_r,
# radius=depth_avg,
# x_offset=x_offset,
# 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)
uni3c_embeds = {
"controlnet": controlnet,
"controlnet_weight": strength,
"start": start_percent,
"end": end_percent,
"render_latent": latents.to(device),
"render_mask": latent_mask,
"camera_embedding": None
}
return (uni3c_embeds,)
NODE_CLASS_MAPPINGS = {
"WanVideoUni3C_ControlnetLoader": WanVideoUni3C_ControlnetLoader,
"WanVideoUni3C_embeds": WanVideoUni3C_embeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoUni3C_ControlnetLoader": "WanVideo Uni3C Controlnet Loader",
"WanVideoUni3C_embeds": "WanVideo Uni3C Embeds",
}
+206
View File
@@ -0,0 +1,206 @@
import imageio
import numpy as np
import torch
from PIL import Image
from scipy.interpolate import UnivariateSpline
from scipy.interpolate import interp1d
def load_video(video_path):
reader = imageio.get_reader(video_path)
total_frames = reader.count_frames()
frames = []
for i in range(total_frames):
frame = reader.get_data(i)
frames.append(Image.fromarray(frame))
reader.close()
return frames
def points_padding(points):
padding = torch.ones_like(points)[..., 0:1]
points = torch.cat([points, padding], dim=-1)
return points
def np_points_padding(points):
padding = np.ones_like(points)[..., 0:1]
points = np.concatenate([points, padding], axis=-1)
return points
def txt_interpolation(input_list, n, mode='smooth'):
x = np.linspace(0, 1, len(input_list))
if mode == 'smooth':
f = UnivariateSpline(x, input_list, k=3)
elif mode == 'linear':
f = interp1d(x, input_list)
else:
raise KeyError(f"Invalid txt interpolation mode: {mode}")
xnew = np.linspace(0, 1, n)
ynew = f(xnew)
return ynew
def traj_map(traj_type):
# pre-defined trajectories
if traj_type == "free1": # Zoom out and rotate to the upper left
cam_traj = "free"
x_offset = 0.0
y_offset = 0.0
z_offset = 0.0
d_theta = -15.0
d_phi = 45.0
d_r = 1.6
elif traj_type == "free2": # Rotate to the right horizontally
cam_traj = "free"
x_offset = -0.05
y_offset = 0.0
z_offset = 0.0
d_theta = 0.0
d_phi = -60.0
d_r = 1.0
elif traj_type == "free3": # Move back to the left
cam_traj = "free"
x_offset = -0.25
y_offset = 0.0
z_offset = 0.0
d_theta = 0.0
d_phi = 0.0
d_r = 1.7
elif traj_type == "free4": # Rotate and approach to the upper right
cam_traj = "free"
x_offset = 0.0
y_offset = 0.0
z_offset = 0.0
d_theta = -15.0
d_phi = -60.0
d_r = 0.75
elif traj_type == "free5": # Large-angle camera movement to the upper right
cam_traj = "free"
x_offset = 0.0
y_offset = 0.0
z_offset = 0.0
d_theta = -15.0
d_phi = -120.0
d_r = 1.6
elif traj_type == "swing1": # Swing shot 1
cam_traj = "swing1"
x_offset = 0.0
y_offset = 0.0
z_offset = 0.0
d_theta = 0.0
d_phi = 0.0
d_r = 1.0
elif traj_type == "swing2": # Swing shot 2
cam_traj = "swing2"
x_offset = 0.0
y_offset = 0.0
z_offset = 0.0
d_theta = 0.0
d_phi = 0.0
d_r = 1.0
elif traj_type == "orbit": # 360-degree counterclockwise rotation
cam_traj = "free"
x_offset = 0.0
y_offset = 0.0
z_offset = 0.0
d_theta = 0.0
d_phi = -360.0
d_r = 1.0
else:
raise NotImplementedError
return cam_traj, x_offset, y_offset, z_offset, d_theta, d_phi, d_r
def set_initial_camera(start_elevation, radius):
c2w_0 = torch.tensor([[1, 0, 0, 0],
[0, 1, 0, 0],
[0, 0, 1, -radius],
[0, 0, 0, 1]], dtype=torch.float32)
elevation_rad = np.deg2rad(start_elevation)
R_elevation = torch.tensor([[1, 0, 0, 0],
[0, np.cos(-elevation_rad), -np.sin(-elevation_rad), 0],
[0, np.sin(-elevation_rad), np.cos(-elevation_rad), 0],
[0, 0, 0, 1]], dtype=torch.float32)
c2w_0 = R_elevation @ c2w_0
w2c_0 = c2w_0.inverse()
return w2c_0, c2w_0
def build_cameras(cam_traj, w2c_0, c2w_0, intrinsic, nframe, focal_length,
d_theta, d_phi, d_r, radius, x_offset, y_offset, z_offset):
# build camera viewpoints according to d_theta,d_phi, d_r
# return: w2cs:[V,4,4], c2ws:[V,4,4], intrinsic:[V,3,3]
if intrinsic.ndim == 2:
intrinsic = intrinsic[None].repeat(nframe, 1, 1)
c2ws = [c2w_0]
w2cs = [w2c_0]
d_thetas, d_phis, d_rs = [], [], []
x_offsets, y_offsets, z_offsets = [], [], []
if cam_traj == "free":
for i in range(nframe - 1):
coef = (i + 1) / (nframe - 1)
d_thetas.append(d_theta * coef)
d_phis.append(d_phi * coef)
d_rs.append(coef * d_r + (1 - coef) * 1.0)
x_offsets.append(radius * x_offset * ((i + 1) / nframe))
y_offsets.append(radius * y_offset * ((i + 1) / nframe))
z_offsets.append(radius * z_offset * ((i + 1) / nframe))
elif cam_traj == "swing1":
phis__ = [0, -5, -25, -30, -20, -8, 0]
thetas__ = [0, -8, -12, -20, -17, -12, -5, -2, 1, 5, 3, 1, 0]
rs__ = [0, 0.2]
d_phis = txt_interpolation(phis__, nframe, mode='smooth')
d_phis[0] = phis__[0]
d_phis[-1] = phis__[-1]
d_thetas = txt_interpolation(thetas__, nframe, mode='smooth')
d_thetas[0] = thetas__[0]
d_thetas[-1] = thetas__[-1]
d_rs = txt_interpolation(rs__, nframe, mode='linear')
d_rs = 1.0 + d_rs
elif cam_traj == "swing2":
phis__ = [0, 5, 25, 30, 20, 10, 0]
thetas__ = [0, -5, -14, -11, 0, 1, 5, 3, 0]
rs__ = [0, -0.03, -0.1, -0.2, -0.17, -0.1, 0]
d_phis = txt_interpolation(phis__, nframe, mode='smooth')
d_phis[0] = phis__[0]
d_phis[-1] = phis__[-1]
d_thetas = txt_interpolation(thetas__, nframe, mode='smooth')
d_thetas[0] = thetas__[0]
d_thetas[-1] = thetas__[-1]
d_rs = txt_interpolation(rs__, nframe, mode='smooth')
d_rs = 1.0 + d_rs
else:
raise NotImplementedError("Unknown trajectory type...")
for i in range(nframe - 1):
d_theta_rad = np.deg2rad(d_thetas[i])
R_theta = torch.tensor([[1, 0, 0, 0],
[0, np.cos(d_theta_rad), -np.sin(d_theta_rad), 0],
[0, np.sin(d_theta_rad), np.cos(d_theta_rad), 0],
[0, 0, 0, 1]], dtype=torch.float32)
d_phi_rad = np.deg2rad(d_phis[i])
R_phi = torch.tensor([[np.cos(d_phi_rad), 0, np.sin(d_phi_rad), 0],
[0, 1, 0, 0],
[-np.sin(d_phi_rad), 0, np.cos(d_phi_rad), 0],
[0, 0, 0, 1]], dtype=torch.float32)
c2w_1 = R_phi @ R_theta @ c2w_0
if i < len(x_offsets) and i < len(y_offsets) and i < len(z_offsets):
c2w_1[:3, -1] += torch.tensor([x_offsets[i], y_offsets[i], z_offsets[i]])
c2w_1[:3, -1] *= d_rs[i]
w2c_1 = c2w_1.inverse()
c2ws.append(c2w_1)
w2cs.append(w2c_1)
intrinsic[i + 1, :2, :2] = intrinsic[i + 1, :2, :2] * focal_length * ((i + 1) / nframe) + \
intrinsic[i + 1, :2, :2] * ((nframe - (i + 1)) / nframe)
w2cs = torch.stack(w2cs, dim=0)
c2ws = torch.stack(c2ws, dim=0)
return w2cs, c2ws, intrinsic
View File
+125
View File
@@ -0,0 +1,125 @@
import cv2
import numpy as np
import torch
def nms(boxes, scores, nms_thr):
"""Single class NMS implemented in Numpy."""
x1 = boxes[:, 0]
y1 = boxes[:, 1]
x2 = boxes[:, 2]
y2 = boxes[:, 3]
areas = (x2 - x1 + 1) * (y2 - y1 + 1)
order = scores.argsort()[::-1]
keep = []
while order.size > 0:
i = order[0]
keep.append(i)
xx1 = np.maximum(x1[i], x1[order[1:]])
yy1 = np.maximum(y1[i], y1[order[1:]])
xx2 = np.minimum(x2[i], x2[order[1:]])
yy2 = np.minimum(y2[i], y2[order[1:]])
w = np.maximum(0.0, xx2 - xx1 + 1)
h = np.maximum(0.0, yy2 - yy1 + 1)
inter = w * h
ovr = inter / (areas[i] + areas[order[1:]] - inter)
inds = np.where(ovr <= nms_thr)[0]
order = order[inds + 1]
return keep
def multiclass_nms(boxes, scores, nms_thr, score_thr):
"""Multiclass NMS implemented in Numpy. Class-aware version."""
final_dets = []
num_classes = scores.shape[1]
for cls_ind in range(num_classes):
cls_scores = scores[:, cls_ind]
valid_score_mask = cls_scores > score_thr
if valid_score_mask.sum() == 0:
continue
else:
valid_scores = cls_scores[valid_score_mask]
valid_boxes = boxes[valid_score_mask]
keep = nms(valid_boxes, valid_scores, nms_thr)
if len(keep) > 0:
cls_inds = np.ones((len(keep), 1)) * cls_ind
dets = np.concatenate(
[valid_boxes[keep], valid_scores[keep, None], cls_inds], 1
)
final_dets.append(dets)
if len(final_dets) == 0:
return None
return np.concatenate(final_dets, 0)
def demo_postprocess(outputs, img_size, p6=False):
grids = []
expanded_strides = []
strides = [8, 16, 32] if not p6 else [8, 16, 32, 64]
hsizes = [img_size[0] // stride for stride in strides]
wsizes = [img_size[1] // stride for stride in strides]
for hsize, wsize, stride in zip(hsizes, wsizes, strides):
xv, yv = np.meshgrid(np.arange(wsize), np.arange(hsize))
grid = np.stack((xv, yv), 2).reshape(1, -1, 2)
grids.append(grid)
shape = grid.shape[:2]
expanded_strides.append(np.full((*shape, 1), stride))
grids = np.concatenate(grids, 1)
expanded_strides = np.concatenate(expanded_strides, 1)
outputs[..., :2] = (outputs[..., :2] + grids) * expanded_strides
outputs[..., 2:4] = np.exp(outputs[..., 2:4]) * expanded_strides
return outputs
def preprocess(img, input_size, swap=(2, 0, 1)):
if len(img.shape) == 3:
padded_img = np.ones((input_size[0], input_size[1], 3), dtype=np.uint8) * 114
else:
padded_img = np.ones(input_size, dtype=np.uint8) * 114
r = min(input_size[0] / img.shape[0], input_size[1] / img.shape[1])
resized_img = cv2.resize(
img,
(int(img.shape[1] * r), int(img.shape[0] * r)),
interpolation=cv2.INTER_LINEAR,
).astype(np.uint8)
padded_img[: int(img.shape[0] * r), : int(img.shape[1] * r)] = resized_img
padded_img = padded_img.transpose(swap)
padded_img = np.ascontiguousarray(padded_img, dtype=np.float32)
return padded_img, r
def inference_detector(model, oriImg, detect_classes=[0]):
input_shape = (640,640)
img, ratio = preprocess(oriImg, input_shape)
device, dtype = next(model.parameters()).device, next(model.parameters()).dtype
input = img[None, :, :, :]
input = torch.from_numpy(input).to(device, dtype)
output = model(input).float().cpu().detach().numpy()
predictions = demo_postprocess(output[0], input_shape)
boxes = predictions[:, :4]
scores = predictions[:, 4:5] * predictions[:, 5:]
boxes_xyxy = np.ones_like(boxes)
boxes_xyxy[:, 0] = boxes[:, 0] - boxes[:, 2]/2.
boxes_xyxy[:, 1] = boxes[:, 1] - boxes[:, 3]/2.
boxes_xyxy[:, 2] = boxes[:, 0] + boxes[:, 2]/2.
boxes_xyxy[:, 3] = boxes[:, 1] + boxes[:, 3]/2.
boxes_xyxy /= ratio
dets = multiclass_nms(boxes_xyxy, scores, nms_thr=0.45, score_thr=0.1)
if dets is None:
return None
final_boxes, final_scores, final_cls_inds = dets[:, :4], dets[:, 4], dets[:, 5]
isscore = final_scores>0.3
iscat = np.isin(final_cls_inds, detect_classes)
isbbox = [ i and j for (i, j) in zip(isscore, iscat)]
final_boxes = final_boxes[isbbox]
return final_boxes
+363
View File
@@ -0,0 +1,363 @@
from typing import List, Tuple
import cv2
import numpy as np
import torch
def preprocess(
img: np.ndarray, out_bbox, input_size: Tuple[int, int] = (192, 256)
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Do preprocessing for DWPose model inference.
Args:
img (np.ndarray): Input image in shape.
input_size (tuple): Input image size in shape (w, h).
Returns:
tuple:
- resized_img (np.ndarray): Preprocessed image.
- center (np.ndarray): Center of image.
- scale (np.ndarray): Scale of image.
"""
# get shape of image
img_shape = img.shape[:2]
out_img, out_center, out_scale = [], [], []
if len(out_bbox) == 0:
out_bbox = [[0, 0, img_shape[1], img_shape[0]]]
for i in range(len(out_bbox)):
x0 = out_bbox[i][0]
y0 = out_bbox[i][1]
x1 = out_bbox[i][2]
y1 = out_bbox[i][3]
bbox = np.array([x0, y0, x1, y1])
# get center and scale
center, scale = bbox_xyxy2cs(bbox, padding=1.25)
# do affine transformation
resized_img, scale = top_down_affine(input_size, scale, center, img)
# normalize image
mean = np.array([123.675, 116.28, 103.53])
std = np.array([58.395, 57.12, 57.375])
resized_img = (resized_img - mean) / std
out_img.append(resized_img)
out_center.append(center)
out_scale.append(scale)
return out_img, out_center, out_scale
def inference(model, img, bs=5):
"""Inference DWPose model implemented in TorchScript.
Args:
model : TorchScript Model.
img : Input image in shape.
Returns:
outputs : Output of DWPose model.
"""
all_out = []
# build input
orig_img_count = len(img)
#Pad zeros to fit batch size
for _ in range(bs - (orig_img_count % bs)):
img.append(np.zeros_like(img[0]))
input = np.stack(img, axis=0).transpose(0, 3, 1, 2)
device, dtype = next(model.parameters()).device, next(model.parameters()).dtype
input = torch.from_numpy(input).to(device, dtype)
out1, out2 = [], []
for i in range(input.shape[0] // bs):
curr_batch_output = model(input[i*bs:(i+1)*bs])
out1.append(curr_batch_output[0].float())
out2.append(curr_batch_output[1].float())
out1, out2 = torch.cat(out1, dim=0)[:orig_img_count], torch.cat(out2, dim=0)[:orig_img_count]
out1, out2 = out1.float().cpu().detach().numpy(), out2.float().cpu().detach().numpy()
all_outputs = out1, out2
for batch_idx in range(len(all_outputs[0])):
outputs = [all_outputs[i][batch_idx:batch_idx+1,...] for i in range(len(all_outputs))]
all_out.append(outputs)
return all_out
def postprocess(outputs: List[np.ndarray],
model_input_size: Tuple[int, int],
center: Tuple[int, int],
scale: Tuple[int, int],
simcc_split_ratio: float = 2.0
) -> Tuple[np.ndarray, np.ndarray]:
"""Postprocess for DWPose model output.
Args:
outputs (np.ndarray): Output of RTMPose model.
model_input_size (tuple): RTMPose model Input image size.
center (tuple): Center of bbox in shape (x, y).
scale (tuple): Scale of bbox in shape (w, h).
simcc_split_ratio (float): Split ratio of simcc.
Returns:
tuple:
- keypoints (np.ndarray): Rescaled keypoints.
- scores (np.ndarray): Model predict scores.
"""
all_key = []
all_score = []
for i in range(len(outputs)):
# use simcc to decode
simcc_x, simcc_y = outputs[i]
keypoints, scores = decode(simcc_x, simcc_y, simcc_split_ratio)
# rescale keypoints
keypoints = keypoints / model_input_size * scale[i] + center[i] - scale[i] / 2
all_key.append(keypoints[0])
all_score.append(scores[0])
return np.array(all_key), np.array(all_score)
def bbox_xyxy2cs(bbox: np.ndarray,
padding: float = 1.) -> Tuple[np.ndarray, np.ndarray]:
"""Transform the bbox format from (x,y,w,h) into (center, scale)
Args:
bbox (ndarray): Bounding box(es) in shape (4,) or (n, 4), formatted
as (left, top, right, bottom)
padding (float): BBox padding factor that will be multilied to scale.
Default: 1.0
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: Center (x, y) of the bbox in shape (2,) or
(n, 2)
- np.ndarray[float32]: Scale (w, h) of the bbox in shape (2,) or
(n, 2)
"""
# convert single bbox from (4, ) to (1, 4)
dim = bbox.ndim
if dim == 1:
bbox = bbox[None, :]
# get bbox center and scale
x1, y1, x2, y2 = np.hsplit(bbox, [1, 2, 3])
center = np.hstack([x1 + x2, y1 + y2]) * 0.5
scale = np.hstack([x2 - x1, y2 - y1]) * padding
if dim == 1:
center = center[0]
scale = scale[0]
return center, scale
def _fix_aspect_ratio(bbox_scale: np.ndarray,
aspect_ratio: float) -> np.ndarray:
"""Extend the scale to match the given aspect ratio.
Args:
scale (np.ndarray): The image scale (w, h) in shape (2, )
aspect_ratio (float): The ratio of ``w/h``
Returns:
np.ndarray: The reshaped image scale in (2, )
"""
w, h = np.hsplit(bbox_scale, [1])
bbox_scale = np.where(w > h * aspect_ratio,
np.hstack([w, w / aspect_ratio]),
np.hstack([h * aspect_ratio, h]))
return bbox_scale
def _rotate_point(pt: np.ndarray, angle_rad: float) -> np.ndarray:
"""Rotate a point by an angle.
Args:
pt (np.ndarray): 2D point coordinates (x, y) in shape (2, )
angle_rad (float): rotation angle in radian
Returns:
np.ndarray: Rotated point in shape (2, )
"""
sn, cs = np.sin(angle_rad), np.cos(angle_rad)
rot_mat = np.array([[cs, -sn], [sn, cs]])
return rot_mat @ pt
def _get_3rd_point(a: np.ndarray, b: np.ndarray) -> np.ndarray:
"""To calculate the affine matrix, three pairs of points are required. This
function is used to get the 3rd point, given 2D points a & b.
The 3rd point is defined by rotating vector `a - b` by 90 degrees
anticlockwise, using b as the rotation center.
Args:
a (np.ndarray): The 1st point (x,y) in shape (2, )
b (np.ndarray): The 2nd point (x,y) in shape (2, )
Returns:
np.ndarray: The 3rd point.
"""
direction = a - b
c = b + np.r_[-direction[1], direction[0]]
return c
def get_warp_matrix(center: np.ndarray,
scale: np.ndarray,
rot: float,
output_size: Tuple[int, int],
shift: Tuple[float, float] = (0., 0.),
inv: bool = False) -> np.ndarray:
"""Calculate the affine transformation matrix that can warp the bbox area
in the input image to the output size.
Args:
center (np.ndarray[2, ]): Center of the bounding box (x, y).
scale (np.ndarray[2, ]): Scale of the bounding box
wrt [width, height].
rot (float): Rotation angle (degree).
output_size (np.ndarray[2, ] | list(2,)): Size of the
destination heatmaps.
shift (0-100%): Shift translation ratio wrt the width/height.
Default (0., 0.).
inv (bool): Option to inverse the affine transform direction.
(inv=False: src->dst or inv=True: dst->src)
Returns:
np.ndarray: A 2x3 transformation matrix
"""
shift = np.array(shift)
src_w = scale[0]
dst_w = output_size[0]
dst_h = output_size[1]
# compute transformation matrix
rot_rad = np.deg2rad(rot)
src_dir = _rotate_point(np.array([0., src_w * -0.5]), rot_rad)
dst_dir = np.array([0., dst_w * -0.5])
# get four corners of the src rectangle in the original image
src = np.zeros((3, 2), dtype=np.float32)
src[0, :] = center + scale * shift
src[1, :] = center + src_dir + scale * shift
src[2, :] = _get_3rd_point(src[0, :], src[1, :])
# get four corners of the dst rectangle in the input image
dst = np.zeros((3, 2), dtype=np.float32)
dst[0, :] = [dst_w * 0.5, dst_h * 0.5]
dst[1, :] = np.array([dst_w * 0.5, dst_h * 0.5]) + dst_dir
dst[2, :] = _get_3rd_point(dst[0, :], dst[1, :])
if inv:
warp_mat = cv2.getAffineTransform(np.float32(dst), np.float32(src))
else:
warp_mat = cv2.getAffineTransform(np.float32(src), np.float32(dst))
return warp_mat
def top_down_affine(input_size: dict, bbox_scale: dict, bbox_center: dict,
img: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
"""Get the bbox image as the model input by affine transform.
Args:
input_size (dict): The input size of the model.
bbox_scale (dict): The bbox scale of the img.
bbox_center (dict): The bbox center of the img.
img (np.ndarray): The original image.
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: img after affine transform.
- np.ndarray[float32]: bbox scale after affine transform.
"""
w, h = input_size
warp_size = (int(w), int(h))
# reshape bbox to fixed aspect ratio
bbox_scale = _fix_aspect_ratio(bbox_scale, aspect_ratio=w / h)
# get the affine matrix
center = bbox_center
scale = bbox_scale
rot = 0
warp_mat = get_warp_matrix(center, scale, rot, output_size=(w, h))
# do affine transform
img = cv2.warpAffine(img, warp_mat, warp_size, flags=cv2.INTER_LINEAR)
return img, bbox_scale
def get_simcc_maximum(simcc_x: np.ndarray,
simcc_y: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
"""Get maximum response location and value from simcc representations.
Note:
instance number: N
num_keypoints: K
heatmap height: H
heatmap width: W
Args:
simcc_x (np.ndarray): x-axis SimCC in shape (K, Wx) or (N, K, Wx)
simcc_y (np.ndarray): y-axis SimCC in shape (K, Wy) or (N, K, Wy)
Returns:
tuple:
- locs (np.ndarray): locations of maximum heatmap responses in shape
(K, 2) or (N, K, 2)
- vals (np.ndarray): values of maximum heatmap responses in shape
(K,) or (N, K)
"""
N, K, Wx = simcc_x.shape
simcc_x = simcc_x.reshape(N * K, -1)
simcc_y = simcc_y.reshape(N * K, -1)
# get maximum value locations
x_locs = np.argmax(simcc_x, axis=1)
y_locs = np.argmax(simcc_y, axis=1)
locs = np.stack((x_locs, y_locs), axis=-1).astype(np.float32)
max_val_x = np.amax(simcc_x, axis=1)
max_val_y = np.amax(simcc_y, axis=1)
# get maximum value across x and y axis
mask = max_val_x > max_val_y
max_val_x[mask] = max_val_y[mask]
vals = max_val_x
locs[vals <= 0.] = -1
# reshape
locs = locs.reshape(N, K, 2)
vals = vals.reshape(N, K)
return locs, vals
def decode(simcc_x: np.ndarray, simcc_y: np.ndarray,
simcc_split_ratio) -> Tuple[np.ndarray, np.ndarray]:
"""Modulate simcc distribution with Gaussian.
Args:
simcc_x (np.ndarray[K, Wx]): model predicted simcc in x.
simcc_y (np.ndarray[K, Wy]): model predicted simcc in y.
simcc_split_ratio (int): The split ratio of simcc.
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: keypoints in shape (K, 2) or (n, K, 2)
- np.ndarray[float32]: scores in shape (K,) or (n, K)
"""
keypoints, scores = get_simcc_maximum(simcc_x, simcc_y)
keypoints /= simcc_split_ratio
return keypoints, scores
def inference_pose(model, out_bbox, oriImg, model_input_size=(288, 384)):
resized_img, center, scale = preprocess(oriImg, out_bbox, model_input_size)
#outputs = inference(session, resized_img, dtype)
outputs = inference(model, resized_img)
keypoints, scores = postprocess(outputs, model_input_size, center, scale)
return keypoints, scores
+127
View File
@@ -0,0 +1,127 @@
import cv2
import numpy as np
import onnxruntime
def nms(boxes, scores, nms_thr):
"""Single class NMS implemented in Numpy."""
x1 = boxes[:, 0]
y1 = boxes[:, 1]
x2 = boxes[:, 2]
y2 = boxes[:, 3]
areas = (x2 - x1 + 1) * (y2 - y1 + 1)
order = scores.argsort()[::-1]
keep = []
while order.size > 0:
i = order[0]
keep.append(i)
xx1 = np.maximum(x1[i], x1[order[1:]])
yy1 = np.maximum(y1[i], y1[order[1:]])
xx2 = np.minimum(x2[i], x2[order[1:]])
yy2 = np.minimum(y2[i], y2[order[1:]])
w = np.maximum(0.0, xx2 - xx1 + 1)
h = np.maximum(0.0, yy2 - yy1 + 1)
inter = w * h
ovr = inter / (areas[i] + areas[order[1:]] - inter)
inds = np.where(ovr <= nms_thr)[0]
order = order[inds + 1]
return keep
def multiclass_nms(boxes, scores, nms_thr, score_thr):
"""Multiclass NMS implemented in Numpy. Class-aware version."""
final_dets = []
num_classes = scores.shape[1]
for cls_ind in range(num_classes):
cls_scores = scores[:, cls_ind]
valid_score_mask = cls_scores > score_thr
if valid_score_mask.sum() == 0:
continue
else:
valid_scores = cls_scores[valid_score_mask]
valid_boxes = boxes[valid_score_mask]
keep = nms(valid_boxes, valid_scores, nms_thr)
if len(keep) > 0:
cls_inds = np.ones((len(keep), 1)) * cls_ind
dets = np.concatenate(
[valid_boxes[keep], valid_scores[keep, None], cls_inds], 1
)
final_dets.append(dets)
if len(final_dets) == 0:
return None
return np.concatenate(final_dets, 0)
def demo_postprocess(outputs, img_size, p6=False):
grids = []
expanded_strides = []
strides = [8, 16, 32] if not p6 else [8, 16, 32, 64]
hsizes = [img_size[0] // stride for stride in strides]
wsizes = [img_size[1] // stride for stride in strides]
for hsize, wsize, stride in zip(hsizes, wsizes, strides):
xv, yv = np.meshgrid(np.arange(wsize), np.arange(hsize))
grid = np.stack((xv, yv), 2).reshape(1, -1, 2)
grids.append(grid)
shape = grid.shape[:2]
expanded_strides.append(np.full((*shape, 1), stride))
grids = np.concatenate(grids, 1)
expanded_strides = np.concatenate(expanded_strides, 1)
outputs[..., :2] = (outputs[..., :2] + grids) * expanded_strides
outputs[..., 2:4] = np.exp(outputs[..., 2:4]) * expanded_strides
return outputs
def preprocess(img, input_size, swap=(2, 0, 1)):
if len(img.shape) == 3:
padded_img = np.ones((input_size[0], input_size[1], 3), dtype=np.uint8) * 114
else:
padded_img = np.ones(input_size, dtype=np.uint8) * 114
r = min(input_size[0] / img.shape[0], input_size[1] / img.shape[1])
resized_img = cv2.resize(
img,
(int(img.shape[1] * r), int(img.shape[0] * r)),
interpolation=cv2.INTER_LINEAR,
).astype(np.uint8)
padded_img[: int(img.shape[0] * r), : int(img.shape[1] * r)] = resized_img
padded_img = padded_img.transpose(swap)
padded_img = np.ascontiguousarray(padded_img, dtype=np.float32)
return padded_img, r
def inference_detector(session, oriImg):
input_shape = (640,640)
img, ratio = preprocess(oriImg, input_shape)
ort_inputs = {session.get_inputs()[0].name: img[None, :, :, :]}
output = session.run(None, ort_inputs)
predictions = demo_postprocess(output[0], input_shape)[0]
boxes = predictions[:, :4]
scores = predictions[:, 4:5] * predictions[:, 5:]
boxes_xyxy = np.ones_like(boxes)
boxes_xyxy[:, 0] = boxes[:, 0] - boxes[:, 2]/2.
boxes_xyxy[:, 1] = boxes[:, 1] - boxes[:, 3]/2.
boxes_xyxy[:, 2] = boxes[:, 0] + boxes[:, 2]/2.
boxes_xyxy[:, 3] = boxes[:, 1] + boxes[:, 3]/2.
boxes_xyxy /= ratio
dets = multiclass_nms(boxes_xyxy, scores, nms_thr=0.45, score_thr=0.1)
if dets is not None:
final_boxes, final_scores, final_cls_inds = dets[:, :4], dets[:, 4], dets[:, 5]
isscore = final_scores>0.3
iscat = final_cls_inds == 0
isbbox = [ i and j for (i, j) in zip(isscore, iscat)]
final_boxes = final_boxes[isbbox]
else:
final_boxes = np.array([])
return final_boxes
+360
View File
@@ -0,0 +1,360 @@
from typing import List, Tuple
import cv2
import numpy as np
import onnxruntime as ort
def preprocess(
img: np.ndarray, out_bbox, input_size: Tuple[int, int] = (192, 256)
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Do preprocessing for RTMPose model inference.
Args:
img (np.ndarray): Input image in shape.
input_size (tuple): Input image size in shape (w, h).
Returns:
tuple:
- resized_img (np.ndarray): Preprocessed image.
- center (np.ndarray): Center of image.
- scale (np.ndarray): Scale of image.
"""
# get shape of image
img_shape = img.shape[:2]
out_img, out_center, out_scale = [], [], []
if len(out_bbox) == 0:
out_bbox = [[0, 0, img_shape[1], img_shape[0]]]
for i in range(len(out_bbox)):
x0 = out_bbox[i][0]
y0 = out_bbox[i][1]
x1 = out_bbox[i][2]
y1 = out_bbox[i][3]
bbox = np.array([x0, y0, x1, y1])
# get center and scale
center, scale = bbox_xyxy2cs(bbox, padding=1.25)
# do affine transformation
resized_img, scale = top_down_affine(input_size, scale, center, img)
# normalize image
mean = np.array([123.675, 116.28, 103.53])
std = np.array([58.395, 57.12, 57.375])
resized_img = (resized_img - mean) / std
out_img.append(resized_img)
out_center.append(center)
out_scale.append(scale)
return out_img, out_center, out_scale
def inference(sess: ort.InferenceSession, img: np.ndarray) -> np.ndarray:
"""Inference RTMPose model.
Args:
sess (ort.InferenceSession): ONNXRuntime session.
img (np.ndarray): Input image in shape.
Returns:
outputs (np.ndarray): Output of RTMPose model.
"""
all_out = []
# build input
for i in range(len(img)):
input = [img[i].transpose(2, 0, 1)]
# build output
sess_input = {sess.get_inputs()[0].name: input}
sess_output = []
for out in sess.get_outputs():
sess_output.append(out.name)
# run model
outputs = sess.run(sess_output, sess_input)
all_out.append(outputs)
return all_out
def postprocess(outputs: List[np.ndarray],
model_input_size: Tuple[int, int],
center: Tuple[int, int],
scale: Tuple[int, int],
simcc_split_ratio: float = 2.0
) -> Tuple[np.ndarray, np.ndarray]:
"""Postprocess for RTMPose model output.
Args:
outputs (np.ndarray): Output of RTMPose model.
model_input_size (tuple): RTMPose model Input image size.
center (tuple): Center of bbox in shape (x, y).
scale (tuple): Scale of bbox in shape (w, h).
simcc_split_ratio (float): Split ratio of simcc.
Returns:
tuple:
- keypoints (np.ndarray): Rescaled keypoints.
- scores (np.ndarray): Model predict scores.
"""
all_key = []
all_score = []
for i in range(len(outputs)):
# use simcc to decode
simcc_x, simcc_y = outputs[i]
keypoints, scores = decode(simcc_x, simcc_y, simcc_split_ratio)
# rescale keypoints
keypoints = keypoints / model_input_size * scale[i] + center[i] - scale[i] / 2
all_key.append(keypoints[0])
all_score.append(scores[0])
return np.array(all_key), np.array(all_score)
def bbox_xyxy2cs(bbox: np.ndarray,
padding: float = 1.) -> Tuple[np.ndarray, np.ndarray]:
"""Transform the bbox format from (x,y,w,h) into (center, scale)
Args:
bbox (ndarray): Bounding box(es) in shape (4,) or (n, 4), formatted
as (left, top, right, bottom)
padding (float): BBox padding factor that will be multilied to scale.
Default: 1.0
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: Center (x, y) of the bbox in shape (2,) or
(n, 2)
- np.ndarray[float32]: Scale (w, h) of the bbox in shape (2,) or
(n, 2)
"""
# convert single bbox from (4, ) to (1, 4)
dim = bbox.ndim
if dim == 1:
bbox = bbox[None, :]
# get bbox center and scale
x1, y1, x2, y2 = np.hsplit(bbox, [1, 2, 3])
center = np.hstack([x1 + x2, y1 + y2]) * 0.5
scale = np.hstack([x2 - x1, y2 - y1]) * padding
if dim == 1:
center = center[0]
scale = scale[0]
return center, scale
def _fix_aspect_ratio(bbox_scale: np.ndarray,
aspect_ratio: float) -> np.ndarray:
"""Extend the scale to match the given aspect ratio.
Args:
scale (np.ndarray): The image scale (w, h) in shape (2, )
aspect_ratio (float): The ratio of ``w/h``
Returns:
np.ndarray: The reshaped image scale in (2, )
"""
w, h = np.hsplit(bbox_scale, [1])
bbox_scale = np.where(w > h * aspect_ratio,
np.hstack([w, w / aspect_ratio]),
np.hstack([h * aspect_ratio, h]))
return bbox_scale
def _rotate_point(pt: np.ndarray, angle_rad: float) -> np.ndarray:
"""Rotate a point by an angle.
Args:
pt (np.ndarray): 2D point coordinates (x, y) in shape (2, )
angle_rad (float): rotation angle in radian
Returns:
np.ndarray: Rotated point in shape (2, )
"""
sn, cs = np.sin(angle_rad), np.cos(angle_rad)
rot_mat = np.array([[cs, -sn], [sn, cs]])
return rot_mat @ pt
def _get_3rd_point(a: np.ndarray, b: np.ndarray) -> np.ndarray:
"""To calculate the affine matrix, three pairs of points are required. This
function is used to get the 3rd point, given 2D points a & b.
The 3rd point is defined by rotating vector `a - b` by 90 degrees
anticlockwise, using b as the rotation center.
Args:
a (np.ndarray): The 1st point (x,y) in shape (2, )
b (np.ndarray): The 2nd point (x,y) in shape (2, )
Returns:
np.ndarray: The 3rd point.
"""
direction = a - b
c = b + np.r_[-direction[1], direction[0]]
return c
def get_warp_matrix(center: np.ndarray,
scale: np.ndarray,
rot: float,
output_size: Tuple[int, int],
shift: Tuple[float, float] = (0., 0.),
inv: bool = False) -> np.ndarray:
"""Calculate the affine transformation matrix that can warp the bbox area
in the input image to the output size.
Args:
center (np.ndarray[2, ]): Center of the bounding box (x, y).
scale (np.ndarray[2, ]): Scale of the bounding box
wrt [width, height].
rot (float): Rotation angle (degree).
output_size (np.ndarray[2, ] | list(2,)): Size of the
destination heatmaps.
shift (0-100%): Shift translation ratio wrt the width/height.
Default (0., 0.).
inv (bool): Option to inverse the affine transform direction.
(inv=False: src->dst or inv=True: dst->src)
Returns:
np.ndarray: A 2x3 transformation matrix
"""
shift = np.array(shift)
src_w = scale[0]
dst_w = output_size[0]
dst_h = output_size[1]
# compute transformation matrix
rot_rad = np.deg2rad(rot)
src_dir = _rotate_point(np.array([0., src_w * -0.5]), rot_rad)
dst_dir = np.array([0., dst_w * -0.5])
# get four corners of the src rectangle in the original image
src = np.zeros((3, 2), dtype=np.float32)
src[0, :] = center + scale * shift
src[1, :] = center + src_dir + scale * shift
src[2, :] = _get_3rd_point(src[0, :], src[1, :])
# get four corners of the dst rectangle in the input image
dst = np.zeros((3, 2), dtype=np.float32)
dst[0, :] = [dst_w * 0.5, dst_h * 0.5]
dst[1, :] = np.array([dst_w * 0.5, dst_h * 0.5]) + dst_dir
dst[2, :] = _get_3rd_point(dst[0, :], dst[1, :])
if inv:
warp_mat = cv2.getAffineTransform(np.float32(dst), np.float32(src))
else:
warp_mat = cv2.getAffineTransform(np.float32(src), np.float32(dst))
return warp_mat
def top_down_affine(input_size: dict, bbox_scale: dict, bbox_center: dict,
img: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
"""Get the bbox image as the model input by affine transform.
Args:
input_size (dict): The input size of the model.
bbox_scale (dict): The bbox scale of the img.
bbox_center (dict): The bbox center of the img.
img (np.ndarray): The original image.
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: img after affine transform.
- np.ndarray[float32]: bbox scale after affine transform.
"""
w, h = input_size
warp_size = (int(w), int(h))
# reshape bbox to fixed aspect ratio
bbox_scale = _fix_aspect_ratio(bbox_scale, aspect_ratio=w / h)
# get the affine matrix
center = bbox_center
scale = bbox_scale
rot = 0
warp_mat = get_warp_matrix(center, scale, rot, output_size=(w, h))
# do affine transform
img = cv2.warpAffine(img, warp_mat, warp_size, flags=cv2.INTER_LINEAR)
return img, bbox_scale
def get_simcc_maximum(simcc_x: np.ndarray,
simcc_y: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
"""Get maximum response location and value from simcc representations.
Note:
instance number: N
num_keypoints: K
heatmap height: H
heatmap width: W
Args:
simcc_x (np.ndarray): x-axis SimCC in shape (K, Wx) or (N, K, Wx)
simcc_y (np.ndarray): y-axis SimCC in shape (K, Wy) or (N, K, Wy)
Returns:
tuple:
- locs (np.ndarray): locations of maximum heatmap responses in shape
(K, 2) or (N, K, 2)
- vals (np.ndarray): values of maximum heatmap responses in shape
(K,) or (N, K)
"""
N, K, Wx = simcc_x.shape
simcc_x = simcc_x.reshape(N * K, -1)
simcc_y = simcc_y.reshape(N * K, -1)
# get maximum value locations
x_locs = np.argmax(simcc_x, axis=1)
y_locs = np.argmax(simcc_y, axis=1)
locs = np.stack((x_locs, y_locs), axis=-1).astype(np.float32)
max_val_x = np.amax(simcc_x, axis=1)
max_val_y = np.amax(simcc_y, axis=1)
# get maximum value across x and y axis
mask = max_val_x > max_val_y
max_val_x[mask] = max_val_y[mask]
vals = max_val_x
locs[vals <= 0.] = -1
# reshape
locs = locs.reshape(N, K, 2)
vals = vals.reshape(N, K)
return locs, vals
def decode(simcc_x: np.ndarray, simcc_y: np.ndarray,
simcc_split_ratio) -> Tuple[np.ndarray, np.ndarray]:
"""Modulate simcc distribution with Gaussian.
Args:
simcc_x (np.ndarray[K, Wx]): model predicted simcc in x.
simcc_y (np.ndarray[K, Wy]): model predicted simcc in y.
simcc_split_ratio (int): The split ratio of simcc.
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: keypoints in shape (K, 2) or (n, K, 2)
- np.ndarray[float32]: scores in shape (K,) or (n, K)
"""
keypoints, scores = get_simcc_maximum(simcc_x, simcc_y)
keypoints /= simcc_split_ratio
return keypoints, scores
def inference_pose(session, out_bbox, oriImg):
h, w = session.get_inputs()[0].shape[2:]
model_input_size = (w, h)
resized_img, center, scale = preprocess(oriImg, out_bbox, model_input_size)
outputs = inference(session, resized_img)
keypoints, scores = postprocess(outputs, model_input_size, center, scale)
return keypoints, scores
+402
View File
@@ -0,0 +1,402 @@
import math
import numpy as np
import colorsys
import cv2
eps = 0.01
def smart_resize(x, s):
Ht, Wt = s
if x.ndim == 2:
Ho, Wo = x.shape
Co = 1
else:
Ho, Wo, Co = x.shape
if Co == 3 or Co == 1:
k = float(Ht + Wt) / float(Ho + Wo)
return cv2.resize(x, (int(Wt), int(Ht)), interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4)
else:
return np.stack([smart_resize(x[:, :, i], s) for i in range(Co)], axis=2)
def smart_resize_k(x, fx, fy):
if x.ndim == 2:
Ho, Wo = x.shape
Co = 1
else:
Ho, Wo, Co = x.shape
Ht, Wt = Ho * fy, Wo * fx
if Co == 3 or Co == 1:
k = float(Ht + Wt) / float(Ho + Wo)
return cv2.resize(x, (int(Wt), int(Ht)), interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4)
else:
return np.stack([smart_resize_k(x[:, :, i], fx, fy) for i in range(Co)], axis=2)
def padRightDownCorner(img, stride, padValue):
h = img.shape[0]
w = img.shape[1]
pad = 4 * [None]
pad[0] = 0 # up
pad[1] = 0 # left
pad[2] = 0 if (h % stride == 0) else stride - (h % stride) # down
pad[3] = 0 if (w % stride == 0) else stride - (w % stride) # right
img_padded = img
pad_up = np.tile(img_padded[0:1, :, :]*0 + padValue, (pad[0], 1, 1))
img_padded = np.concatenate((pad_up, img_padded), axis=0)
pad_left = np.tile(img_padded[:, 0:1, :]*0 + padValue, (1, pad[1], 1))
img_padded = np.concatenate((pad_left, img_padded), axis=1)
pad_down = np.tile(img_padded[-2:-1, :, :]*0 + padValue, (pad[2], 1, 1))
img_padded = np.concatenate((img_padded, pad_down), axis=0)
pad_right = np.tile(img_padded[:, -2:-1, :]*0 + padValue, (1, pad[3], 1))
img_padded = np.concatenate((img_padded, pad_right), axis=1)
return img_padded, pad
def transfer(model, model_weights):
transfered_model_weights = {}
for weights_name in model.state_dict().keys():
transfered_model_weights[weights_name] = model_weights['.'.join(weights_name.split('.')[1:])]
return transfered_model_weights
def draw_bodypose(canvas, candidate, subset):
H, W, C = canvas.shape
candidate = np.array(candidate)
subset = np.array(subset)
stickwidth = 4
limbSeq = [[2, 3], [2, 6], [3, 4], [4, 5], [6, 7], [7, 8], [2, 9], [9, 10], \
[10, 11], [2, 12], [12, 13], [13, 14], [2, 1], [1, 15], [15, 17], \
[1, 16], [16, 18], [3, 17], [6, 18]]
colors = [[255, 0, 0], [255, 85, 0], [255, 170, 0], [255, 255, 0], [170, 255, 0], [85, 255, 0], [0, 255, 0], \
[0, 255, 85], [0, 255, 170], [0, 255, 255], [0, 170, 255], [0, 85, 255], [0, 0, 255], [85, 0, 255], \
[170, 0, 255], [255, 0, 255], [255, 0, 170], [255, 0, 85]]
for i in range(17):
for n in range(len(subset)):
index = subset[n][np.array(limbSeq[i]) - 1]
if -1 in index:
continue
Y = candidate[index.astype(int), 0] * float(W)
X = candidate[index.astype(int), 1] * float(H)
mX = np.mean(X)
mY = np.mean(Y)
length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5
angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1]))
polygon = cv2.ellipse2Poly((int(mY), int(mX)), (int(length / 2), stickwidth), int(angle), 0, 360, 1)
cv2.fillConvexPoly(canvas, polygon, colors[i])
canvas = (canvas * 0.6).astype(np.uint8)
for i in range(18):
for n in range(len(subset)):
index = int(subset[n][i])
if index == -1:
continue
x, y = candidate[index][0:2]
x = int(x * W)
y = int(y * H)
cv2.circle(canvas, (int(x), int(y)), 4, colors[i], thickness=-1)
return canvas
def alpha_blend_color(color, alpha):
"""blend color according to point conf
"""
return [int(c * alpha) for c in color]
def draw_body_and_foot(canvas, candidate, subset, score, stick_width=4, draw_body=True, draw_feet=True, body_keypoint_size=4, draw_head=True):
H, W, C = canvas.shape
candidate = np.array(candidate)
subset = np.array(subset)
limbSeq_and_colors = []
if draw_body:
limbSeq_and_colors = [
([2, 3], [255, 0, 0]), # Neck to Right Shoulder
([2, 6], [255, 85, 0]), # Neck to Left Shoulder
([3, 4], [255, 170, 0]), # Right Shoulder to Right Elbow
([4, 5], [255, 255, 0]), # Right Elbow to Right Wrist
([6, 7], [170, 255, 0]), # Left Shoulder to Left Elbow
([7, 8], [85, 255, 0]), # Left Elbow to Left Wrist
([2, 9], [0, 255, 0]), # Neck to Right Hip
([9, 10], [0, 255, 85]), # Right Hip to Right Knee
([10, 11], [0, 255, 170]), # Right Knee to Right Ankle
([2, 12], [0, 255, 255]), # Neck to Left Hip
([12, 13], [0, 170, 255]), # Left Hip to Left Knee
([13, 14], [0, 85, 255]), # Left Knee to Left Ankle
]
else:
limbSeq_and_colors = [
([2, 3], [0, 0, 0]), # Neck to Right Shoulder
([2, 6], [0, 0, 0]), # Neck to Left Shoulder
([3, 4], [0, 0, 0]), # Right Shoulder to Right Elbow
([4, 5], [0, 0, 0]), # Right Elbow to Right Wrist
([6, 7], [0, 0, 0]), # Left Shoulder to Left Elbow
([7, 8], [0, 0, 0]), # Left Elbow to Left Wrist
([2, 9], [0, 0, 0]), # Neck to Right Hip
([9, 10], [0, 0, 0]), # Right Hip to Right Knee
([10, 11], [0, 0, 0]), # Right Knee to Right Ankle
([2, 12], [0, 0, 0]), # Neck to Left Hip
([12, 13], [0, 0, 0]), # Left Hip to Left Knee
([13, 14], [0, 0, 0]), # Left Knee to Left Ankle
]
# Conditionally add head-related elements
if draw_head:
head_elements = [
([2, 1], [0, 0, 255]), # Neck to Nose
([1, 15], [0, 0, 255]), # Nose to Right Eye
([15, 17], [85, 0, 255]), # Right Eye to Right Ear
([1, 16], [170, 0, 255]), # Nose to Left Eye
([16, 18], [255, 0, 255]), # Left Eye to Left Ear
([3, 17], [255, 0, 170]), # Right Shoulder to Right Ear
([6, 18], [255, 0, 85]) # Left Shoulder to Left Ear
]
else:
head_elements = [
([2, 1], [0, 0, 0]), # Neck to Nose
([1, 15], [0, 0, 0]), # Nose to Right Eye
([15, 17], [0, 0, 0]), # Right Eye to Right Ear
([1, 16], [0, 0, 0]), # Nose to Left Eye
([16, 18], [0, 0, 0]), # Left Eye to Left Ear
([3, 17], [0, 0, 0]), # Right Shoulder to Right Ear
([6, 18], [0, 0, 0]) # Left Shoulder to Left Ear
]
if draw_feet:
limbSeq_and_colors += [
([14, 19], [170, 255, 255]), # Left Ankle to Right Foot
([11, 20], [255, 255, 0]), # Right Ankle to Left Foot
]
# Append head elements based on the condition
limbSeq_and_colors += head_elements
for limb_info in limbSeq_and_colors[:17]:
limbSeq, color = limb_info
for n in range(len(subset)):
index = subset[n][np.array(limbSeq) - 1]
conf = score[n][np.array(limbSeq) - 1]
if conf[0] < 0.3 or conf[1] < 0.3:
continue
Y = candidate[index.astype(int), 0] * float(W)
X = candidate[index.astype(int), 1] * float(H)
mX = np.mean(X)
mY = np.mean(Y)
length = np.sqrt((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2)
angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1]))
polygon = cv2.ellipse2Poly((int(mY), int(mX)), (int(length / 2), stick_width), int(angle), 0, 360, 1)
cv2.fillConvexPoly(canvas, polygon, alpha_blend_color(color, conf[0] * conf[1]))
canvas = (canvas * 0.6).astype(np.uint8)
for limb_info in limbSeq_and_colors[:18]:
limbSeq, color = limb_info
for i in limbSeq:
for n in range(len(subset)):
index = int(subset[n][i - 1])
if index == -1:
continue
x, y = candidate[index][0:2]
conf = score[n][i - 1]
x = int(x * W)
y = int(y * H)
cv2.circle(canvas, (x, y), 4, alpha_blend_color(color, conf), thickness=-1)
return canvas
def draw_handpose(canvas, all_hand_peaks, draw_hands=True, hand_keypoint_size=4):
H, W, C = canvas.shape
edges = [[0, 1], [1, 2], [2, 3], [3, 4], [0, 5], [5, 6], [6, 7], [7, 8], [0, 9], [9, 10], \
[10, 11], [11, 12], [0, 13], [13, 14], [14, 15], [15, 16], [0, 17], [17, 18], [18, 19], [19, 20]]
for peaks in all_hand_peaks:
peaks = np.array(peaks)
if draw_hands:
for ie, e in enumerate(edges):
x1, y1 = peaks[e[0]]
x2, y2 = peaks[e[1]]
x1 = int(x1 * W)
y1 = int(y1 * H)
x2 = int(x2 * W)
y2 = int(y2 * H)
if x1 > eps and y1 > eps and x2 > eps and y2 > eps:
h = (ie / float(len(edges))) % 1.0
s, v = 1.0, 1.0
r, g, b = colorsys.hsv_to_rgb(h, s, v)
color = (int(255 * r), int(255 * g), int(255 * b))
cv2.line(canvas, (x1, y1), (x2, y2), color, thickness=2)
if hand_keypoint_size > 0:
for i, keypoint in enumerate(peaks):
x, y = keypoint
x = int(x * W)
y = int(y * H)
if x > eps and y > eps:
cv2.circle(canvas, (x, y), hand_keypoint_size, (0, 0, 255), thickness=-1)
return canvas
def draw_facepose(canvas, all_lmks):
H, W, C = canvas.shape
for lmks in all_lmks:
lmks = np.array(lmks)
for lmk in lmks:
x, y = lmk
x = int(x * W)
y = int(y * H)
if x > eps and y > eps:
cv2.circle(canvas, (x, y), 3, (255, 255, 255), thickness=-1)
return canvas
# detect hand according to body pose keypoints
# please refer to https://github.com/CMU-Perceptual-Computing-Lab/openpose/blob/master/src/openpose/hand/handDetector.cpp
def handDetect(candidate, subset, oriImg):
# right hand: wrist 4, elbow 3, shoulder 2
# left hand: wrist 7, elbow 6, shoulder 5
ratioWristElbow = 0.33
detect_result = []
image_height, image_width = oriImg.shape[0:2]
for person in subset.astype(int):
# if any of three not detected
has_left = np.sum(person[[5, 6, 7]] == -1) == 0
has_right = np.sum(person[[2, 3, 4]] == -1) == 0
if not (has_left or has_right):
continue
hands = []
#left hand
if has_left:
left_shoulder_index, left_elbow_index, left_wrist_index = person[[5, 6, 7]]
x1, y1 = candidate[left_shoulder_index][:2]
x2, y2 = candidate[left_elbow_index][:2]
x3, y3 = candidate[left_wrist_index][:2]
hands.append([x1, y1, x2, y2, x3, y3, True])
# right hand
if has_right:
right_shoulder_index, right_elbow_index, right_wrist_index = person[[2, 3, 4]]
x1, y1 = candidate[right_shoulder_index][:2]
x2, y2 = candidate[right_elbow_index][:2]
x3, y3 = candidate[right_wrist_index][:2]
hands.append([x1, y1, x2, y2, x3, y3, False])
for x1, y1, x2, y2, x3, y3, is_left in hands:
x = x3 + ratioWristElbow * (x3 - x2)
y = y3 + ratioWristElbow * (y3 - y2)
distanceWristElbow = math.sqrt((x3 - x2) ** 2 + (y3 - y2) ** 2)
distanceElbowShoulder = math.sqrt((x2 - x1) ** 2 + (y2 - y1) ** 2)
width = 1.5 * max(distanceWristElbow, 0.9 * distanceElbowShoulder)
# x-y refers to the center --> offset to topLeft point
# handRectangle.x -= handRectangle.width / 2.f;
# handRectangle.y -= handRectangle.height / 2.f;
x -= width / 2
y -= width / 2 # width = height
# overflow the image
if x < 0: x = 0
if y < 0: y = 0
width1 = width
width2 = width
if x + width > image_width: width1 = image_width - x
if y + width > image_height: width2 = image_height - y
width = min(width1, width2)
# the max hand box value is 20 pixels
if width >= 20:
detect_result.append([int(x), int(y), int(width), is_left])
'''
return value: [[x, y, w, True if left hand else False]].
width=height since the network require squared input.
x, y is the coordinate of top left
'''
return detect_result
# Written by Lvmin
def faceDetect(candidate, subset, oriImg):
# left right eye ear 14 15 16 17
detect_result = []
image_height, image_width = oriImg.shape[0:2]
for person in subset.astype(int):
has_head = person[0] > -1
if not has_head:
continue
has_left_eye = person[14] > -1
has_right_eye = person[15] > -1
has_left_ear = person[16] > -1
has_right_ear = person[17] > -1
if not (has_left_eye or has_right_eye or has_left_ear or has_right_ear):
continue
head, left_eye, right_eye, left_ear, right_ear = person[[0, 14, 15, 16, 17]]
width = 0.0
x0, y0 = candidate[head][:2]
if has_left_eye:
x1, y1 = candidate[left_eye][:2]
d = max(abs(x0 - x1), abs(y0 - y1))
width = max(width, d * 3.0)
if has_right_eye:
x1, y1 = candidate[right_eye][:2]
d = max(abs(x0 - x1), abs(y0 - y1))
width = max(width, d * 3.0)
if has_left_ear:
x1, y1 = candidate[left_ear][:2]
d = max(abs(x0 - x1), abs(y0 - y1))
width = max(width, d * 1.5)
if has_right_ear:
x1, y1 = candidate[right_ear][:2]
d = max(abs(x0 - x1), abs(y0 - y1))
width = max(width, d * 1.5)
x, y = x0, y0
x -= width
y -= width
if x < 0:
x = 0
if y < 0:
y = 0
width1 = width * 2
width2 = width * 2
if x + width > image_width:
width1 = image_width - x
if y + width > image_height:
width2 = image_height - y
width = min(width1, width2)
if width >= 20:
detect_result.append([int(x), int(y), int(width)])
return detect_result
# get max index of 2d array
def npmax(array):
arrayindex = array.argmax(1)
arrayvalue = array.max(1)
i = arrayvalue.argmax()
j = arrayindex[i]
return i, j
+42
View File
@@ -0,0 +1,42 @@
import numpy as np
from .jit_det import inference_detector as inference_jit_yolox
from .jit_pose import inference_pose as inference_jit_pose
import os
class Wholebody:
def __init__(self, model_det, model_pose):
self.model_det = model_det
self.model_pose = model_pose
def __call__(self, oriImg):
det_result = inference_jit_yolox(self.model_det, oriImg, detect_classes=[0])
keypoints, scores = inference_jit_pose(self.model_pose, det_result, oriImg)
keypoints_info = np.concatenate(
(keypoints, scores[..., None]), axis=-1)
# compute neck joint
neck = np.mean(keypoints_info[:, [5, 6]], axis=1)
# neck score when visualizing pred
neck[:, 2:4] = np.logical_and(
keypoints_info[:, 5, 2:4] > 0.3,
keypoints_info[:, 6, 2:4] > 0.3).astype(int)
new_keypoints_info = np.insert(
keypoints_info, 17, neck, axis=1)
mmpose_idx = [
17, 6, 8, 10, 7, 9, 12, 14, 16, 13, 15, 2, 1, 4, 3
]
openpose_idx = [
1, 2, 3, 4, 6, 7, 8, 9, 10, 12, 13, 14, 15, 16, 17
]
new_keypoints_info[:, openpose_idx] = \
new_keypoints_info[:, mmpose_idx]
keypoints_info = new_keypoints_info
keypoints, scores = keypoints_info[
..., :2], keypoints_info[..., 2]
return keypoints, scores
+836
View File
@@ -0,0 +1,836 @@
import torch.nn as nn
from ..utils import log
import comfy.model_management as mm
from comfy.utils import ProgressBar
from tqdm import tqdm
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))
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(concat_dim * 4, randomref_dim, 3, stride=2, padding=1),
)
state_dict_new = {}
for key in list(state_dict.keys()):
if "dwpose_embedding" in key:
state_dict_new[key.split("dwpose_embedding.")[1]] = state_dict.pop(key)
transformer.dwpose_embedding.load_state_dict(state_dict_new, strict=True)
state_dict_new = {}
for key in list(state_dict.keys()):
if "randomref_embedding_pose" in key:
state_dict_new[key.split("randomref_embedding_pose.")[1]] = state_dict.pop(key)
transformer.randomref_embedding_pose.load_state_dict(state_dict_new,strict=True)
return transformer
# Openpose
# Original from CMU https://github.com/CMU-Perceptual-Computing-Lab/openpose
# 2nd Edited by https://github.com/Hzzone/pytorch-openpose
# 3rd Edited by ControlNet
# 4th Edited by ControlNet (added face and correct hands)
import os
import torch
import numpy as np
import copy
import torch
import numpy as np
import math
from .dwpose.wholebody import Wholebody
def smoothing_factor(t_e, cutoff):
r = 2 * math.pi * cutoff * t_e
return r / (r + 1)
def exponential_smoothing(a, x, x_prev):
return a * x + (1 - a) * x_prev
class OneEuroFilter:
def __init__(self, t0, x0, dx0=0.0, min_cutoff=1.0, beta=0.0,
d_cutoff=1.0):
"""Initialize the one euro filter."""
# The parameters.
self.min_cutoff = float(min_cutoff)
self.beta = float(beta)
self.d_cutoff = float(d_cutoff)
# Previous values.
self.x_prev = x0
self.dx_prev = float(dx0)
self.t_prev = float(t0)
def __call__(self, t, x):
"""Compute the filtered signal."""
t_e = t - self.t_prev
# The filtered derivative of the signal.
a_d = smoothing_factor(t_e, self.d_cutoff)
dx = (x - self.x_prev) / t_e
dx_hat = exponential_smoothing(a_d, dx, self.dx_prev)
# The filtered signal.
cutoff = self.min_cutoff + self.beta * abs(dx_hat)
a = smoothing_factor(t_e, cutoff)
x_hat = exponential_smoothing(a, x, self.x_prev)
# Memorize the previous values.
self.x_prev = x_hat
self.dx_prev = dx_hat
self.t_prev = t
return x_hat
class DWposeDetector:
def __init__(self, model_det, model_pose):
self.pose_estimation = Wholebody(model_det, model_pose)
def __call__(self, oriImg, score_threshold=0.3):
oriImg = oriImg.copy()
H, W, C = oriImg.shape
with torch.no_grad():
candidate, subset = self.pose_estimation(oriImg)
candidate = candidate[0][np.newaxis, :, :]
subset = subset[0][np.newaxis, :]
nums, keys, locs = candidate.shape
candidate[..., 0] /= float(W)
candidate[..., 1] /= float(H)
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:
score[i][j] = int(18*i+j)
else:
score[i][j] = -1
un_visible = subset<score_threshold
candidate[un_visible] = -1
bodyfoot_score = subset[:,:24].copy()
for i in range(len(bodyfoot_score)):
for j in range(len(bodyfoot_score[i])):
if bodyfoot_score[i][j] > score_threshold:
bodyfoot_score[i][j] = int(18*i+j)
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.])
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.])
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
else:
bodyfoot[i][18] = np.array([-1., -1.])
if -1 not in bodyfoot[i][21] and -1 not in bodyfoot[i][22]:
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)
foot = candidate[:,18:24]
faces = candidate[:,24:92]
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)
# 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,
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']
faces = pose['faces']
hands = pose['hands']
candidate = bodies['candidate']
subset = bodies['subset']
score=bodies['score']
canvas = np.zeros(shape=(H, W, 3), dtype=np.uint8)
canvas = draw_body_and_foot(canvas, candidate, subset, score, draw_body=draw_body, stick_width=stick_width, draw_feet=draw_feet, draw_head=draw_head, body_keypoint_size=body_keypoint_size)
canvas = draw_handpose(canvas, hands, draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size)
canvas_without_face = copy.deepcopy(canvas)
canvas = draw_facepose(canvas, faces)
return canvas_without_face, canvas
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))
if ref_image is not None:
try:
pose_ref = dwpose_model(ref_image.squeeze(0), score_threshold=score_threshold)
except:
raise ValueError("No pose detected in reference image")
prev_pose = None
for img in tqdm(pose_images, desc="Pose Extraction", unit="image", total=len(pose_images)):
try:
pose = dwpose_model(img, score_threshold=score_threshold)
if handle_not_detected == "repeat":
prev_pose = pose
except:
if prev_pose is not None:
pose = prev_pose
else:
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']
candidate = bodies['candidate']
if ref_image is not None:
ref_bodies = pose_ref['bodies']
ref_faces = pose_ref['faces']
ref_hands = pose_ref['hands']
ref_candidate = ref_bodies['candidate']
ref_2_x = ref_candidate[2][0]
ref_2_y = ref_candidate[2][1]
ref_5_x = ref_candidate[5][0]
ref_5_y = ref_candidate[5][1]
ref_8_x = ref_candidate[8][0]
ref_8_y = ref_candidate[8][1]
ref_11_x = ref_candidate[11][0]
ref_11_y = ref_candidate[11][1]
ref_center1 = 0.5*(ref_candidate[2]+ref_candidate[5])
ref_center2 = 0.5*(ref_candidate[8]+ref_candidate[11])
zero_2_x = candidate[2][0]
zero_2_y = candidate[2][1]
zero_5_x = candidate[5][0]
zero_5_y = candidate[5][1]
zero_8_x = candidate[8][0]
zero_8_y = candidate[8][1]
zero_11_x = candidate[11][0]
zero_11_y = candidate[11][1]
zero_center1 = 0.5*(candidate[2]+candidate[5])
zero_center2 = 0.5*(candidate[8]+candidate[11])
x_ratio = (ref_5_x-ref_2_x)/(zero_5_x-zero_2_x)
y_ratio = (ref_center2[1]-ref_center1[1])/(zero_center2[1]-zero_center1[1])
results_vis[0]['bodies']['candidate'][:,0] *= x_ratio
results_vis[0]['bodies']['candidate'][:,1] *= y_ratio
results_vis[0]['faces'][:,:,0] *= x_ratio
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
neck_ratio = l_neck_ref / l_neck_0
x_offset_neck = (candidate[1][0]-candidate[0][0])*(1.-neck_ratio)
y_offset_neck = (candidate[1][1]-candidate[0][1])*(1.-neck_ratio)
results_vis[0]['bodies']['candidate'][0,0] += x_offset_neck
results_vis[0]['bodies']['candidate'][0,1] += y_offset_neck
results_vis[0]['bodies']['candidate'][14,0] += x_offset_neck
results_vis[0]['bodies']['candidate'][14,1] += y_offset_neck
results_vis[0]['bodies']['candidate'][15,0] += x_offset_neck
results_vis[0]['bodies']['candidate'][15,1] += y_offset_neck
results_vis[0]['bodies']['candidate'][16,0] += x_offset_neck
results_vis[0]['bodies']['candidate'][16,1] += y_offset_neck
results_vis[0]['bodies']['candidate'][17,0] += x_offset_neck
results_vis[0]['bodies']['candidate'][17,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
shoulder2_ratio = l_shoulder2_ref / l_shoulder2_0
x_offset_shoulder2 = (candidate[1][0]-candidate[2][0])*(1.-shoulder2_ratio)
y_offset_shoulder2 = (candidate[1][1]-candidate[2][1])*(1.-shoulder2_ratio)
results_vis[0]['bodies']['candidate'][2,0] += x_offset_shoulder2
results_vis[0]['bodies']['candidate'][2,1] += y_offset_shoulder2
results_vis[0]['bodies']['candidate'][3,0] += x_offset_shoulder2
results_vis[0]['bodies']['candidate'][3,1] += y_offset_shoulder2
results_vis[0]['bodies']['candidate'][4,0] += x_offset_shoulder2
results_vis[0]['bodies']['candidate'][4,1] += y_offset_shoulder2
results_vis[0]['hands'][1,:,0] += x_offset_shoulder2
results_vis[0]['hands'][1,:,1] += y_offset_shoulder2
########shoulder5########
l_shoulder5_ref = ((ref_candidate[5][0] - ref_candidate[1][0]) ** 2 + (ref_candidate[5][1] - ref_candidate[1][1]) ** 2) ** 0.5
l_shoulder5_0 = ((candidate[5][0] - candidate[1][0]) ** 2 + (candidate[5][1] - candidate[1][1]) ** 2) ** 0.5
shoulder5_ratio = l_shoulder5_ref / l_shoulder5_0
x_offset_shoulder5 = (candidate[1][0]-candidate[5][0])*(1.-shoulder5_ratio)
y_offset_shoulder5 = (candidate[1][1]-candidate[5][1])*(1.-shoulder5_ratio)
results_vis[0]['bodies']['candidate'][5,0] += x_offset_shoulder5
results_vis[0]['bodies']['candidate'][5,1] += y_offset_shoulder5
results_vis[0]['bodies']['candidate'][6,0] += x_offset_shoulder5
results_vis[0]['bodies']['candidate'][6,1] += y_offset_shoulder5
results_vis[0]['bodies']['candidate'][7,0] += x_offset_shoulder5
results_vis[0]['bodies']['candidate'][7,1] += y_offset_shoulder5
results_vis[0]['hands'][0,:,0] += x_offset_shoulder5
results_vis[0]['hands'][0,:,1] += y_offset_shoulder5
########arm3########
l_arm3_ref = ((ref_candidate[3][0] - ref_candidate[2][0]) ** 2 + (ref_candidate[3][1] - ref_candidate[2][1]) ** 2) ** 0.5
l_arm3_0 = ((candidate[3][0] - candidate[2][0]) ** 2 + (candidate[3][1] - candidate[2][1]) ** 2) ** 0.5
arm3_ratio = l_arm3_ref / l_arm3_0
x_offset_arm3 = (candidate[2][0]-candidate[3][0])*(1.-arm3_ratio)
y_offset_arm3 = (candidate[2][1]-candidate[3][1])*(1.-arm3_ratio)
results_vis[0]['bodies']['candidate'][3,0] += x_offset_arm3
results_vis[0]['bodies']['candidate'][3,1] += y_offset_arm3
results_vis[0]['bodies']['candidate'][4,0] += x_offset_arm3
results_vis[0]['bodies']['candidate'][4,1] += y_offset_arm3
results_vis[0]['hands'][1,:,0] += x_offset_arm3
results_vis[0]['hands'][1,:,1] += y_offset_arm3
########arm4########
l_arm4_ref = ((ref_candidate[4][0] - ref_candidate[3][0]) ** 2 + (ref_candidate[4][1] - ref_candidate[3][1]) ** 2) ** 0.5
l_arm4_0 = ((candidate[4][0] - candidate[3][0]) ** 2 + (candidate[4][1] - candidate[3][1]) ** 2) ** 0.5
arm4_ratio = l_arm4_ref / l_arm4_0
x_offset_arm4 = (candidate[3][0]-candidate[4][0])*(1.-arm4_ratio)
y_offset_arm4 = (candidate[3][1]-candidate[4][1])*(1.-arm4_ratio)
results_vis[0]['bodies']['candidate'][4,0] += x_offset_arm4
results_vis[0]['bodies']['candidate'][4,1] += y_offset_arm4
results_vis[0]['hands'][1,:,0] += x_offset_arm4
results_vis[0]['hands'][1,:,1] += y_offset_arm4
########arm6########
l_arm6_ref = ((ref_candidate[6][0] - ref_candidate[5][0]) ** 2 + (ref_candidate[6][1] - ref_candidate[5][1]) ** 2) ** 0.5
l_arm6_0 = ((candidate[6][0] - candidate[5][0]) ** 2 + (candidate[6][1] - candidate[5][1]) ** 2) ** 0.5
arm6_ratio = l_arm6_ref / l_arm6_0
x_offset_arm6 = (candidate[5][0]-candidate[6][0])*(1.-arm6_ratio)
y_offset_arm6 = (candidate[5][1]-candidate[6][1])*(1.-arm6_ratio)
results_vis[0]['bodies']['candidate'][6,0] += x_offset_arm6
results_vis[0]['bodies']['candidate'][6,1] += y_offset_arm6
results_vis[0]['bodies']['candidate'][7,0] += x_offset_arm6
results_vis[0]['bodies']['candidate'][7,1] += y_offset_arm6
results_vis[0]['hands'][0,:,0] += x_offset_arm6
results_vis[0]['hands'][0,:,1] += y_offset_arm6
########arm7########
l_arm7_ref = ((ref_candidate[7][0] - ref_candidate[6][0]) ** 2 + (ref_candidate[7][1] - ref_candidate[6][1]) ** 2) ** 0.5
l_arm7_0 = ((candidate[7][0] - candidate[6][0]) ** 2 + (candidate[7][1] - candidate[6][1]) ** 2) ** 0.5
arm7_ratio = l_arm7_ref / l_arm7_0
x_offset_arm7 = (candidate[6][0]-candidate[7][0])*(1.-arm7_ratio)
y_offset_arm7 = (candidate[6][1]-candidate[7][1])*(1.-arm7_ratio)
results_vis[0]['bodies']['candidate'][7,0] += x_offset_arm7
results_vis[0]['bodies']['candidate'][7,1] += y_offset_arm7
results_vis[0]['hands'][0,:,0] += x_offset_arm7
results_vis[0]['hands'][0,:,1] += y_offset_arm7
########head14########
l_head14_ref = ((ref_candidate[14][0] - ref_candidate[0][0]) ** 2 + (ref_candidate[14][1] - ref_candidate[0][1]) ** 2) ** 0.5
l_head14_0 = ((candidate[14][0] - candidate[0][0]) ** 2 + (candidate[14][1] - candidate[0][1]) ** 2) ** 0.5
head14_ratio = l_head14_ref / l_head14_0
x_offset_head14 = (candidate[0][0]-candidate[14][0])*(1.-head14_ratio)
y_offset_head14 = (candidate[0][1]-candidate[14][1])*(1.-head14_ratio)
results_vis[0]['bodies']['candidate'][14,0] += x_offset_head14
results_vis[0]['bodies']['candidate'][14,1] += y_offset_head14
results_vis[0]['bodies']['candidate'][16,0] += x_offset_head14
results_vis[0]['bodies']['candidate'][16,1] += y_offset_head14
########head15########
l_head15_ref = ((ref_candidate[15][0] - ref_candidate[0][0]) ** 2 + (ref_candidate[15][1] - ref_candidate[0][1]) ** 2) ** 0.5
l_head15_0 = ((candidate[15][0] - candidate[0][0]) ** 2 + (candidate[15][1] - candidate[0][1]) ** 2) ** 0.5
head15_ratio = l_head15_ref / l_head15_0
x_offset_head15 = (candidate[0][0]-candidate[15][0])*(1.-head15_ratio)
y_offset_head15 = (candidate[0][1]-candidate[15][1])*(1.-head15_ratio)
results_vis[0]['bodies']['candidate'][15,0] += x_offset_head15
results_vis[0]['bodies']['candidate'][15,1] += y_offset_head15
results_vis[0]['bodies']['candidate'][17,0] += x_offset_head15
results_vis[0]['bodies']['candidate'][17,1] += y_offset_head15
########head16########
l_head16_ref = ((ref_candidate[16][0] - ref_candidate[14][0]) ** 2 + (ref_candidate[16][1] - ref_candidate[14][1]) ** 2) ** 0.5
l_head16_0 = ((candidate[16][0] - candidate[14][0]) ** 2 + (candidate[16][1] - candidate[14][1]) ** 2) ** 0.5
head16_ratio = l_head16_ref / l_head16_0
x_offset_head16 = (candidate[14][0]-candidate[16][0])*(1.-head16_ratio)
y_offset_head16 = (candidate[14][1]-candidate[16][1])*(1.-head16_ratio)
results_vis[0]['bodies']['candidate'][16,0] += x_offset_head16
results_vis[0]['bodies']['candidate'][16,1] += y_offset_head16
########head17########
l_head17_ref = ((ref_candidate[17][0] - ref_candidate[15][0]) ** 2 + (ref_candidate[17][1] - ref_candidate[15][1]) ** 2) ** 0.5
l_head17_0 = ((candidate[17][0] - candidate[15][0]) ** 2 + (candidate[17][1] - candidate[15][1]) ** 2) ** 0.5
head17_ratio = l_head17_ref / l_head17_0
x_offset_head17 = (candidate[15][0]-candidate[17][0])*(1.-head17_ratio)
y_offset_head17 = (candidate[15][1]-candidate[17][1])*(1.-head17_ratio)
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
ll1_ratio = l_ll1_ref / l_ll1_0
x_offset_ll1 = (candidate[9][0]-candidate[8][0])*(ll1_ratio-1.)
y_offset_ll1 = (candidate[9][1]-candidate[8][1])*(ll1_ratio-1.)
results_vis[0]['bodies']['candidate'][9,0] += x_offset_ll1
results_vis[0]['bodies']['candidate'][9,1] += y_offset_ll1
results_vis[0]['bodies']['candidate'][10,0] += x_offset_ll1
results_vis[0]['bodies']['candidate'][10,1] += y_offset_ll1
results_vis[0]['bodies']['candidate'][19,0] += x_offset_ll1
results_vis[0]['bodies']['candidate'][19,1] += y_offset_ll1
l_ll2_ref = ((ref_candidate[9][0] - ref_candidate[10][0]) ** 2 + (ref_candidate[9][1] - ref_candidate[10][1]) ** 2) ** 0.5
l_ll2_0 = ((candidate[9][0] - candidate[10][0]) ** 2 + (candidate[9][1] - candidate[10][1]) ** 2) ** 0.5
ll2_ratio = l_ll2_ref / l_ll2_0
x_offset_ll2 = (candidate[10][0]-candidate[9][0])*(ll2_ratio-1.)
y_offset_ll2 = (candidate[10][1]-candidate[9][1])*(ll2_ratio-1.)
results_vis[0]['bodies']['candidate'][10,0] += x_offset_ll2
results_vis[0]['bodies']['candidate'][10,1] += y_offset_ll2
results_vis[0]['bodies']['candidate'][19,0] += x_offset_ll2
results_vis[0]['bodies']['candidate'][19,1] += y_offset_ll2
########right leg########
l_rl1_ref = ((ref_candidate[11][0] - ref_candidate[12][0]) ** 2 + (ref_candidate[11][1] - ref_candidate[12][1]) ** 2) ** 0.5
l_rl1_0 = ((candidate[11][0] - candidate[12][0]) ** 2 + (candidate[11][1] - candidate[12][1]) ** 2) ** 0.5
rl1_ratio = l_rl1_ref / l_rl1_0
x_offset_rl1 = (candidate[12][0]-candidate[11][0])*(rl1_ratio-1.)
y_offset_rl1 = (candidate[12][1]-candidate[11][1])*(rl1_ratio-1.)
results_vis[0]['bodies']['candidate'][12,0] += x_offset_rl1
results_vis[0]['bodies']['candidate'][12,1] += y_offset_rl1
results_vis[0]['bodies']['candidate'][13,0] += x_offset_rl1
results_vis[0]['bodies']['candidate'][13,1] += y_offset_rl1
results_vis[0]['bodies']['candidate'][18,0] += x_offset_rl1
results_vis[0]['bodies']['candidate'][18,1] += y_offset_rl1
l_rl2_ref = ((ref_candidate[12][0] - ref_candidate[13][0]) ** 2 + (ref_candidate[12][1] - ref_candidate[13][1]) ** 2) ** 0.5
l_rl2_0 = ((candidate[12][0] - candidate[13][0]) ** 2 + (candidate[12][1] - candidate[13][1]) ** 2) ** 0.5
rl2_ratio = l_rl2_ref / l_rl2_0
x_offset_rl2 = (candidate[13][0]-candidate[12][0])*(rl2_ratio-1.)
y_offset_rl2 = (candidate[13][1]-candidate[12][1])*(rl2_ratio-1.)
results_vis[0]['bodies']['candidate'][13,0] += x_offset_rl2
results_vis[0]['bodies']['candidate'][13,1] += y_offset_rl2
results_vis[0]['bodies']['candidate'][18,0] += x_offset_rl2
results_vis[0]['bodies']['candidate'][18,1] += y_offset_rl2
offset = ref_candidate[1] - results_vis[0]['bodies']['candidate'][1]
results_vis[0]['bodies']['candidate'] += offset[np.newaxis, :]
results_vis[0]['faces'] += offset[np.newaxis, np.newaxis, :]
results_vis[0]['hands'] += offset[np.newaxis, np.newaxis, :]
for i in range(1, len(results_vis)):
results_vis[i]['bodies']['candidate'][:,0] *= x_ratio
results_vis[i]['bodies']['candidate'][:,1] *= y_ratio
results_vis[i]['faces'][:,:,0] *= x_ratio
results_vis[i]['faces'][:,:,1] *= y_ratio
results_vis[i]['hands'][:,:,0] *= x_ratio
results_vis[i]['hands'][:,:,1] *= y_ratio
########neck########
x_offset_neck = (results_vis[i]['bodies']['candidate'][1][0]-results_vis[i]['bodies']['candidate'][0][0])*(1.-neck_ratio)
y_offset_neck = (results_vis[i]['bodies']['candidate'][1][1]-results_vis[i]['bodies']['candidate'][0][1])*(1.-neck_ratio)
results_vis[i]['bodies']['candidate'][0,0] += x_offset_neck
results_vis[i]['bodies']['candidate'][0,1] += y_offset_neck
results_vis[i]['bodies']['candidate'][14,0] += x_offset_neck
results_vis[i]['bodies']['candidate'][14,1] += y_offset_neck
results_vis[i]['bodies']['candidate'][15,0] += x_offset_neck
results_vis[i]['bodies']['candidate'][15,1] += y_offset_neck
results_vis[i]['bodies']['candidate'][16,0] += x_offset_neck
results_vis[i]['bodies']['candidate'][16,1] += y_offset_neck
results_vis[i]['bodies']['candidate'][17,0] += x_offset_neck
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)
results_vis[i]['bodies']['candidate'][2,0] += x_offset_shoulder2
results_vis[i]['bodies']['candidate'][2,1] += y_offset_shoulder2
results_vis[i]['bodies']['candidate'][3,0] += x_offset_shoulder2
results_vis[i]['bodies']['candidate'][3,1] += y_offset_shoulder2
results_vis[i]['bodies']['candidate'][4,0] += x_offset_shoulder2
results_vis[i]['bodies']['candidate'][4,1] += y_offset_shoulder2
results_vis[i]['hands'][1,:,0] += x_offset_shoulder2
results_vis[i]['hands'][1,:,1] += y_offset_shoulder2
########shoulder5########
x_offset_shoulder5 = (results_vis[i]['bodies']['candidate'][1][0]-results_vis[i]['bodies']['candidate'][5][0])*(1.-shoulder5_ratio)
y_offset_shoulder5 = (results_vis[i]['bodies']['candidate'][1][1]-results_vis[i]['bodies']['candidate'][5][1])*(1.-shoulder5_ratio)
results_vis[i]['bodies']['candidate'][5,0] += x_offset_shoulder5
results_vis[i]['bodies']['candidate'][5,1] += y_offset_shoulder5
results_vis[i]['bodies']['candidate'][6,0] += x_offset_shoulder5
results_vis[i]['bodies']['candidate'][6,1] += y_offset_shoulder5
results_vis[i]['bodies']['candidate'][7,0] += x_offset_shoulder5
results_vis[i]['bodies']['candidate'][7,1] += y_offset_shoulder5
results_vis[i]['hands'][0,:,0] += x_offset_shoulder5
results_vis[i]['hands'][0,:,1] += y_offset_shoulder5
########arm3########
x_offset_arm3 = (results_vis[i]['bodies']['candidate'][2][0]-results_vis[i]['bodies']['candidate'][3][0])*(1.-arm3_ratio)
y_offset_arm3 = (results_vis[i]['bodies']['candidate'][2][1]-results_vis[i]['bodies']['candidate'][3][1])*(1.-arm3_ratio)
results_vis[i]['bodies']['candidate'][3,0] += x_offset_arm3
results_vis[i]['bodies']['candidate'][3,1] += y_offset_arm3
results_vis[i]['bodies']['candidate'][4,0] += x_offset_arm3
results_vis[i]['bodies']['candidate'][4,1] += y_offset_arm3
results_vis[i]['hands'][1,:,0] += x_offset_arm3
results_vis[i]['hands'][1,:,1] += y_offset_arm3
########arm4########
x_offset_arm4 = (results_vis[i]['bodies']['candidate'][3][0]-results_vis[i]['bodies']['candidate'][4][0])*(1.-arm4_ratio)
y_offset_arm4 = (results_vis[i]['bodies']['candidate'][3][1]-results_vis[i]['bodies']['candidate'][4][1])*(1.-arm4_ratio)
results_vis[i]['bodies']['candidate'][4,0] += x_offset_arm4
results_vis[i]['bodies']['candidate'][4,1] += y_offset_arm4
results_vis[i]['hands'][1,:,0] += x_offset_arm4
results_vis[i]['hands'][1,:,1] += y_offset_arm4
########arm6########
x_offset_arm6 = (results_vis[i]['bodies']['candidate'][5][0]-results_vis[i]['bodies']['candidate'][6][0])*(1.-arm6_ratio)
y_offset_arm6 = (results_vis[i]['bodies']['candidate'][5][1]-results_vis[i]['bodies']['candidate'][6][1])*(1.-arm6_ratio)
results_vis[i]['bodies']['candidate'][6,0] += x_offset_arm6
results_vis[i]['bodies']['candidate'][6,1] += y_offset_arm6
results_vis[i]['bodies']['candidate'][7,0] += x_offset_arm6
results_vis[i]['bodies']['candidate'][7,1] += y_offset_arm6
results_vis[i]['hands'][0,:,0] += x_offset_arm6
results_vis[i]['hands'][0,:,1] += y_offset_arm6
########arm7########
x_offset_arm7 = (results_vis[i]['bodies']['candidate'][6][0]-results_vis[i]['bodies']['candidate'][7][0])*(1.-arm7_ratio)
y_offset_arm7 = (results_vis[i]['bodies']['candidate'][6][1]-results_vis[i]['bodies']['candidate'][7][1])*(1.-arm7_ratio)
results_vis[i]['bodies']['candidate'][7,0] += x_offset_arm7
results_vis[i]['bodies']['candidate'][7,1] += y_offset_arm7
results_vis[i]['hands'][0,:,0] += x_offset_arm7
results_vis[i]['hands'][0,:,1] += y_offset_arm7
########head14########
x_offset_head14 = (results_vis[i]['bodies']['candidate'][0][0]-results_vis[i]['bodies']['candidate'][14][0])*(1.-head14_ratio)
y_offset_head14 = (results_vis[i]['bodies']['candidate'][0][1]-results_vis[i]['bodies']['candidate'][14][1])*(1.-head14_ratio)
results_vis[i]['bodies']['candidate'][14,0] += x_offset_head14
results_vis[i]['bodies']['candidate'][14,1] += y_offset_head14
results_vis[i]['bodies']['candidate'][16,0] += x_offset_head14
results_vis[i]['bodies']['candidate'][16,1] += y_offset_head14
########head15########
x_offset_head15 = (results_vis[i]['bodies']['candidate'][0][0]-results_vis[i]['bodies']['candidate'][15][0])*(1.-head15_ratio)
y_offset_head15 = (results_vis[i]['bodies']['candidate'][0][1]-results_vis[i]['bodies']['candidate'][15][1])*(1.-head15_ratio)
results_vis[i]['bodies']['candidate'][15,0] += x_offset_head15
results_vis[i]['bodies']['candidate'][15,1] += y_offset_head15
results_vis[i]['bodies']['candidate'][17,0] += x_offset_head15
results_vis[i]['bodies']['candidate'][17,1] += y_offset_head15
########head16########
x_offset_head16 = (results_vis[i]['bodies']['candidate'][14][0]-results_vis[i]['bodies']['candidate'][16][0])*(1.-head16_ratio)
y_offset_head16 = (results_vis[i]['bodies']['candidate'][14][1]-results_vis[i]['bodies']['candidate'][16][1])*(1.-head16_ratio)
results_vis[i]['bodies']['candidate'][16,0] += x_offset_head16
results_vis[i]['bodies']['candidate'][16,1] += y_offset_head16
########head17########
x_offset_head17 = (results_vis[i]['bodies']['candidate'][15][0]-results_vis[i]['bodies']['candidate'][17][0])*(1.-head17_ratio)
y_offset_head17 = (results_vis[i]['bodies']['candidate'][15][1]-results_vis[i]['bodies']['candidate'][17][1])*(1.-head17_ratio)
results_vis[i]['bodies']['candidate'][17,0] += x_offset_head17
results_vis[i]['bodies']['candidate'][17,1] += y_offset_head17
# ########MovingAverage########
########left leg########
x_offset_ll1 = (results_vis[i]['bodies']['candidate'][9][0]-results_vis[i]['bodies']['candidate'][8][0])*(ll1_ratio-1.)
y_offset_ll1 = (results_vis[i]['bodies']['candidate'][9][1]-results_vis[i]['bodies']['candidate'][8][1])*(ll1_ratio-1.)
results_vis[i]['bodies']['candidate'][9,0] += x_offset_ll1
results_vis[i]['bodies']['candidate'][9,1] += y_offset_ll1
results_vis[i]['bodies']['candidate'][10,0] += x_offset_ll1
results_vis[i]['bodies']['candidate'][10,1] += y_offset_ll1
results_vis[i]['bodies']['candidate'][19,0] += x_offset_ll1
results_vis[i]['bodies']['candidate'][19,1] += y_offset_ll1
x_offset_ll2 = (results_vis[i]['bodies']['candidate'][10][0]-results_vis[i]['bodies']['candidate'][9][0])*(ll2_ratio-1.)
y_offset_ll2 = (results_vis[i]['bodies']['candidate'][10][1]-results_vis[i]['bodies']['candidate'][9][1])*(ll2_ratio-1.)
results_vis[i]['bodies']['candidate'][10,0] += x_offset_ll2
results_vis[i]['bodies']['candidate'][10,1] += y_offset_ll2
results_vis[i]['bodies']['candidate'][19,0] += x_offset_ll2
results_vis[i]['bodies']['candidate'][19,1] += y_offset_ll2
########right leg########
x_offset_rl1 = (results_vis[i]['bodies']['candidate'][12][0]-results_vis[i]['bodies']['candidate'][11][0])*(rl1_ratio-1.)
y_offset_rl1 = (results_vis[i]['bodies']['candidate'][12][1]-results_vis[i]['bodies']['candidate'][11][1])*(rl1_ratio-1.)
results_vis[i]['bodies']['candidate'][12,0] += x_offset_rl1
results_vis[i]['bodies']['candidate'][12,1] += y_offset_rl1
results_vis[i]['bodies']['candidate'][13,0] += x_offset_rl1
results_vis[i]['bodies']['candidate'][13,1] += y_offset_rl1
results_vis[i]['bodies']['candidate'][18,0] += x_offset_rl1
results_vis[i]['bodies']['candidate'][18,1] += y_offset_rl1
x_offset_rl2 = (results_vis[i]['bodies']['candidate'][13][0]-results_vis[i]['bodies']['candidate'][12][0])*(rl2_ratio-1.)
y_offset_rl2 = (results_vis[i]['bodies']['candidate'][13][1]-results_vis[i]['bodies']['candidate'][12][1])*(rl2_ratio-1.)
results_vis[i]['bodies']['candidate'][13,0] += x_offset_rl2
results_vis[i]['bodies']['candidate'][13,1] += y_offset_rl2
results_vis[i]['bodies']['candidate'][18,0] += x_offset_rl2
results_vis[i]['bodies']['candidate'][18,1] += y_offset_rl2
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:
dwpose_woface, dwpose_wface = draw_pose(results_vis[i], H=height, W=width, stick_width=stick_width,
draw_body=draw_body, draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size,
draw_feet=draw_feet, body_keypoint_size=body_keypoint_size, draw_head=draw_head)
result = torch.from_numpy(dwpose_woface)
#except:
# result = torch.zeros((height, width, 3), dtype=torch.uint8)
dwpose_woface_list.append(result)
dwpose_woface_tensor = torch.stack(dwpose_woface_list, dim=0)
dwpose_woface_ref_tensor = None
if ref_image is not None:
dwpose_woface_ref, dwpose_wface_ref = draw_pose(pose_ref, H=height, W=width, stick_width=stick_width,
draw_body=draw_body, draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size,
draw_feet=draw_feet, body_keypoint_size=body_keypoint_size, draw_head=draw_head)
dwpose_woface_ref_tensor = torch.from_numpy(dwpose_woface_ref)
return dwpose_woface_tensor, dwpose_woface_ref_tensor
class WanVideoUniAnimateDWPoseDetector:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"pose_images": ("IMAGE", {"tooltip": "Pose images"}),
"score_threshold": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Score threshold for pose detection"}),
"stick_width": ("INT", {"default": 4, "min": 1, "max": 100, "step": 1, "tooltip": "Stick width for drawing keypoints"}),
"draw_body": ("BOOLEAN", {"default": True, "tooltip": "Draw body keypoints"}),
"body_keypoint_size": ("INT", {"default": 4, "min": 0, "max": 100, "step": 1, "tooltip": "Body keypoint size"}),
"draw_feet": ("BOOLEAN", {"default": True, "tooltip": "Draw feet keypoints"}),
"draw_hands": ("BOOLEAN", {"default": True, "tooltip": "Draw hand keypoints"}),
"hand_keypoint_size": ("INT", {"default": 4, "min": 0, "max": 100, "step": 1, "tooltip": "Hand keypoint size"}),
"colorspace": (["RGB", "BGR"], {"tooltip": "Color space for the output image"}),
"handle_not_detected": (["empty", "repeat"], {"default": "empty", "tooltip": "How to handle undetected poses, empty inserts black and repeat inserts previous detection"}),
"draw_head": ("BOOLEAN", {"default": True, "tooltip": "Draw head keypoints"}),
},
"optional": {
"reference_pose_image": ("IMAGE", {"tooltip": "Reference pose image"}),
},
}
RETURN_TYPES = ("IMAGE", "IMAGE", )
RETURN_NAMES = ("poses", "reference_pose",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
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"
script_directory = os.path.dirname(os.path.abspath(__file__))
model_base_path = os.path.join(script_directory, "models", "DWPose")
model_det=os.path.join(model_base_path, yolo_model)
model_pose=os.path.join(model_base_path, dw_pose_model)
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",
allow_patterns=[f"*{yolo_model}*"],
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",
allow_patterns=[f"*{dw_pose_model}*"],
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)
#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:
ref = reference_pose_image
ref_np = ref.cpu().numpy() * 255
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_hands=draw_hands, hand_keypoint_size=hand_keypoint_size, handle_not_detected=handle_not_detected, draw_head=draw_head)
poses = poses / 255.0
if reference_pose_image is not None:
reference_pose = reference_pose.unsqueeze(0) / 255.0
else:
reference_pose = torch.zeros(1, 64, 64, 3, device=torch.device("cpu"))
if colorspace == "BGR":
poses=torch.flip(poses, dims=[-1])
return (poses, reference_pose, )
class WanVideoUniAnimatePoseInput:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"pose_images": ("IMAGE", {"tooltip": "Pose images"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.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 for the pose control"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage for the pose control"}),
},
"optional": {
"reference_pose_image": ("IMAGE", {"tooltip": "Reference pose image"}),
},
}
RETURN_TYPES = ("UNIANIMATE_POSE", )
RETURN_NAMES = ("unianimate_poses",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, pose_images, strength, start_percent, end_percent, reference_pose_image=None):
pose = pose_images.permute(3, 0, 1, 2).unsqueeze(0).contiguous()
ref = None
if reference_pose_image is not None:
ref = reference_pose_image.permute(0, 3, 1, 2).contiguous()
unianim_poses = {
"pose": pose,
"ref": ref,
"strength": strength,
"start_percent": start_percent,
"end_percent": end_percent
}
return (unianim_poses,)
NODE_CLASS_MAPPINGS = {
"WanVideoUniAnimatePoseInput": WanVideoUniAnimatePoseInput,
"WanVideoUniAnimateDWPoseDetector": WanVideoUniAnimateDWPoseDetector,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoUniAnimatePoseInput": "WanVideo UniAnimate Pose Input",
"WanVideoUniAnimateDWPoseDetector": "WanVideo UniAnimate DWPose Detector",
}
+68 -7
View File
@@ -58,16 +58,21 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
name = name.replace("._orig_mod.", ".") # torch compiled modules have this prefix
if low_mem_load:
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
if "modulation" in name:
if "patch_embedding" in name:
dtype_to_use = torch.float32
if name.startswith("diffusion_model."):
name_no_prefix = name[len("diffusion_model."):]
key = "{}.{}".format(name_no_prefix, param)
set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[key])
try:
set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[key])
except:
continue
model.patch_weight_to_device("{}.{}".format(name, param), device_to=device_to)
if low_mem_load:
set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=model.model.diffusion_model.state_dict()[key])
try:
set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=model.model.diffusion_model.state_dict()[key])
except:
continue
m.comfy_patched_weights = True
model.current_weight_patches_uuid = model.patches_uuid
@@ -75,9 +80,12 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
for name, param in model.model.diffusion_model.named_parameters():
if param.device != transformer_load_device:
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
if "modulation" in name:
if "patch_embedding" in name:
dtype_to_use = torch.float32
set_module_tensor_to_device(model.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[name])
try:
set_module_tensor_to_device(model.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[name])
except:
continue
return model
@@ -164,4 +172,57 @@ def encode_image_(clip_vision, image):
pixel_values = clip_preprocess(image, size=224, crop=True).float()
out = clip_vision.visual(pixel_values)
return out
return out
# Code based on https://github.com/WikiChao/FreSca (MIT License)
import torch
import torch.fft as fft
def fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20):
"""
Apply frequency-dependent scaling to an image tensor using Fourier transforms.
Parameters:
x: Input tensor of shape (B, C, H, W)
scale_low: Scaling factor for low-frequency components (default: 1.0)
scale_high: Scaling factor for high-frequency components (default: 1.5)
freq_cutoff: Number of frequency indices around center to consider as low-frequency (default: 20)
Returns:
x_filtered: Filtered version of x in spatial domain with frequency-specific scaling applied.
"""
# Preserve input dtype and device
dtype, device = x.dtype, x.device
# Convert to float32 for FFT computations
x = x.to(torch.float32)
# 1) Apply FFT and shift low frequencies to center
x_freq = fft.fftn(x, dim=(-2, -1))
x_freq = fft.fftshift(x_freq, dim=(-2, -1))
# 2) Create a mask to scale frequencies differently
C, B, H, W = x_freq.shape
crow, ccol = H // 2, W // 2
# Initialize mask with high-frequency scaling factor
mask = torch.ones((C, B, H, W), device=device) * scale_high
# Apply low-frequency scaling factor to center region
mask[
...,
crow - freq_cutoff : crow + freq_cutoff,
ccol - freq_cutoff : ccol + freq_cutoff,
] = scale_low
# 3) Apply frequency-specific scaling
x_freq = x_freq * mask
# 4) Convert back to spatial domain
x_freq = fft.ifftshift(x_freq, dim=(-2, -1))
x_filtered = fft.ifftn(x_freq, dim=(-2, -1)).real
# 5) Restore original dtype
x_filtered = x_filtered.to(dtype)
return x_filtered
+7 -4
View File
@@ -17,7 +17,10 @@ try:
from sageattention import sageattn
@torch.compiler.disable()
def sageattn_func(q, k, v, attn_mask=None, dropout_p=0, is_causal=False):
return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal)
if q.dtype == torch.float32:
return sageattn(q.to(torch.float16), k.to(torch.float16), v.to(torch.float16), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal).to(torch.float32)
else:
return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal)
except Exception as e:
print(f"Warning: Could not load sageattention: {str(e)}")
if isinstance(e, ModuleNotFoundError):
@@ -196,9 +199,9 @@ def attention(
elif attention_mode == 'sageattn':
attn_mask = None
q = q.transpose(1, 2).to(dtype)
k = k.transpose(1, 2).to(dtype)
v = v.transpose(1, 2).to(dtype)
q = q.transpose(1, 2)
k = k.transpose(1, 2)
v = v.transpose(1, 2)
out = sageattn_func(
q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p)
+629 -85
View File
File diff suppressed because it is too large Load Diff
+2
View File
@@ -516,4 +516,6 @@ class T5EncoderModel:
mask = mask.to(device)
seq_lens = mask.gt(0).sum(dim=1).long()
context = self.model(ids, mask)
for u, v in zip(context, seq_lens):
u[v:] = 0.0 # set padding to 0.0
return [u[:v] for u, v in zip(context, seq_lens)]
+63
View File
@@ -0,0 +1,63 @@
#https://github.com/aigc-apps/VideoX-Fun/blob/wan_fun_v1.1/videox_fun/models/wan_camera_adapter.py
import torch
import torch.nn as nn
class SimpleAdapter(nn.Module):
def __init__(self, in_dim, out_dim, kernel_size, stride, num_residual_blocks=1):
super(SimpleAdapter, self).__init__()
# Pixel Unshuffle: reduce spatial dimensions by a factor of 8
self.pixel_unshuffle = nn.PixelUnshuffle(downscale_factor=8)
# Convolution: reduce spatial dimensions by a factor
# of 2 (without overlap)
self.conv = nn.Conv2d(in_dim * 64, out_dim, kernel_size=kernel_size, stride=stride, padding=0)
# Residual blocks for feature extraction
self.residual_blocks = nn.Sequential(
*[ResidualBlock(out_dim) for _ in range(num_residual_blocks)]
)
def forward(self, x):
# Reshape to merge the frame dimension into batch
bs, c, f, h, w = x.size()
x = x.permute(0, 2, 1, 3, 4).contiguous().view(bs * f, c, h, w)
# Pixel Unshuffle operation
x_unshuffled = self.pixel_unshuffle(x)
# Convolution operation
x_conv = self.conv(x_unshuffled)
# Feature extraction with residual blocks
out = self.residual_blocks(x_conv)
# Reshape to restore original bf dimension
out = out.view(bs, f, out.size(1), out.size(2), out.size(3))
# Permute dimensions to reorder (if needed), e.g., swap channels and feature frames
out = out.permute(0, 2, 1, 3, 4)
return out
class ResidualBlock(nn.Module):
def __init__(self, dim):
super(ResidualBlock, self).__init__()
self.conv1 = nn.Conv2d(dim, dim, kernel_size=3, padding=1)
self.relu = nn.ReLU(inplace=True)
self.conv2 = nn.Conv2d(dim, dim, kernel_size=3, padding=1)
def forward(self, x):
residual = x
out = self.relu(self.conv1(x))
out = self.conv2(out)
out += residual
return out
# Example usage
# in_dim = 3
# out_dim = 64
# adapter = SimpleAdapterWithReshape(in_dim, out_dim)
# x = torch.randn(1, in_dim, 4, 64, 64) # e.g., batch size = 1, channels = 3, frames/features = 4
# output = adapter(x)
# print(output.shape) # Should reflect transformed dimensions
+95
View File
@@ -0,0 +1,95 @@
"""
The following code is copied from https://github.com/modelscope/DiffSynth-Studio/blob/main/diffsynth/schedulers/flow_match.py
"""
import torch
class FlowMatchScheduler():
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0, sigma_max=1.0, sigma_min=0.003 / 1.002, inverse_timesteps=False, extra_one_step=False, reverse_sigmas=False):
self.num_train_timesteps = num_train_timesteps
self.shift = shift
self.sigma_max = sigma_max
self.sigma_min = sigma_min
self.inverse_timesteps = inverse_timesteps
self.extra_one_step = extra_one_step
self.reverse_sigmas = reverse_sigmas
self.set_timesteps(num_inference_steps)
def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, training=False):
sigma_start = self.sigma_min + \
(self.sigma_max - self.sigma_min) * denoising_strength
if self.extra_one_step:
self.sigmas = torch.linspace(
sigma_start, self.sigma_min, num_inference_steps + 1)[:-1]
else:
self.sigmas = torch.linspace(
sigma_start, self.sigma_min, num_inference_steps)
if self.inverse_timesteps:
self.sigmas = torch.flip(self.sigmas, dims=[0])
self.sigmas = self.shift * self.sigmas / \
(1 + (self.shift - 1) * self.sigmas)
if self.reverse_sigmas:
self.sigmas = 1 - self.sigmas
self.timesteps = self.sigmas * self.num_train_timesteps
if training:
x = self.timesteps
y = torch.exp(-2 * ((x - num_inference_steps / 2) /
num_inference_steps) ** 2)
y_shifted = y - y.min()
bsmntw_weighing = y_shifted * \
(num_inference_steps / y_shifted.sum())
self.linear_timesteps_weights = bsmntw_weighing
def step(self, model_output, timestep, sample, to_final=False):
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)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
if to_final or (timestep_id + 1 >= len(self.timesteps)).any():
sigma_ = 1 if (
self.inverse_timesteps or self.reverse_sigmas) else 0
else:
sigma_ = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
prev_sample = sample + model_output * (sigma_ - sigma)
return prev_sample
def add_noise(self, original_samples, noise, timestep):
"""
Diffusion forward corruption process.
Input:
- clean_latent: the clean latent with shape [B*T, C, H, W]
- noise: the noise with shape [B*T, C, H, W]
- timestep: the timestep with shape [B*T]
Output: the corrupted latent with shape [B*T, C, H, W]
"""
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)
def training_target(self, sample, noise, timestep):
target = noise - sample
return target
def training_weight(self, timestep):
"""
Input:
- timestep: the timestep with shape [B*T]
Output: the corresponding weighting [B*T]
"""
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
self.linear_timesteps_weights = self.linear_timesteps_weights.to(timestep.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(1) - timestep.unsqueeze(0)).abs(), dim=0)
weights = self.linear_timesteps_weights[timestep_id]
return weights
+587
View File
@@ -0,0 +1,587 @@
# Copyright 2024 Stability AI, Katherine Crowson and 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.
import math
from dataclasses import dataclass
from typing import List, Optional, Tuple, Union
import numpy as np
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.utils import BaseOutput, is_scipy_available, logging
from diffusers.utils.torch_utils import randn_tensor
from diffusers.schedulers.scheduling_utils import SchedulerMixin
if is_scipy_available():
import scipy.stats
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@dataclass
class FlowMatchLCMSchedulerOutput(BaseOutput):
"""
Output class for the scheduler's `step` function output.
Args:
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
denoising loop.
"""
prev_sample: torch.FloatTensor
class FlowMatchLCMScheduler(SchedulerMixin, ConfigMixin):
"""
LCM scheduler for Flow Matching.
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
methods the library implements for all schedulers such as loading and saving.
Args:
num_train_timesteps (`int`, defaults to 1000):
The number of diffusion steps to train the model.
shift (`float`, defaults to 1.0):
The shift value for the timestep schedule.
use_dynamic_shifting (`bool`, defaults to False):
Whether to apply timestep shifting on-the-fly based on the image resolution.
base_shift (`float`, defaults to 0.5):
Value to stabilize image generation. Increasing `base_shift` reduces variation and image is more consistent
with desired output.
max_shift (`float`, defaults to 1.15):
Value change allowed to latent vectors. Increasing `max_shift` encourages more variation and image may be
more exaggerated or stylized.
base_image_seq_len (`int`, defaults to 256):
The base image sequence length.
max_image_seq_len (`int`, defaults to 4096):
The maximum image sequence length.
invert_sigmas (`bool`, defaults to False):
Whether to invert the sigmas.
shift_terminal (`float`, defaults to None):
The end value of the shifted timestep schedule.
use_karras_sigmas (`bool`, defaults to False):
Whether to use Karras sigmas for step sizes in the noise schedule during sampling.
use_exponential_sigmas (`bool`, defaults to False):
Whether to use exponential sigmas for step sizes in the noise schedule during sampling.
use_beta_sigmas (`bool`, defaults to False):
Whether to use beta sigmas for step sizes in the noise schedule during sampling.
time_shift_type (`str`, defaults to "exponential"):
The type of dynamic resolution-dependent timestep shifting to apply. Either "exponential" or "linear".
scale_factors ('list', defaults to None)
It defines how to scale the latents at which predictions are made.
upscale_mode ('str', defaults to 'bicubic')
Upscaling method, applied if scale-wise generation is considered
"""
_compatibles = []
order = 1
@register_to_config
def __init__(
self,
num_train_timesteps: int = 1000,
shift: float = 1.0,
use_dynamic_shifting: bool = False,
base_shift: Optional[float] = 0.5,
max_shift: Optional[float] = 1.15,
base_image_seq_len: Optional[int] = 256,
max_image_seq_len: Optional[int] = 4096,
invert_sigmas: bool = False,
shift_terminal: Optional[float] = None,
use_karras_sigmas: Optional[bool] = False,
use_exponential_sigmas: Optional[bool] = False,
use_beta_sigmas: Optional[bool] = False,
time_shift_type: str = "exponential",
scale_factors: Optional[List[float]] = None,
upscale_mode: Optional[str] = 'bicubic'
):
if self.config.use_beta_sigmas and not is_scipy_available():
raise ImportError("Make sure to install scipy if you want to use beta sigmas.")
if sum([self.config.use_beta_sigmas, self.config.use_exponential_sigmas, self.config.use_karras_sigmas]) > 1:
raise ValueError(
"Only one of `config.use_beta_sigmas`, `config.use_exponential_sigmas`, `config.use_karras_sigmas` can be used."
)
if time_shift_type not in {"exponential", "linear"}:
raise ValueError("`time_shift_type` must either be 'exponential' or 'linear'.")
timesteps = np.linspace(1, num_train_timesteps, num_train_timesteps, dtype=np.float32)[::-1].copy()
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
sigmas = timesteps / num_train_timesteps
if not use_dynamic_shifting:
# when use_dynamic_shifting is True, we apply the timestep shifting on the fly based on the image resolution
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
self.timesteps = sigmas * num_train_timesteps
self._step_index = None
self._begin_index = None
self._shift = shift
self._init_size = None
self._scale_factors = scale_factors
self._upscale_mode = upscale_mode
self.sigmas = sigmas.to("cpu") # to avoid too much CPU/GPU communication
self.sigma_min = self.sigmas[-1].item()
self.sigma_max = self.sigmas[0].item()
@property
def shift(self):
"""
The value used for shifting.
"""
return self._shift
@property
def step_index(self):
"""
The index counter for current timestep. It will increase 1 after each scheduler step.
"""
return self._step_index
@property
def begin_index(self):
"""
The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
"""
return self._begin_index
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
def set_begin_index(self, begin_index: int = 0):
"""
Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
Args:
begin_index (`int`):
The begin index for the scheduler.
"""
self._begin_index = begin_index
def set_shift(self, shift: float):
self._shift = shift
def set_scale_factors(self, scale_factors: list, upscale_mode):
"""
Sets scale factors for a scale-wise generation regime.
Args:
scale_factors (`list`):
The scale factors for each step
upscale_mode (`str`):
Upscaling method
"""
self._scale_factors = scale_factors
self._upscale_mode = upscale_mode
def scale_noise(
self,
sample: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
noise: Optional[torch.FloatTensor] = None,
) -> torch.FloatTensor:
"""
Forward process in flow-matching
Args:
sample (`torch.FloatTensor`):
The input sample.
timestep (`int`, *optional*):
The current timestep in the diffusion chain.
Returns:
`torch.FloatTensor`:
A scaled input sample.
"""
# Make sure sigmas and timesteps have the same device and dtype as original_samples
sigmas = self.sigmas.to(device=sample.device, dtype=sample.dtype)
if sample.device.type == "mps" and torch.is_floating_point(timestep):
# mps does not support float64
schedule_timesteps = self.timesteps.to(sample.device, dtype=torch.float32)
timestep = timestep.to(sample.device, dtype=torch.float32)
else:
schedule_timesteps = self.timesteps.to(sample.device)
timestep = timestep.to(sample.device)
# self.begin_index is None when scheduler is used for training, or pipeline does not implement set_begin_index
if self.begin_index is None:
step_indices = [self.index_for_timestep(t, schedule_timesteps) for t in timestep]
elif self.step_index is not None:
# add_noise is called after first denoising step (for inpainting)
step_indices = [self.step_index] * timestep.shape[0]
else:
# add noise is called before first denoising step to create initial latent(img2img)
step_indices = [self.begin_index] * timestep.shape[0]
sigma = sigmas[step_indices].flatten()
while len(sigma.shape) < len(sample.shape):
sigma = sigma.unsqueeze(-1)
sample = sigma * noise + (1.0 - sigma) * sample
return sample
def _sigma_to_t(self, sigma):
return sigma * self.config.num_train_timesteps
def time_shift(self, mu: float, sigma: float, t: torch.Tensor):
if self.config.time_shift_type == "exponential":
return self._time_shift_exponential(mu, sigma, t)
elif self.config.time_shift_type == "linear":
return self._time_shift_linear(mu, sigma, t)
def stretch_shift_to_terminal(self, t: torch.Tensor) -> torch.Tensor:
r"""
Stretches and shifts the timestep schedule to ensure it terminates at the configured `shift_terminal` config
value.
Reference:
https://github.com/Lightricks/LTX-Video/blob/a01a171f8fe3d99dce2728d60a73fecf4d4238ae/ltx_video/schedulers/rf.py#L51
Args:
t (`torch.Tensor`):
A tensor of timesteps to be stretched and shifted.
Returns:
`torch.Tensor`:
A tensor of adjusted timesteps such that the final value equals `self.config.shift_terminal`.
"""
one_minus_z = 1 - t
scale_factor = one_minus_z[-1] / (1 - self.config.shift_terminal)
stretched_t = 1 - (one_minus_z / scale_factor)
return stretched_t
def set_timesteps(
self,
num_inference_steps: Optional[int] = None,
device: Union[str, torch.device] = None,
sigmas: Optional[List[float]] = None,
mu: Optional[float] = None,
timesteps: Optional[List[float]] = None,
):
"""
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
Args:
num_inference_steps (`int`, *optional*):
The number of diffusion steps used when generating samples with a pre-trained model.
device (`str` or `torch.device`, *optional*):
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
sigmas (`List[float]`, *optional*):
Custom values for sigmas to be used for each diffusion step. If `None`, the sigmas are computed
automatically.
mu (`float`, *optional*):
Determines the amount of shifting applied to sigmas when performing resolution-dependent timestep
shifting.
timesteps (`List[float]`, *optional*):
Custom values for timesteps to be used for each diffusion step. If `None`, the timesteps are computed
automatically.
"""
if self.config.use_dynamic_shifting and mu is None:
raise ValueError("`mu` must be passed when `use_dynamic_shifting` is set to be `True`")
if sigmas is not None and timesteps is not None:
if len(sigmas) != len(timesteps):
raise ValueError("`sigmas` and `timesteps` should have the same length")
if num_inference_steps is not None:
if (sigmas is not None and len(sigmas) != num_inference_steps) or (
timesteps is not None and len(timesteps) != num_inference_steps
):
raise ValueError(
"`sigmas` and `timesteps` should have the same length as num_inference_steps, if `num_inference_steps` is provided"
)
else:
num_inference_steps = len(sigmas) if sigmas is not None else len(timesteps)
self.num_inference_steps = num_inference_steps
# 1. Prepare default sigmas
is_timesteps_provided = timesteps is not None
if is_timesteps_provided:
timesteps = np.array(timesteps).astype(np.float32)
if sigmas is None:
if timesteps is None:
timesteps = np.linspace(
self._sigma_to_t(self.sigma_max), self._sigma_to_t(self.sigma_min), num_inference_steps
)
sigmas = timesteps / self.config.num_train_timesteps
else:
sigmas = np.array(sigmas).astype(np.float32)
num_inference_steps = len(sigmas)
# 2. Perform timestep shifting. Either no shifting is applied, or resolution-dependent shifting of
# "exponential" or "linear" type is applied
if self.config.use_dynamic_shifting:
sigmas = self.time_shift(mu, 1.0, sigmas)
else:
sigmas = self.shift * sigmas / (1 + (self.shift - 1) * sigmas)
# 3. If required, stretch the sigmas schedule to terminate at the configured `shift_terminal` value
if self.config.shift_terminal:
sigmas = self.stretch_shift_to_terminal(sigmas)
# 4. If required, convert sigmas to one of karras, exponential, or beta sigma schedules
if self.config.use_karras_sigmas:
sigmas = self._convert_to_karras(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
elif self.config.use_exponential_sigmas:
sigmas = self._convert_to_exponential(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
elif self.config.use_beta_sigmas:
sigmas = self._convert_to_beta(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
# 5. Convert sigmas and timesteps to tensors and move to specified device
sigmas = torch.from_numpy(sigmas).to(dtype=torch.float32, device=device)
if not is_timesteps_provided:
timesteps = sigmas * self.config.num_train_timesteps
else:
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32, device=device)
# 6. Append the terminal sigma value.
# If a model requires inverted sigma schedule for denoising but timesteps without inversion, the
# `invert_sigmas` flag can be set to `True`. This case is only required in Mochi
if self.config.invert_sigmas:
sigmas = 1.0 - sigmas
timesteps = sigmas * self.config.num_train_timesteps
sigmas = torch.cat([sigmas, torch.ones(1, device=sigmas.device)])
else:
sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)])
self.timesteps = timesteps
self.sigmas = sigmas
self._step_index = None
self._begin_index = None
def index_for_timestep(self, timestep, schedule_timesteps=None):
if schedule_timesteps is None:
schedule_timesteps = self.timesteps
indices = (schedule_timesteps == timestep).nonzero()
# The sigma index that is taken for the **very** first `step`
# is always the second index (or the last index if there is only 1)
# This way we can ensure we don't accidentally skip a sigma in
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
pos = 1 if len(indices) > 1 else 0
return indices[pos].item()
def _init_step_index(self, timestep):
if self.begin_index is None:
if isinstance(timestep, torch.Tensor):
timestep = timestep.to(self.timesteps.device)
self._step_index = self.index_for_timestep(timestep)
else:
self._step_index = self._begin_index
def step(
self,
model_output: torch.FloatTensor,
timestep: Union[float, torch.FloatTensor],
sample: torch.FloatTensor,
s_churn: float = 0.0,
s_tmin: float = 0.0,
s_tmax: float = float("inf"),
s_noise: float = 1.0,
generator: Optional[torch.Generator] = None,
per_token_timesteps: Optional[torch.Tensor] = None,
return_dict: bool = True,
) -> Union[FlowMatchLCMSchedulerOutput, Tuple]:
"""
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
process from the learned model outputs (most often the predicted noise).
Args:
model_output (`torch.FloatTensor`):
The direct output from learned diffusion model.
timestep (`float`):
The current discrete timestep in the diffusion chain.
sample (`torch.FloatTensor`):
A current instance of a sample created by the diffusion process.
s_churn (`float`):
s_tmin (`float`):
s_tmax (`float`):
s_noise (`float`, defaults to 1.0):
Scaling factor for noise added to the sample.
generator (`torch.Generator`, *optional*):
A random number generator.
per_token_timesteps (`torch.Tensor`, *optional*):
The timesteps for each token in the sample.
return_dict (`bool`):
Whether or not to return a
[`~schedulers.scheduling_flow_match_lcm.FlowMatchLCMSchedulerOutput`] or tuple.
Returns:
[`~schedulers.scheduling_flow_match_lcm.FlowMatchLCMSchedulerOutput`] or `tuple`:
If return_dict is `True`,
[`~schedulers.scheduling_flow_match_lcm.FlowMatchLCMSchedulerOutput`] is returned,
otherwise a tuple is returned where the first element is the sample tensor.
"""
if (
isinstance(timestep, int)
or isinstance(timestep, torch.IntTensor)
or isinstance(timestep, torch.LongTensor)
):
raise ValueError(
(
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `FlowMatchLCMScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."
),
)
if (
self._scale_factors
and self._upscale_mode
and len(self.timesteps) != len(self._scale_factors) + 1
):
raise ValueError(
"`_scale_factors` should have the same length as `timesteps` - 1, if `_scale_factors` are set."
)
if self._init_size is None or self.step_index is None:
self._init_size = model_output.size()[2:]
if self.step_index is None:
self._init_step_index(timestep)
# Upcast to avoid precision issues when computing prev_sample
sample = sample.to(torch.float32)
sigma = self.sigmas[self.step_index]
sigma_next = self.sigmas[self.step_index + 1]
x0_pred = (sample - sigma * model_output)
if self._scale_factors and self._upscale_mode:
if self._step_index < len(self._scale_factors):
size = [
round(self._scale_factors[self._step_index] * size)
for size in self._init_size
]
x0_pred = torch.nn.functional.interpolate(
x0_pred,
size=size,
mode=self._upscale_mode
)
noise = randn_tensor(
x0_pred.shape, generator=generator, device=x0_pred.device, dtype=x0_pred.dtype
)
prev_sample = (1 - sigma_next) * x0_pred + sigma_next * noise
# upon completion increase step index by one
self._step_index += 1
if per_token_timesteps is None:
# Cast sample back to model compatible dtype
prev_sample = prev_sample.to(model_output.dtype)
if not return_dict:
return (prev_sample,)
return FlowMatchLCMSchedulerOutput(prev_sample=prev_sample)
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_karras
def _convert_to_karras(self, in_sigmas: torch.Tensor, num_inference_steps) -> torch.Tensor:
"""Constructs the noise schedule of Karras et al. (2022)."""
# Hack to make sure that other schedulers which copy this function don't break
# TODO: Add this logic to the other schedulers
if hasattr(self.config, "sigma_min"):
sigma_min = self.config.sigma_min
else:
sigma_min = None
if hasattr(self.config, "sigma_max"):
sigma_max = self.config.sigma_max
else:
sigma_max = None
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
rho = 7.0 # 7.0 is the value used in the paper
ramp = np.linspace(0, 1, num_inference_steps)
min_inv_rho = sigma_min ** (1 / rho)
max_inv_rho = sigma_max ** (1 / rho)
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho
return sigmas
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_exponential
def _convert_to_exponential(self, in_sigmas: torch.Tensor, num_inference_steps: int) -> torch.Tensor:
"""Constructs an exponential noise schedule."""
# Hack to make sure that other schedulers which copy this function don't break
# TODO: Add this logic to the other schedulers
if hasattr(self.config, "sigma_min"):
sigma_min = self.config.sigma_min
else:
sigma_min = None
if hasattr(self.config, "sigma_max"):
sigma_max = self.config.sigma_max
else:
sigma_max = None
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
sigmas = np.exp(np.linspace(math.log(sigma_max), math.log(sigma_min), num_inference_steps))
return sigmas
# Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_beta
def _convert_to_beta(
self, in_sigmas: torch.Tensor, num_inference_steps: int, alpha: float = 0.6, beta: float = 0.6
) -> torch.Tensor:
"""From "Beta Sampling is All You Need" [arXiv:2407.12173] (Lee et. al, 2024)"""
# Hack to make sure that other schedulers which copy this function don't break
# TODO: Add this logic to the other schedulers
if hasattr(self.config, "sigma_min"):
sigma_min = self.config.sigma_min
else:
sigma_min = None
if hasattr(self.config, "sigma_max"):
sigma_max = self.config.sigma_max
else:
sigma_max = None
sigma_min = sigma_min if sigma_min is not None else in_sigmas[-1].item()
sigma_max = sigma_max if sigma_max is not None else in_sigmas[0].item()
sigmas = np.array(
[
sigma_min + (ppf * (sigma_max - sigma_min))
for ppf in [
scipy.stats.beta.ppf(timestep, alpha, beta)
for timestep in 1 - np.linspace(0, 1, num_inference_steps)
]
]
)
return sigmas
def _time_shift_exponential(self, mu, sigma, t):
return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma)
def _time_shift_linear(self, mu, sigma, t):
return mu / (mu + (1 / t - 1) ** sigma)
def __len__(self):
return self.config.num_train_timesteps