Compare commits
142
Commits
syncammaster
...
causvid
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
af696b02f3 | ||
|
|
307437e495 | ||
|
|
985268bb69 | ||
|
|
e5bbb95804 | ||
|
|
8410a6977b | ||
|
|
ced1ddaa1a | ||
|
|
5917f51837 | ||
|
|
6eddec54a6 | ||
|
|
9b380ec3c0 | ||
|
|
39412cf422 | ||
|
|
820ac3008e | ||
|
|
d3a09a1a6b | ||
|
|
87a980f6bf | ||
|
|
a290f75b5e | ||
|
|
9c27705dac | ||
|
|
efb87445d5 | ||
|
|
6139017535 | ||
|
|
45dca5c41c | ||
|
|
bdd3828596 | ||
|
|
821881429b | ||
|
|
bd6c853051 | ||
|
|
f4a1157f71 | ||
|
|
e2f42b5773 | ||
|
|
9e8978731b | ||
|
|
7d63a771d8 | ||
|
|
da3c7def22 | ||
|
|
abf61ad833 | ||
|
|
774b055452 | ||
|
|
8bc74daad9 | ||
|
|
cd2884d88a | ||
|
|
c4a8f6d835 | ||
|
|
5c676d3a49 | ||
|
|
ef40577b70 | ||
|
|
fb73022a06 | ||
|
|
0323a9d2c7 | ||
|
|
129f368380 | ||
|
|
87ae18e203 | ||
|
|
07c7fc6c2a | ||
|
|
81e7022ab5 | ||
|
|
934dc127a8 | ||
|
|
ea01ae9ed9 | ||
|
|
0f7de6bac2 | ||
|
|
23e3368de3 | ||
|
|
ece2917a41 | ||
|
|
ef1ed29178 | ||
|
|
d9ca90c1f0 | ||
|
|
75de496d6a | ||
|
|
fe7a3d5b46 | ||
|
|
370837233a | ||
|
|
7875031efe | ||
|
|
fa7015315a | ||
|
|
d547482ad6 | ||
|
|
fd562e9730 | ||
|
|
7bba1e8e17 | ||
|
|
1c3060f59f | ||
|
|
4494bda603 | ||
|
|
ec9f310486 | ||
|
|
6aec1d97a8 | ||
|
|
555c67a8cd | ||
|
|
3cb4902ef1 | ||
|
|
1e6a780729 | ||
|
|
5aed18d658 | ||
|
|
1701375214 | ||
|
|
abdac5ec3a | ||
|
|
bfed2afc2e | ||
|
|
ab870a198f | ||
|
|
f2bc29b931 | ||
|
|
5f7164788e | ||
|
|
03edd8cbf2 | ||
|
|
77f78ceded | ||
|
|
4a41fb0aaf | ||
|
|
b78505014c | ||
|
|
e8074288df | ||
|
|
e3afc7fc75 | ||
|
|
1b3592f6b4 | ||
|
|
ed068e1fe3 | ||
|
|
fe25d21fba | ||
|
|
2d0e0fe698 | ||
|
|
ca0974d5f0 | ||
|
|
a9ec21bb8f | ||
|
|
a7fa683c2c | ||
|
|
1f6fd1e5a8 | ||
|
|
e334f78124 | ||
|
|
c8a084d27a | ||
|
|
8117c6f033 | ||
|
|
d96787b312 | ||
|
|
4e562b5d92 | ||
|
|
5683d8306a | ||
|
|
df95c85283 | ||
|
|
4fc4159c88 | ||
|
|
fc7ab666a2 | ||
|
|
caebbafab8 | ||
|
|
03eeef32a3 | ||
|
|
c09b981d1e | ||
|
|
e75e6d6933 | ||
|
|
031d8ac817 | ||
|
|
33b8f1c082 | ||
|
|
863829d083 | ||
|
|
c7b2635f01 | ||
|
|
96a2172e13 | ||
|
|
b3f3b0cd93 | ||
|
|
4602b9c885 | ||
|
|
e0e5fcf713 | ||
|
|
2b0bf44994 | ||
|
|
24876cffc1 | ||
|
|
1f53574387 | ||
|
|
6fadcbd957 | ||
|
|
a623f87dca | ||
|
|
6099ad393b | ||
|
|
949c887e0c | ||
|
|
5109a74839 | ||
|
|
6ba7bc811e | ||
|
|
f93e54d21b | ||
|
|
3ae0bc1fec | ||
|
|
e3ea2bf392 | ||
|
|
884069121f | ||
|
|
f371c0a736 | ||
|
|
65f5505fca | ||
|
|
e5a326c981 | ||
|
|
7d005201a2 | ||
|
|
604f0e2714 | ||
|
|
04cc8fa80e | ||
|
|
18aa47cc74 | ||
|
|
20088d0fbf | ||
|
|
3692f3580a | ||
|
|
52d6f00770 | ||
|
|
b5b4e44512 | ||
|
|
7b6c34e26d | ||
|
|
b9b8ec7bf8 | ||
|
|
2cfae9fa8f | ||
|
|
90c9415d31 | ||
|
|
421e375f13 | ||
|
|
19044adc78 | ||
|
|
3813c615d3 | ||
|
|
6d521a9f4e | ||
|
|
7462356743 | ||
|
|
00de4f5e0e | ||
|
|
cc8450b1a7 | ||
|
|
8257cd1f8a | ||
|
|
761d188dba | ||
|
|
425fcfcf2c | ||
|
|
7c81d7ce10 |
+2
-1
@@ -9,4 +9,5 @@ logs/
|
|||||||
.idea
|
.idea
|
||||||
tools/
|
tools/
|
||||||
.vscode/
|
.vscode/
|
||||||
convert_*
|
convert_*
|
||||||
|
*.pt
|
||||||
@@ -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)
|
||||||
@@ -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
@@ -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
@@ -1,7 +1,35 @@
|
|||||||
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
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 .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_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"]
|
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||||
@@ -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",
|
||||||
|
}
|
||||||
@@ -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",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -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)
|
||||||
|
|
||||||
@@ -75,7 +75,7 @@ def enable_vram_management_recursively(model: torch.nn.Module, module_map: dict,
|
|||||||
for name, module in model.named_children():
|
for name, module in model.named_children():
|
||||||
for source_module, target_module in module_map.items():
|
for source_module, target_module in module_map.items():
|
||||||
if isinstance(module, source_module):
|
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
|
continue
|
||||||
|
|
||||||
num_param = sum(p.numel() for p in module.parameters())
|
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",
|
"id": "c6e410bc-5e2c-460b-ae81-c91b6094fbb1",
|
||||||
"revision": 0,
|
"revision": 0,
|
||||||
"last_node_id": 204,
|
"last_node_id": 206,
|
||||||
"last_link_id": 336,
|
"last_link_id": 341,
|
||||||
"nodes": [
|
"nodes": [
|
||||||
{
|
{
|
||||||
"id": 42,
|
"id": 42,
|
||||||
@@ -101,7 +101,7 @@
|
|||||||
200
|
200
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 21,
|
"order": 22,
|
||||||
"mode": 2,
|
"mode": 2,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
@@ -258,7 +258,7 @@
|
|||||||
174
|
174
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 34,
|
"order": 39,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
@@ -371,7 +371,7 @@
|
|||||||
200
|
200
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 22,
|
"order": 23,
|
||||||
"mode": 2,
|
"mode": 2,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
@@ -413,7 +413,7 @@
|
|||||||
86
|
86
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 24,
|
"order": 25,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
@@ -647,9 +647,7 @@
|
|||||||
"ver": "0.3.27",
|
"ver": "0.3.27",
|
||||||
"Node name for S&R": "PreviewImage"
|
"Node name for S&R": "PreviewImage"
|
||||||
},
|
},
|
||||||
"widgets_values": [
|
"widgets_values": []
|
||||||
""
|
|
||||||
]
|
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"id": 125,
|
"id": 125,
|
||||||
@@ -663,15 +661,12 @@
|
|||||||
190.28567504882812
|
190.28567504882812
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 40,
|
"order": 41,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
"name": "text",
|
"name": "text",
|
||||||
"type": "STRING",
|
"type": "STRING",
|
||||||
"widget": {
|
|
||||||
"name": "text"
|
|
||||||
},
|
|
||||||
"link": 215
|
"link": 215
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
@@ -689,7 +684,7 @@
|
|||||||
"Node name for S&R": "ShowText|pysssss"
|
"Node name for S&R": "ShowText|pysssss"
|
||||||
},
|
},
|
||||||
"widgets_values": [
|
"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."
|
"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
|
261.5306701660156
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 41,
|
"order": 42,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
@@ -842,7 +837,7 @@
|
|||||||
"flags": {
|
"flags": {
|
||||||
"collapsed": true
|
"collapsed": true
|
||||||
},
|
},
|
||||||
"order": 23,
|
"order": 24,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
@@ -914,7 +909,7 @@
|
|||||||
266
|
266
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 28,
|
"order": 29,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
@@ -922,28 +917,22 @@
|
|||||||
"type": "IMAGE",
|
"type": "IMAGE",
|
||||||
"link": 244
|
"link": 244
|
||||||
},
|
},
|
||||||
{
|
|
||||||
"name": "get_image_size",
|
|
||||||
"shape": 7,
|
|
||||||
"type": "IMAGE",
|
|
||||||
"link": null
|
|
||||||
},
|
|
||||||
{
|
{
|
||||||
"name": "width_input",
|
"name": "width_input",
|
||||||
"shape": 7,
|
"shape": 7,
|
||||||
"type": "INT",
|
"type": "INT",
|
||||||
"widget": {
|
|
||||||
"name": "width_input"
|
|
||||||
},
|
|
||||||
"link": null
|
"link": null
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"name": "height_input",
|
"name": "height_input",
|
||||||
"shape": 7,
|
"shape": 7,
|
||||||
"type": "INT",
|
"type": "INT",
|
||||||
"widget": {
|
"link": null
|
||||||
"name": "height_input"
|
},
|
||||||
},
|
{
|
||||||
|
"name": "get_image_size",
|
||||||
|
"shape": 7,
|
||||||
|
"type": "IMAGE",
|
||||||
"link": null
|
"link": null
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
@@ -978,8 +967,6 @@
|
|||||||
"lanczos",
|
"lanczos",
|
||||||
false,
|
false,
|
||||||
16,
|
16,
|
||||||
0,
|
|
||||||
0,
|
|
||||||
"center"
|
"center"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
@@ -1025,50 +1012,6 @@
|
|||||||
"color": "#2a363b",
|
"color": "#2a363b",
|
||||||
"bgcolor": "#3f5159"
|
"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,
|
"id": 58,
|
||||||
"type": "WanVideoEncode",
|
"type": "WanVideoEncode",
|
||||||
@@ -1139,7 +1082,7 @@
|
|||||||
"flags": {
|
"flags": {
|
||||||
"collapsed": true
|
"collapsed": true
|
||||||
},
|
},
|
||||||
"order": 42,
|
"order": 43,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
@@ -1177,7 +1120,7 @@
|
|||||||
555.8994140625
|
555.8994140625
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 35,
|
"order": 40,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
@@ -1192,9 +1135,7 @@
|
|||||||
"ver": "0.3.27",
|
"ver": "0.3.27",
|
||||||
"Node name for S&R": "PreviewImage"
|
"Node name for S&R": "PreviewImage"
|
||||||
},
|
},
|
||||||
"widgets_values": [
|
"widgets_values": []
|
||||||
""
|
|
||||||
]
|
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
"id": 128,
|
"id": 128,
|
||||||
@@ -1292,7 +1233,7 @@
|
|||||||
274
|
274
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 39,
|
"order": 44,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
@@ -1304,9 +1245,6 @@
|
|||||||
"name": "caption",
|
"name": "caption",
|
||||||
"shape": 7,
|
"shape": 7,
|
||||||
"type": "STRING",
|
"type": "STRING",
|
||||||
"widget": {
|
|
||||||
"name": "caption"
|
|
||||||
},
|
|
||||||
"link": null
|
"link": null
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -1341,8 +1279,7 @@
|
|||||||
"black",
|
"black",
|
||||||
"FreeMonoBoldOblique.otf",
|
"FreeMonoBoldOblique.otf",
|
||||||
"input",
|
"input",
|
||||||
"up",
|
"up"
|
||||||
""
|
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -1353,7 +1290,7 @@
|
|||||||
-1155.6121826171875
|
-1155.6121826171875
|
||||||
],
|
],
|
||||||
"size": [
|
"size": [
|
||||||
390.5999755859375,
|
421.6000061035156,
|
||||||
202
|
202
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
@@ -1383,82 +1320,6 @@
|
|||||||
128
|
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,
|
"id": 127,
|
||||||
"type": "WanVideoExperimentalArgs",
|
"type": "WanVideoExperimentalArgs",
|
||||||
@@ -1471,7 +1332,7 @@
|
|||||||
130
|
130
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 15,
|
"order": 14,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [],
|
"inputs": [],
|
||||||
"outputs": [
|
"outputs": [
|
||||||
@@ -1509,7 +1370,7 @@
|
|||||||
"flags": {
|
"flags": {
|
||||||
"collapsed": true
|
"collapsed": true
|
||||||
},
|
},
|
||||||
"order": 16,
|
"order": 15,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [],
|
"inputs": [],
|
||||||
"outputs": [
|
"outputs": [
|
||||||
@@ -1543,7 +1404,7 @@
|
|||||||
"flags": {
|
"flags": {
|
||||||
"collapsed": true
|
"collapsed": true
|
||||||
},
|
},
|
||||||
"order": 17,
|
"order": 16,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [],
|
"inputs": [],
|
||||||
"outputs": [
|
"outputs": [
|
||||||
@@ -1575,7 +1436,7 @@
|
|||||||
46
|
46
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 27,
|
"order": 28,
|
||||||
"mode": 2,
|
"mode": 2,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
@@ -1617,7 +1478,7 @@
|
|||||||
"flags": {
|
"flags": {
|
||||||
"collapsed": true
|
"collapsed": true
|
||||||
},
|
},
|
||||||
"order": 44,
|
"order": 45,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
@@ -1657,7 +1518,7 @@
|
|||||||
"flags": {
|
"flags": {
|
||||||
"collapsed": true
|
"collapsed": true
|
||||||
},
|
},
|
||||||
"order": 18,
|
"order": 17,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [],
|
"inputs": [],
|
||||||
"outputs": [
|
"outputs": [
|
||||||
@@ -1689,7 +1550,7 @@
|
|||||||
688.150634765625
|
688.150634765625
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 30,
|
"order": 34,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
@@ -1754,6 +1615,12 @@
|
|||||||
"shape": 7,
|
"shape": 7,
|
||||||
"type": "EXPERIMENTALARGS",
|
"type": "EXPERIMENTALARGS",
|
||||||
"link": 334
|
"link": 334
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "sigmas",
|
||||||
|
"shape": 7,
|
||||||
|
"type": "SIGMAS",
|
||||||
|
"link": null
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"outputs": [
|
"outputs": [
|
||||||
@@ -1781,8 +1648,7 @@
|
|||||||
0,
|
0,
|
||||||
1,
|
1,
|
||||||
false,
|
false,
|
||||||
"comfy",
|
"comfy"
|
||||||
""
|
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -1797,7 +1663,7 @@
|
|||||||
178
|
178
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 19,
|
"order": 18,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [],
|
"inputs": [],
|
||||||
"outputs": [
|
"outputs": [
|
||||||
@@ -1832,10 +1698,10 @@
|
|||||||
],
|
],
|
||||||
"size": [
|
"size": [
|
||||||
908.9017944335938,
|
908.9017944335938,
|
||||||
912.1107788085938
|
334
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 43,
|
"order": 46,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [
|
"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,
|
"id": 22,
|
||||||
"type": "WanVideoModelLoader",
|
"type": "WanVideoModelLoader",
|
||||||
@@ -1959,7 +1778,7 @@
|
|||||||
234
|
234
|
||||||
],
|
],
|
||||||
"flags": {},
|
"flags": {},
|
||||||
"order": 20,
|
"order": 19,
|
||||||
"mode": 0,
|
"mode": 0,
|
||||||
"inputs": [
|
"inputs": [
|
||||||
{
|
{
|
||||||
@@ -2017,6 +1836,248 @@
|
|||||||
],
|
],
|
||||||
"color": "#223",
|
"color": "#223",
|
||||||
"bgcolor": "#335"
|
"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": [
|
"links": [
|
||||||
@@ -2060,14 +2121,6 @@
|
|||||||
1,
|
1,
|
||||||
"CONDITIONING"
|
"CONDITIONING"
|
||||||
],
|
],
|
||||||
[
|
|
||||||
102,
|
|
||||||
56,
|
|
||||||
0,
|
|
||||||
74,
|
|
||||||
0,
|
|
||||||
"*"
|
|
||||||
],
|
|
||||||
[
|
[
|
||||||
210,
|
210,
|
||||||
28,
|
28,
|
||||||
@@ -2257,7 +2310,7 @@
|
|||||||
157,
|
157,
|
||||||
0,
|
0,
|
||||||
56,
|
56,
|
||||||
0,
|
1,
|
||||||
"LATENT"
|
"LATENT"
|
||||||
],
|
],
|
||||||
[
|
[
|
||||||
@@ -2291,6 +2344,30 @@
|
|||||||
155,
|
155,
|
||||||
6,
|
6,
|
||||||
"TEACACHEARGS"
|
"TEACACHEARGS"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
338,
|
||||||
|
157,
|
||||||
|
0,
|
||||||
|
205,
|
||||||
|
0,
|
||||||
|
"LATENT"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
340,
|
||||||
|
205,
|
||||||
|
0,
|
||||||
|
56,
|
||||||
|
0,
|
||||||
|
"CAMERAPOSES"
|
||||||
|
],
|
||||||
|
[
|
||||||
|
341,
|
||||||
|
205,
|
||||||
|
0,
|
||||||
|
74,
|
||||||
|
0,
|
||||||
|
"*"
|
||||||
]
|
]
|
||||||
],
|
],
|
||||||
"groups": [
|
"groups": [
|
||||||
@@ -2337,13 +2414,12 @@
|
|||||||
"config": {},
|
"config": {},
|
||||||
"extra": {
|
"extra": {
|
||||||
"ds": {
|
"ds": {
|
||||||
"scale": 0.7400249944258357,
|
"scale": 0.611590904484162,
|
||||||
"offset": [
|
"offset": [
|
||||||
1311.6629502036258,
|
1176.5764579377562,
|
||||||
1366.2150672288358
|
1095.9393193240473
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
"linkExtensions": [],
|
|
||||||
"node_versions": {
|
"node_versions": {
|
||||||
"ComfyUI-WanVideoWrapper": "5a2383621a05825d0d0437781afcb8552d9590fd",
|
"ComfyUI-WanVideoWrapper": "5a2383621a05825d0d0437781afcb8552d9590fd",
|
||||||
"comfy-core": "0.3.26",
|
"comfy-core": "0.3.26",
|
||||||
@@ -2352,7 +2428,8 @@
|
|||||||
"VHS_latentpreview": true,
|
"VHS_latentpreview": true,
|
||||||
"VHS_latentpreviewrate": 0,
|
"VHS_latentpreviewrate": 0,
|
||||||
"VHS_MetadataImage": true,
|
"VHS_MetadataImage": true,
|
||||||
"VHS_KeepIntermediate": true
|
"VHS_KeepIntermediate": true,
|
||||||
|
"frontendVersion": "1.16.7"
|
||||||
},
|
},
|
||||||
"version": 0.4
|
"version": 0.4
|
||||||
}
|
}
|
||||||
+2733
-2408
File diff suppressed because it is too large
Load Diff
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large
Load Diff
File diff suppressed because 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
@@ -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
|
||||||
|
)
|
||||||
@@ -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
@@ -7,8 +7,8 @@ def fp8_linear_forward(cls, original_dtype, input):
|
|||||||
weight_dtype = cls.weight.dtype
|
weight_dtype = cls.weight.dtype
|
||||||
if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
|
if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
|
||||||
if len(input.shape) == 3:
|
if len(input.shape) == 3:
|
||||||
target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn
|
#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)
|
inn = input.reshape(-1, input.shape[2]).to(weight_dtype)
|
||||||
w = cls.weight.t()
|
w = cls.weight.t()
|
||||||
|
|
||||||
scale = torch.ones((1), device=input.device, dtype=torch.float32)
|
scale = torch.ones((1), device=input.device, dtype=torch.float32)
|
||||||
|
|||||||
@@ -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
@@ -84,7 +84,7 @@ def get_previewer(device, latent_format):
|
|||||||
taew_sd = comfy.utils.load_torch_file(taehv_path)
|
taew_sd = comfy.utils.load_torch_file(taehv_path)
|
||||||
taesd = TAEHV(taew_sd).to(device)
|
taesd = TAEHV(taew_sd).to(device)
|
||||||
previewer = TAESDPreviewerImpl(taesd)
|
previewer = TAESDPreviewerImpl(taesd)
|
||||||
previewer = WrappedPreviewer(previewer, rate=16)
|
previewer = WrappedPreviewer(previewer, rate=3)
|
||||||
|
|
||||||
if previewer is None:
|
if previewer is None:
|
||||||
if latent_format.latent_rgb_factors is not 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
|
self.c_index = (self.c_index + num_previews) % num_images
|
||||||
return None
|
return None
|
||||||
def process_previews(self, image_tensor, ind, leng):
|
def process_previews(self, image_tensor, ind, leng):
|
||||||
max_size = 256
|
max_size = 512
|
||||||
image_tensor = self.decode_latent_to_preview(image_tensor)
|
image_tensor = self.decode_latent_to_preview(image_tensor)
|
||||||
if image_tensor.size(1) > max_size or image_tensor.size(2) > max_size:
|
if image_tensor.size(1) > max_size or image_tensor.size(2) > max_size:
|
||||||
image_tensor = image_tensor.movedim(-1,0)
|
image_tensor = image_tensor.movedim(-1,0)
|
||||||
if image_tensor.size(2) < image_tensor.size(3):
|
if image_tensor.size(2) < image_tensor.size(3):
|
||||||
height = (max_size * 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:
|
else:
|
||||||
width = (max_size * image_tensor.size(3)) // image_tensor.size(2)
|
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.movedim(0,-1)
|
||||||
|
#image_tensor = image_tensor.repeat_interleave(2, dim=0)
|
||||||
|
|
||||||
previews_ubyte = (image_tensor.clamp(0, 1)
|
previews_ubyte = (image_tensor.clamp(0, 1)
|
||||||
.mul(0xFF) # to 0..255
|
.mul(0xFF) # to 0..255
|
||||||
).to(device="cpu", dtype=torch.uint8)
|
).to(device="cpu", dtype=torch.uint8)
|
||||||
@@ -189,10 +191,10 @@ class WrappedPreviewer(LatentPreviewer):
|
|||||||
#NOTE: send sync already uses call_soon_threadsafe
|
#NOTE: send sync already uses call_soon_threadsafe
|
||||||
serv.send_sync(server.BinaryEventTypes.PREVIEW_IMAGE,
|
serv.send_sync(server.BinaryEventTypes.PREVIEW_IMAGE,
|
||||||
message.getvalue(), serv.client_id)
|
message.getvalue(), serv.client_id)
|
||||||
if self.rate == 16:
|
#if self.rate == 16:
|
||||||
ind = (ind + 1) % ((leng-1) * 4 - 1)
|
ind = (ind + 1) % ((leng-1) * 4 - 1)
|
||||||
else:
|
#else:
|
||||||
ind = (ind + 1) % leng
|
# ind = (ind + 1) % leng
|
||||||
|
|
||||||
# Send SwarmUI preview if detected
|
# Send SwarmUI preview if detected
|
||||||
if self.swarmui_env:
|
if self.swarmui_env:
|
||||||
|
|||||||
+2
-2
@@ -1,7 +1,7 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "ComfyUI-WanVideoWrapper"
|
name = "ComfyUI-WanVideoWrapper"
|
||||||
description = "ComfyUI diffusers wrapper nodes for WanVideo"
|
description = "ComfyUI wrapper nodes for WanVideo"
|
||||||
version = "1.1.5"
|
version = "1.1.9"
|
||||||
license = {file = "LICENSE"}
|
license = {file = "LICENSE"}
|
||||||
dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.32.0", "ftfy"]
|
dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.32.0", "ftfy"]
|
||||||
|
|
||||||
|
|||||||
+105
-18
@@ -10,9 +10,16 @@ class Camera(object):
|
|||||||
c2w_mat = np.array(c2w).reshape(4, 4)
|
c2w_mat = np.array(c2w).reshape(4, 4)
|
||||||
self.c2w_mat = c2w_mat
|
self.c2w_mat = c2w_mat
|
||||||
self.w2c_mat = np.linalg.inv(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
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {"required": {
|
return {"required": {
|
||||||
@@ -32,18 +39,16 @@ class WanVideoReCamMasterCameraEmbed:
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "CAMERAPOSES",)
|
RETURN_TYPES = ("CAMERAPOSES",)
|
||||||
RETURN_NAMES = ("camera_embeds", "camera_poses",)
|
RETURN_NAMES = ("camera_poses",)
|
||||||
FUNCTION = "process"
|
FUNCTION = "process"
|
||||||
CATEGORY = "WanVideoWrapper"
|
CATEGORY = "WanVideoWrapper"
|
||||||
DESCRIPTION = "https://github.com/KwaiVGI/ReCamMaster"
|
DESCRIPTION = "https://github.com/KwaiVGI/ReCamMaster"
|
||||||
|
|
||||||
def process(self, camera_type, latents):
|
def process(self, camera_type, latents):
|
||||||
# load camera
|
|
||||||
import json
|
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:
|
with open(camera_data_path, 'r') as file:
|
||||||
cam_data = json.load(file)
|
cam_data = json.load(file)
|
||||||
|
|
||||||
@@ -65,10 +70,96 @@ class WanVideoReCamMasterCameraEmbed:
|
|||||||
}
|
}
|
||||||
|
|
||||||
cam_idx = list(range(num_frames))[::4]
|
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)
|
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 = []
|
c2ws = []
|
||||||
for c2w in traj:
|
for c2w in camera_poses:
|
||||||
c2w = c2w[:, [1, 2, 0, 3]]
|
c2w = c2w[:, [1, 2, 0, 3]]
|
||||||
c2w[:3, 1] *= -1.
|
c2w[:3, 1] *= -1.
|
||||||
c2w[:3, 3] /= 100
|
c2w[:3, 3] /= 100
|
||||||
@@ -93,15 +184,7 @@ class WanVideoReCamMasterCameraEmbed:
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return (embeds, traj,)
|
return (embeds, camera_poses,)
|
||||||
|
|
||||||
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)
|
|
||||||
|
|
||||||
def get_relative_pose(self, cam_params):
|
def get_relative_pose(self, cam_params):
|
||||||
abs_w2cs = [cam_param.w2c_mat for cam_param in 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 = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"WanVideoReCamMasterCameraEmbed": WanVideoReCamMasterCameraEmbed,
|
"WanVideoReCamMasterCameraEmbed": WanVideoReCamMasterCameraEmbed,
|
||||||
"ReCamMasterPoseVisualizer": ReCamMasterPoseVisualizer,
|
"ReCamMasterPoseVisualizer": ReCamMasterPoseVisualizer,
|
||||||
|
"WanVideoReCamMasterGenerateOrbitCamera": WanVideoReCamMasterGenerateOrbitCamera,
|
||||||
|
"WanVideoReCamMasterDefaultCamera": WanVideoReCamMasterDefaultCamera,
|
||||||
}
|
}
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
"WanVideoReCamMasterCameraEmbed": "WanVideo ReCamMaster Camera Embed",
|
"WanVideoReCamMasterCameraEmbed": "WanVideo ReCamMaster Camera Embed",
|
||||||
"ReCamMasterPoseVisualizer": "ReCamMaster Pose Visualizer",
|
"ReCamMasterPoseVisualizer": "ReCamMaster Pose Visualizer",
|
||||||
|
"WanVideoReCamMasterGenerateOrbitCamera": "WanVideo ReCamMaster Generate Orbit Camera",
|
||||||
|
"WanVideoReCamMasterDefaultCamera": "WanVideo ReCamMaster Default Camera",
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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",
|
||||||
|
}
|
||||||
@@ -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
|
||||||
@@ -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
@@ -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
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
|
|
||||||
|
|
||||||
@@ -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",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -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
|
name = name.replace("._orig_mod.", ".") # torch compiled modules have this prefix
|
||||||
if low_mem_load:
|
if low_mem_load:
|
||||||
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
|
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
|
||||||
if "modulation" in name:
|
if "patch_embedding" in name:
|
||||||
dtype_to_use = torch.float32
|
dtype_to_use = torch.float32
|
||||||
if name.startswith("diffusion_model."):
|
if name.startswith("diffusion_model."):
|
||||||
name_no_prefix = name[len("diffusion_model."):]
|
name_no_prefix = name[len("diffusion_model."):]
|
||||||
key = "{}.{}".format(name_no_prefix, param)
|
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)
|
model.patch_weight_to_device("{}.{}".format(name, param), device_to=device_to)
|
||||||
if low_mem_load:
|
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
|
m.comfy_patched_weights = True
|
||||||
|
|
||||||
model.current_weight_patches_uuid = model.patches_uuid
|
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():
|
for name, param in model.model.diffusion_model.named_parameters():
|
||||||
if param.device != transformer_load_device:
|
if param.device != transformer_load_device:
|
||||||
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
|
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
|
||||||
if "modulation" in name:
|
if "patch_embedding" in name:
|
||||||
dtype_to_use = torch.float32
|
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
|
return model
|
||||||
|
|
||||||
|
|
||||||
@@ -164,4 +172,57 @@ def encode_image_(clip_vision, image):
|
|||||||
pixel_values = clip_preprocess(image, size=224, crop=True).float()
|
pixel_values = clip_preprocess(image, size=224, crop=True).float()
|
||||||
out = clip_vision.visual(pixel_values)
|
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
|
||||||
@@ -17,7 +17,10 @@ try:
|
|||||||
from sageattention import sageattn
|
from sageattention import sageattn
|
||||||
@torch.compiler.disable()
|
@torch.compiler.disable()
|
||||||
def sageattn_func(q, k, v, attn_mask=None, dropout_p=0, is_causal=False):
|
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:
|
except Exception as e:
|
||||||
print(f"Warning: Could not load sageattention: {str(e)}")
|
print(f"Warning: Could not load sageattention: {str(e)}")
|
||||||
if isinstance(e, ModuleNotFoundError):
|
if isinstance(e, ModuleNotFoundError):
|
||||||
@@ -196,9 +199,9 @@ def attention(
|
|||||||
elif attention_mode == 'sageattn':
|
elif attention_mode == 'sageattn':
|
||||||
attn_mask = None
|
attn_mask = None
|
||||||
|
|
||||||
q = q.transpose(1, 2).to(dtype)
|
q = q.transpose(1, 2)
|
||||||
k = k.transpose(1, 2).to(dtype)
|
k = k.transpose(1, 2)
|
||||||
v = v.transpose(1, 2).to(dtype)
|
v = v.transpose(1, 2)
|
||||||
|
|
||||||
out = sageattn_func(
|
out = sageattn_func(
|
||||||
q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p)
|
q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p)
|
||||||
|
|||||||
+629
-85
File diff suppressed because it is too large
Load Diff
@@ -516,4 +516,6 @@ class T5EncoderModel:
|
|||||||
mask = mask.to(device)
|
mask = mask.to(device)
|
||||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||||
context = self.model(ids, mask)
|
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)]
|
return [u[:v] for u, v in zip(context, seq_lens)]
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
Reference in New Issue
Block a user