Squashed commit of the following:
commit 6a01a8a1d80d36b5b8ac979a36069c8eb0c2f9a7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date: Thu Aug 14 15:40:50 2025 +0300
Update wanvideo_2_1_I2V_FantasyPortrait_example_01.json
commit e3cf4bf5bc13321e6d8f91fcb7ee92210a7adf01
Merge: bbf14ec f3d5f6b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date: Thu Aug 14 15:08:53 2025 +0300
Merge branch 'main' into fantasy_portrait
commit bbf14ec9e9965c1a8582eea02b50913e79a036d0
Author: kijai <40791699+kijai@users.noreply.github.com>
Date: Thu Aug 14 02:16:11 2025 +0300
update
commit 8192f9f4302b3933e48641a8ef313929e3e263aa
Author: kijai <40791699+kijai@users.noreply.github.com>
Date: Thu Aug 14 01:46:15 2025 +0300
progress bar, fix context windows
commit 39fab8ad4d950a974ba49bd177814f30f518b478
Author: kijai <40791699+kijai@users.noreply.github.com>
Date: Thu Aug 14 01:29:15 2025 +0300
Update nodes.py
commit 36f472c0134e6342ab8c2062a7e3f0b0c829003b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date: Thu Aug 14 01:14:24 2025 +0300
Add start/end percent
commit 16f5922c6bc575754412c9b473907377562b956c
Author: kijai <40791699+kijai@users.noreply.github.com>
Date: Thu Aug 14 00:58:57 2025 +0300
init
This commit is contained in:
+12
@@ -2,6 +2,7 @@ from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .recammaster.nodes import NODE_CLASS_MAPPINGS as RECAM_MASTER_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .skyreels.nodes import NODE_CLASS_MAPPINGS as SKYREELS_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as SKYREELS_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .fantasytalking.nodes import NODE_CLASS_MAPPINGS as FANTASYTALKING_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FANTASYTALKING_NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
from .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
|
||||
@@ -15,8 +16,17 @@ from .nodes_deprecated import NODE_CLASS_MAPPINGS as DEPRECATED_NODE_CLASS_MAPPI
|
||||
try:
|
||||
from .qwen.qwen import NODE_CLASS_MAPPINGS as QWEN_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as QWEN_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except ImportError:
|
||||
QWEN_NODE_CLASS_MAPPINGS = {}
|
||||
QWEN_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
print("Qwen not available due to missing dependencies, probably transformers")
|
||||
|
||||
try:
|
||||
from .fantasyportrait.nodes import NODE_CLASS_MAPPINGS as FANTASYPORTRAIT_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FANTASYPORTRAIT_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except ImportError:
|
||||
print("FantasyPortrait not available due to missing dependencies, probably safetensors or torch")
|
||||
FANTASYPORTRAIT_NODE_CLASS_MAPPINGS = {}
|
||||
FANTASYPORTRAIT_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
try:
|
||||
from .unianimate.nodes import NODE_CLASS_MAPPINGS as UNIANIMATE_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except ImportError:
|
||||
@@ -28,6 +38,7 @@ NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(FANTASYTALKING_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(FANTASYPORTRAIT_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(FUN_CAMERA_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(UNI3C_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(CONTROLNET_NODE_CLASS_MAPPINGS)
|
||||
@@ -43,6 +54,7 @@ NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(SKYREELS_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(FANTASYTALKING_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(FANTASYPORTRAIT_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(FUN_CAMERA_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(UNI3C_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(CONTROLNET_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,506 @@
|
||||
import math
|
||||
import os.path as osp
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
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, dx0=0.0, d_cutoff=1.0):
|
||||
self.d_cutoff = float(d_cutoff)
|
||||
self.dx_prev = float(dx0)
|
||||
|
||||
def __call__(self, x, x_prev, fcmin=1.0, min_cutoff=1.0, beta=0.0):
|
||||
if x_prev is None:
|
||||
return x
|
||||
# t_e = 1
|
||||
a_d = smoothing_factor(fcmin, self.d_cutoff)
|
||||
dx = (x - x_prev) / fcmin
|
||||
dx_hat = exponential_smoothing(a_d, dx, self.dx_prev)
|
||||
cutoff = min_cutoff + beta * abs(dx_hat)
|
||||
a = smoothing_factor(fcmin, cutoff)
|
||||
x_hat = exponential_smoothing(a, x, x_prev)
|
||||
self.dx_prev = dx_hat
|
||||
return x_hat
|
||||
|
||||
|
||||
def cult_dis(old_kpts, new_kpts):
|
||||
dis = np.sqrt(
|
||||
np.square(new_kpts[:, 0] - old_kpts[:, 0])
|
||||
+ np.square(new_kpts[:, 1] - old_kpts[:, 1])
|
||||
)
|
||||
return dis
|
||||
|
||||
|
||||
class Smoother222(object):
|
||||
def __init__(self):
|
||||
# face config
|
||||
self.face_idx = list(range(0, 33))
|
||||
self.face_down_idx = list(range(9, 24))
|
||||
self.filter_face = OneEuroFilter()
|
||||
# nose config
|
||||
self.nose_idx = list(range(33, 48))
|
||||
self.filter_nose = OneEuroFilter()
|
||||
# eyebrow config
|
||||
self.eyebrow_idx = list(range(48, 74))
|
||||
self.filter_eyebrow = OneEuroFilter()
|
||||
# eye config
|
||||
self.left_eye_idx = list(range(74, 96))
|
||||
self.filter_left_eye = OneEuroFilter()
|
||||
self.right_eye_idx = list(range(96, 118))
|
||||
self.filter_right_eye = OneEuroFilter()
|
||||
# mouth config
|
||||
self.mouth_idx = list(range(118, 182))
|
||||
self.filter_mouth = OneEuroFilter()
|
||||
# pupil config
|
||||
self.left_pupil_idx = list(range(182, 202))
|
||||
self.filter_left_pupil = OneEuroFilter()
|
||||
self.right_pupil_idx = list(range(202, 222))
|
||||
self.filter_right_pupil = OneEuroFilter()
|
||||
self.prev_points = None
|
||||
|
||||
def smooth(self, new_points, face_dis):
|
||||
if self.prev_points is None:
|
||||
self.prev_points = new_points.copy()
|
||||
return new_points
|
||||
dis = cult_dis(self.prev_points, new_points) / face_dis
|
||||
smooth_points = new_points.copy()
|
||||
|
||||
# smooth face
|
||||
if np.mean(dis[self.face_down_idx]) < 0.005:
|
||||
ratio_tmp = np.mean(dis[self.face_down_idx]) / 0.005
|
||||
fcmin_tmp = 0.05 * ratio_tmp
|
||||
beta_tmp = 0.05 * ratio_tmp
|
||||
smooth_points[self.face_idx] = self.filter_face(
|
||||
new_points[self.face_idx],
|
||||
self.prev_points[self.face_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
elif np.mean(dis[self.face_down_idx]) < 0.02:
|
||||
ratio_tmp = (np.mean(dis[self.face_down_idx]) - 0.005) / (0.02 - 0.005)
|
||||
fcmin_tmp = 0.05 + (0.3 - 0.05) * ratio_tmp
|
||||
beta_tmp = 0.05 + (0.3 - 0.05) * ratio_tmp
|
||||
smooth_points[self.face_idx] = self.filter_face(
|
||||
new_points[self.face_idx],
|
||||
self.prev_points[self.face_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
else:
|
||||
smooth_points[self.face_idx] = self.filter_face(
|
||||
new_points[self.face_idx],
|
||||
self.prev_points[self.face_idx],
|
||||
fcmin=0.3,
|
||||
beta=0.3,
|
||||
)
|
||||
# smooth nose
|
||||
if np.mean(dis[self.nose_idx]) < 0.003:
|
||||
# stable
|
||||
ratio_tmp = np.mean(dis[self.nose_idx]) / 0.003
|
||||
fcmin_tmp = 0.03 * ratio_tmp
|
||||
beta_tmp = 0.03 * ratio_tmp
|
||||
smooth_points[self.nose_idx] = self.filter_nose(
|
||||
new_points[self.nose_idx],
|
||||
self.prev_points[self.nose_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
elif np.mean(dis[self.nose_idx]) < 0.02:
|
||||
ratio_tmp = (np.mean(dis[self.nose_idx]) - 0.003) / (0.02 - 0.003)
|
||||
fcmin_tmp = 0.03 + (0.7 - 0.03) * ratio_tmp
|
||||
beta_tmp = 0.03 + (0.7 - 0.03) * ratio_tmp
|
||||
smooth_points[self.nose_idx] = self.filter_nose(
|
||||
new_points[self.nose_idx],
|
||||
self.prev_points[self.nose_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
else:
|
||||
# filter
|
||||
smooth_points[self.nose_idx] = self.filter_nose(
|
||||
new_points[self.nose_idx],
|
||||
self.prev_points[self.nose_idx],
|
||||
fcmin=0.7,
|
||||
beta=0.7,
|
||||
)
|
||||
# smooth eyebrow
|
||||
if np.mean(dis[self.eyebrow_idx]) < 0.003:
|
||||
# stable
|
||||
ratio_tmp = np.mean(dis[self.eyebrow_idx]) / 0.003
|
||||
fcmin_tmp = 0.02 * ratio_tmp
|
||||
beta_tmp = 0.02 * ratio_tmp
|
||||
smooth_points[self.eyebrow_idx] = self.filter_eyebrow(
|
||||
new_points[self.eyebrow_idx],
|
||||
self.prev_points[self.eyebrow_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
elif np.mean(dis[self.eyebrow_idx]) < 0.02:
|
||||
# filter
|
||||
ratio_tmp = (np.mean(dis[self.eyebrow_idx]) - 0.003) / (0.02 - 0.003)
|
||||
fcmin_tmp = 0.02 + (0.5 - 0.02) * ratio_tmp
|
||||
beta_tmp = 0.02 + (0.5 - 0.02) * ratio_tmp
|
||||
smooth_points[self.eyebrow_idx] = self.filter_eyebrow(
|
||||
new_points[self.eyebrow_idx],
|
||||
self.prev_points[self.eyebrow_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
else:
|
||||
# filter
|
||||
smooth_points[self.eyebrow_idx] = self.filter_eyebrow(
|
||||
new_points[self.eyebrow_idx],
|
||||
self.prev_points[self.eyebrow_idx],
|
||||
fcmin=0.5,
|
||||
beta=0.5,
|
||||
)
|
||||
# smooth eye
|
||||
if np.mean(dis[self.left_eye_idx]) < 0.003:
|
||||
# stable
|
||||
ratio_tmp = np.mean(dis[self.left_eye_idx]) / 0.003
|
||||
fcmin_tmp = 0.03 * ratio_tmp
|
||||
beta_tmp = 0.03 * ratio_tmp
|
||||
smooth_points[self.left_eye_idx] = self.filter_left_eye(
|
||||
new_points[self.left_eye_idx],
|
||||
self.prev_points[self.left_eye_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
elif np.mean(dis[self.left_eye_idx]) < 0.02:
|
||||
# filter
|
||||
ratio_tmp = (np.mean(dis[self.left_eye_idx]) - 0.003) / (0.02 - 0.003)
|
||||
fcmin_tmp = 0.03 + (0.7 - 0.03) * ratio_tmp
|
||||
beta_tmp = 0.03 + (0.7 - 0.03) * ratio_tmp
|
||||
smooth_points[self.left_eye_idx] = self.filter_left_eye(
|
||||
new_points[self.left_eye_idx],
|
||||
self.prev_points[self.left_eye_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
else:
|
||||
# fast
|
||||
smooth_points[self.left_eye_idx] = self.filter_left_eye(
|
||||
new_points[self.left_eye_idx],
|
||||
self.prev_points[self.left_eye_idx],
|
||||
fcmin=0.7,
|
||||
beta=0.7,
|
||||
)
|
||||
if np.mean(dis[self.right_eye_idx]) < 0.003:
|
||||
# stable
|
||||
ratio_tmp = np.mean(dis[self.right_eye_idx]) / 0.003
|
||||
fcmin_tmp = 0.03 * ratio_tmp
|
||||
beta_tmp = 0.03 * ratio_tmp
|
||||
smooth_points[self.right_eye_idx] = self.filter_right_eye(
|
||||
new_points[self.right_eye_idx],
|
||||
self.prev_points[self.right_eye_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
elif np.mean(dis[self.right_eye_idx]) < 0.02:
|
||||
# filter
|
||||
ratio_tmp = (np.mean(dis[self.right_eye_idx]) - 0.003) / (0.02 - 0.003)
|
||||
fcmin_tmp = 0.03 + (0.7 - 0.03) * ratio_tmp
|
||||
beta_tmp = 0.03 + (0.7 - 0.03) * ratio_tmp
|
||||
smooth_points[self.right_eye_idx] = self.filter_right_eye(
|
||||
new_points[self.right_eye_idx],
|
||||
self.prev_points[self.right_eye_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
else:
|
||||
# fast
|
||||
smooth_points[self.right_eye_idx] = self.filter_right_eye(
|
||||
new_points[self.right_eye_idx],
|
||||
self.prev_points[self.right_eye_idx],
|
||||
fcmin=0.7,
|
||||
beta=0.7,
|
||||
)
|
||||
|
||||
# smooth mouth
|
||||
if np.mean(dis[self.mouth_idx]) < 0.003:
|
||||
# stable
|
||||
ratio_tmp = np.mean(dis[self.mouth_idx]) / 0.003
|
||||
fcmin_tmp = 0.05 * ratio_tmp
|
||||
beta_tmp = 0.05 * ratio_tmp
|
||||
smooth_points[self.mouth_idx] = self.filter_mouth(
|
||||
new_points[self.mouth_idx],
|
||||
self.prev_points[self.mouth_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
elif np.mean(dis[self.mouth_idx]) < 0.02:
|
||||
# filter
|
||||
ratio_tmp = (np.mean(dis[self.mouth_idx]) - 0.003) / (0.02 - 0.003)
|
||||
fcmin_tmp = 0.05 + (0.7 - 0.05) * ratio_tmp
|
||||
beta_tmp = 0.05 + (0.7 - 0.05) * ratio_tmp
|
||||
smooth_points[self.mouth_idx] = self.filter_mouth(
|
||||
new_points[self.mouth_idx],
|
||||
self.prev_points[self.mouth_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
else:
|
||||
# fast
|
||||
smooth_points[self.mouth_idx] = self.filter_mouth(
|
||||
new_points[self.mouth_idx],
|
||||
self.prev_points[self.mouth_idx],
|
||||
fcmin=0.7,
|
||||
beta=0.7,
|
||||
)
|
||||
|
||||
# smooth pupil
|
||||
if np.mean(dis[self.left_pupil_idx]) < 0.003:
|
||||
# stable
|
||||
ratio_tmp = np.mean(dis[self.left_pupil_idx]) / 0.003
|
||||
fcmin_tmp = 0.03 * ratio_tmp
|
||||
beta_tmp = 0.03 * ratio_tmp
|
||||
smooth_points[self.left_pupil_idx] = self.filter_left_pupil(
|
||||
new_points[self.left_pupil_idx],
|
||||
self.prev_points[self.left_pupil_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
elif np.mean(dis[self.left_pupil_idx]) < 0.02:
|
||||
# filter
|
||||
ratio_tmp = (np.mean(dis[self.left_pupil_idx]) - 0.003) / (0.02 - 0.003)
|
||||
fcmin_tmp = 0.03 + (0.7 - 0.03) * ratio_tmp
|
||||
beta_tmp = 0.03 + (0.7 - 0.03) * ratio_tmp
|
||||
smooth_points[self.left_pupil_idx] = self.filter_left_pupil(
|
||||
new_points[self.left_pupil_idx],
|
||||
self.prev_points[self.left_pupil_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
else:
|
||||
# fast
|
||||
smooth_points[self.left_pupil_idx] = self.filter_left_pupil(
|
||||
new_points[self.left_pupil_idx],
|
||||
self.prev_points[self.left_pupil_idx],
|
||||
fcmin=0.7,
|
||||
beta=0.7,
|
||||
)
|
||||
if np.mean(dis[self.right_pupil_idx]) < 0.003:
|
||||
# stable
|
||||
ratio_tmp = np.mean(dis[self.right_pupil_idx]) / 0.003
|
||||
fcmin_tmp = 0.03 * ratio_tmp
|
||||
beta_tmp = 0.03 * ratio_tmp
|
||||
smooth_points[self.right_pupil_idx] = self.filter_right_pupil(
|
||||
new_points[self.right_pupil_idx],
|
||||
self.prev_points[self.right_pupil_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
elif np.mean(dis[self.right_pupil_idx]) < 0.02:
|
||||
# filter
|
||||
ratio_tmp = (np.mean(dis[self.right_pupil_idx]) - 0.003) / (0.02 - 0.003)
|
||||
fcmin_tmp = 0.03 + (0.7 - 0.03) * ratio_tmp
|
||||
beta_tmp = 0.03 + (0.7 - 0.03) * ratio_tmp
|
||||
smooth_points[self.right_pupil_idx] = self.filter_right_pupil(
|
||||
new_points[self.right_pupil_idx],
|
||||
self.prev_points[self.right_pupil_idx],
|
||||
fcmin=fcmin_tmp,
|
||||
beta=beta_tmp,
|
||||
)
|
||||
else:
|
||||
# fast
|
||||
smooth_points[self.right_pupil_idx] = self.filter_right_pupil(
|
||||
new_points[self.right_pupil_idx],
|
||||
self.prev_points[self.right_pupil_idx],
|
||||
fcmin=0.7,
|
||||
beta=0.7,
|
||||
)
|
||||
|
||||
# update pre points
|
||||
self.prev_points = smooth_points
|
||||
return smooth_points
|
||||
|
||||
|
||||
class CameraDemo(object):
|
||||
def __init__(self, face_alignment_module, reset=False):
|
||||
self.face_alignment_module = face_alignment_module
|
||||
self.face_prob_th = 0.0001
|
||||
self.min_face = 96
|
||||
self.face_image_size = self.face_alignment_module.face_image_size
|
||||
self.trackingFaces = []
|
||||
self.reset = reset
|
||||
|
||||
def reset_track(self):
|
||||
self.trackingFaces = []
|
||||
|
||||
def forward(self, src_image, reset=False, pre_rect=None):
|
||||
|
||||
if self.reset or reset:
|
||||
self.trackingFaces = []
|
||||
|
||||
if len(self.trackingFaces) == 0:
|
||||
if pre_rect is not None:
|
||||
detected_faces = [pre_rect]
|
||||
else:
|
||||
detected_faces, _, _ = self.face_alignment_module.face_detector.detect(
|
||||
src_image
|
||||
)
|
||||
for face_rect in detected_faces:
|
||||
new_tracking_object = {
|
||||
"face_rect": face_rect,
|
||||
"rotate_angle": 0.0,
|
||||
"pre_kpt_222": None,
|
||||
"face_dis": np.sqrt(
|
||||
np.square((face_rect[2] - face_rect[0]))
|
||||
+ np.square((face_rect[3] - face_rect[1]))
|
||||
),
|
||||
"smoother_222": Smoother222(),
|
||||
"prob": 0,
|
||||
}
|
||||
self.trackingFaces.append(new_tracking_object)
|
||||
else:
|
||||
detected_faces, _, _ = self.face_alignment_module.face_detector.detect(
|
||||
src_image
|
||||
)
|
||||
for face_rect in detected_faces:
|
||||
new_tracking_object = {
|
||||
"face_rect": face_rect,
|
||||
"rotate_angle": 0.0,
|
||||
"pre_kpt_222": None,
|
||||
"face_dis": np.sqrt(
|
||||
np.square((face_rect[2] - face_rect[0]))
|
||||
+ np.square((face_rect[3] - face_rect[1]))
|
||||
),
|
||||
"smoother_222": Smoother222(),
|
||||
"prob": 0,
|
||||
}
|
||||
self.trackingFaces.append(new_tracking_object)
|
||||
|
||||
delete_idx_list = []
|
||||
for face_idx, tracking_face in enumerate(self.trackingFaces):
|
||||
if tracking_face["pre_kpt_222"] is not None:
|
||||
result_dict = self.face_alignment_module.forward(
|
||||
src_image, pre_pts=tracking_face["pre_kpt_222"], iterations=3
|
||||
)
|
||||
else:
|
||||
result_dict = self.face_alignment_module.forward(
|
||||
src_image, face_box=tracking_face["face_rect"], iterations=3
|
||||
)
|
||||
|
||||
if result_dict["prob"] < self.face_prob_th:
|
||||
if not face_idx in delete_idx_list:
|
||||
delete_idx_list.append(face_idx)
|
||||
continue
|
||||
|
||||
landmarks_final = tracking_face["smoother_222"].smooth(
|
||||
result_dict["pt222"], tracking_face["face_dis"]
|
||||
)
|
||||
tracking_face["pre_kpt_222"] = landmarks_final
|
||||
|
||||
left_eye_corner = landmarks_final[74]
|
||||
right_eye_corner = landmarks_final[96]
|
||||
|
||||
radian = np.arctan2(
|
||||
right_eye_corner[1] - left_eye_corner[1],
|
||||
right_eye_corner[0] - left_eye_corner[0] + 0.00000001,
|
||||
)
|
||||
rotate_angle = np.rad2deg(radian)
|
||||
face_x_min, face_x_max = np.min(landmarks_final[:, 0]), np.max(
|
||||
landmarks_final[:, 0]
|
||||
)
|
||||
face_y_min, face_y_max = np.min(landmarks_final[:, 1]), np.max(
|
||||
landmarks_final[:, 1]
|
||||
)
|
||||
face_bbox = [face_x_min, face_y_min, face_x_max, face_y_max]
|
||||
face_dis = np.linalg.norm(landmarks_final[0] - landmarks_final[32])
|
||||
|
||||
if (
|
||||
face_x_max - face_x_min < self.min_face
|
||||
or face_y_max - face_y_min < self.min_face
|
||||
):
|
||||
if not face_idx in delete_idx_list:
|
||||
delete_idx_list.append(face_idx)
|
||||
|
||||
euler_pred = result_dict["euler_rad"]
|
||||
pitch = np.rad2deg(euler_pred[0])
|
||||
yaw = np.rad2deg(euler_pred[1])
|
||||
roll = np.rad2deg(euler_pred[2])
|
||||
# print("pitch, yaw, roll", pitch, yaw, roll)
|
||||
|
||||
# one filter model
|
||||
max_euler = abs(pitch) + (abs(yaw) * 0.6)
|
||||
face_dis *= 1.0 + max_euler / 18.0
|
||||
|
||||
# two filter model
|
||||
tracking_face["face_rect"] = face_bbox
|
||||
tracking_face["rotate_angle"] = rotate_angle
|
||||
tracking_face["face_dis"] = face_dis
|
||||
tracking_face["prob"] = result_dict["prob"]
|
||||
tracking_face["pitch"] = pitch
|
||||
tracking_face["yaw"] = yaw
|
||||
tracking_face["roll"] = roll
|
||||
tracking_face["euler_rad"] = result_dict["euler_rad"]
|
||||
|
||||
if len(self.trackingFaces) > 1:
|
||||
for face_idx, tracking_face_target in enumerate(self.trackingFaces):
|
||||
if face_idx in delete_idx_list:
|
||||
continue
|
||||
for idx, tracking_face in enumerate(self.trackingFaces):
|
||||
if idx in delete_idx_list:
|
||||
continue
|
||||
if face_idx == idx:
|
||||
continue
|
||||
iou_temp = self.count_iou(
|
||||
tracking_face_target["face_rect"], tracking_face["face_rect"]
|
||||
)
|
||||
# prog 2
|
||||
if iou_temp > 0.12:
|
||||
if (
|
||||
self.area(tracking_face_target["face_rect"])
|
||||
- self.area(tracking_face["face_rect"])
|
||||
< 0
|
||||
):
|
||||
if not face_idx in delete_idx_list:
|
||||
delete_idx_list.append(face_idx)
|
||||
else:
|
||||
if not idx in delete_idx_list:
|
||||
delete_idx_list.append(idx)
|
||||
|
||||
idx_offset = 0
|
||||
for delete_idx in sorted(delete_idx_list):
|
||||
self.trackingFaces.pop(delete_idx - idx_offset)
|
||||
idx_offset += 1
|
||||
|
||||
return self.trackingFaces
|
||||
|
||||
def count_iou(self, boxA, boxB):
|
||||
# determine the (x, y)-coordinates of the intersection rectangle
|
||||
xA = max(boxA[0], boxB[0])
|
||||
yA = max(boxA[1], boxB[1])
|
||||
xB = min(boxA[2], boxB[2])
|
||||
yB = min(boxA[3], boxB[3])
|
||||
|
||||
# compute the area of intersection rectangle
|
||||
interArea = abs(max((xB - xA, 0)) * max((yB - yA), 0))
|
||||
if interArea == 0:
|
||||
return 0
|
||||
# compute the area of both the prediction and ground-truth
|
||||
# rectangles
|
||||
boxAArea = abs((boxA[2] - boxA[0]) * (boxA[3] - boxA[1]))
|
||||
boxBArea = abs((boxB[2] - boxB[0]) * (boxB[3] - boxB[1]))
|
||||
|
||||
# compute the intersection over union by taking the intersection
|
||||
# area and dividing it by the sum of prediction + ground-truth
|
||||
# areas - the interesection area
|
||||
iou = interArea / float(boxAArea + boxBArea - interArea)
|
||||
|
||||
# return the intersection over union value
|
||||
return iou
|
||||
|
||||
def area(self, bbox):
|
||||
w = bbox[3] - bbox[1]
|
||||
h = bbox[2] - bbox[0]
|
||||
return w * h
|
||||
@@ -0,0 +1,117 @@
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
from .face_det import FaceDet
|
||||
from .face_utils import (create_onnx_session, get_warp_mat_bbox,
|
||||
get_warp_mat_bbox_by_gt_pts_float, transform_points)
|
||||
|
||||
|
||||
class FaceAlignment(object):
|
||||
def __init__(self, gpu_id=None, alignment_model_path="", det_model_path=""):
|
||||
expand_ratio = 0.15
|
||||
|
||||
self.face_alignment_net_222 = create_onnx_session(
|
||||
alignment_model_path, gpu_id=gpu_id
|
||||
)
|
||||
self.onnx_input_name_222 = self.face_alignment_net_222.get_inputs()[0].name
|
||||
self.onnx_output_name_222 = [
|
||||
output.name for output in self.face_alignment_net_222.get_outputs()
|
||||
]
|
||||
self.face_image_size = 128
|
||||
|
||||
self.face_detector = FaceDet(det_model_path, gpu_id=gpu_id)
|
||||
self.expand_ratio = expand_ratio
|
||||
|
||||
def onnx_infer(self, input_uint8):
|
||||
assert input_uint8.shape[0] == input_uint8.shape[1] == self.face_image_size
|
||||
onnx_input = (
|
||||
input_uint8.transpose((2, 0, 1)).astype(np.float32)[np.newaxis, :, :, :]
|
||||
/ 255.0
|
||||
)
|
||||
landmark, euler, prob = self.face_alignment_net_222.run(
|
||||
self.onnx_output_name_222, {self.onnx_input_name_222: onnx_input}
|
||||
)
|
||||
|
||||
landmark = (
|
||||
np.reshape(landmark[0], (2, -1)).transpose((1, 0)) * self.face_image_size
|
||||
)
|
||||
left_eye_corner = landmark[74]
|
||||
right_eye_corner = landmark[96]
|
||||
radian = np.arctan2(
|
||||
right_eye_corner[1] - left_eye_corner[1],
|
||||
right_eye_corner[0] - left_eye_corner[0] + 0.00000001,
|
||||
)
|
||||
euler_rad = np.array([euler[0, 0], euler[0, 1], radian], dtype=np.float32)
|
||||
prob = prob[0]
|
||||
|
||||
return landmark, euler_rad, prob
|
||||
|
||||
def forward(self, src_image, face_box=None, pre_pts=None, iterations=3):
|
||||
if pre_pts is None:
|
||||
if face_box is None:
|
||||
# Detect max size face
|
||||
bounding_boxes, _, score = self.face_detector.detect(src_image)
|
||||
print("facedet score", score)
|
||||
if len(bounding_boxes) == 0:
|
||||
return None
|
||||
bbox = np.zeros(4, dtype=np.float32)
|
||||
if len(bounding_boxes) >= 1:
|
||||
max_area = 0.0
|
||||
for each_bbox in bounding_boxes:
|
||||
area = (each_bbox[2] - each_bbox[0]) * (
|
||||
each_bbox[3] - each_bbox[1]
|
||||
)
|
||||
if area > max_area:
|
||||
bbox[:4] = each_bbox[:4]
|
||||
max_area = area
|
||||
else:
|
||||
bbox = bounding_boxes[0, :4]
|
||||
else:
|
||||
bbox = face_box.copy()
|
||||
M_Face = get_warp_mat_bbox(
|
||||
bbox, 0, self.face_image_size, expand_ratio=self.expand_ratio
|
||||
)
|
||||
else:
|
||||
left_eye_corner = pre_pts[74]
|
||||
right_eye_corner = pre_pts[96]
|
||||
|
||||
radian = np.arctan2(
|
||||
right_eye_corner[1] - left_eye_corner[1],
|
||||
right_eye_corner[0] - left_eye_corner[0] + 0.00000001,
|
||||
)
|
||||
M_Face = get_warp_mat_bbox_by_gt_pts_float(
|
||||
pre_pts,
|
||||
np.rad2deg(radian),
|
||||
self.face_image_size,
|
||||
expand_ratio=self.expand_ratio,
|
||||
)
|
||||
|
||||
face_input = cv2.warpAffine(
|
||||
src_image, M_Face, (self.face_image_size, self.face_image_size)
|
||||
)
|
||||
landmarks, euler, prob = self.onnx_infer(face_input)
|
||||
landmarks = transform_points(landmarks, M_Face, invert=True)
|
||||
|
||||
# Repeat
|
||||
for i in range(iterations - 1):
|
||||
M_Face = get_warp_mat_bbox_by_gt_pts_float(
|
||||
landmarks,
|
||||
np.rad2deg(euler[2]),
|
||||
self.face_image_size,
|
||||
expand_ratio=self.expand_ratio,
|
||||
)
|
||||
face_input = cv2.warpAffine(
|
||||
src_image, M_Face, (self.face_image_size, self.face_image_size)
|
||||
)
|
||||
landmarks, euler, prob = self.onnx_infer(face_input)
|
||||
landmarks = transform_points(landmarks, M_Face, invert=True)
|
||||
|
||||
return_dict = {
|
||||
"pt222": landmarks,
|
||||
"euler_rad": euler,
|
||||
"prob": prob,
|
||||
"M_Face": M_Face,
|
||||
"face_input": face_input,
|
||||
}
|
||||
|
||||
return return_dict
|
||||
@@ -0,0 +1,320 @@
|
||||
import os.path as osp
|
||||
from abc import ABCMeta, abstractmethod
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
from scipy.special import softmax
|
||||
|
||||
from .face_utils import create_onnx_session
|
||||
|
||||
_COLORS = (
|
||||
np.array(
|
||||
[
|
||||
0.000,
|
||||
0.447,
|
||||
0.741,
|
||||
]
|
||||
)
|
||||
.astype(np.float32)
|
||||
.reshape(-1, 3)
|
||||
)
|
||||
|
||||
|
||||
def get_resize_matrix(raw_shape, dst_shape, keep_ratio):
|
||||
"""
|
||||
Get resize matrix for resizing raw img to input size
|
||||
:param raw_shape: (width, height) of raw image
|
||||
:param dst_shape: (width, height) of input image
|
||||
:param keep_ratio: whether keep original ratio
|
||||
:return: 3x3 Matrix
|
||||
"""
|
||||
r_w, r_h = raw_shape
|
||||
d_w, d_h = dst_shape
|
||||
Rs = np.eye(3)
|
||||
if keep_ratio:
|
||||
C = np.eye(3)
|
||||
C[0, 2] = -r_w / 2
|
||||
C[1, 2] = -r_h / 2
|
||||
|
||||
if r_w / r_h < d_w / d_h:
|
||||
ratio = d_h / r_h
|
||||
else:
|
||||
ratio = d_w / r_w
|
||||
Rs[0, 0] *= ratio
|
||||
Rs[1, 1] *= ratio
|
||||
|
||||
T = np.eye(3)
|
||||
T[0, 2] = 0.5 * d_w
|
||||
T[1, 2] = 0.5 * d_h
|
||||
return T @ Rs @ C
|
||||
else:
|
||||
Rs[0, 0] *= d_w / r_w
|
||||
Rs[1, 1] *= d_h / r_h
|
||||
return Rs
|
||||
|
||||
|
||||
def warp_boxes(boxes, M, width, height):
|
||||
"""Apply transform to boxes
|
||||
Copy from nanodet/data/transform/warp.py
|
||||
"""
|
||||
n = len(boxes)
|
||||
if n:
|
||||
# warp points
|
||||
xy = np.ones((n * 4, 3))
|
||||
xy[:, :2] = boxes[:, [0, 1, 2, 3, 0, 3, 2, 1]].reshape(
|
||||
n * 4, 2
|
||||
) # x1y1, x2y2, x1y2, x2y1
|
||||
xy = xy @ M.T # transform
|
||||
xy = (xy[:, :2] / xy[:, 2:3]).reshape(n, 8) # rescale
|
||||
# create new boxes
|
||||
x = xy[:, [0, 2, 4, 6]]
|
||||
y = xy[:, [1, 3, 5, 7]]
|
||||
xy = np.concatenate((x.min(1), y.min(1), x.max(1), y.max(1))).reshape(4, n).T
|
||||
# clip boxes
|
||||
xy[:, [0, 2]] = xy[:, [0, 2]].clip(0, width)
|
||||
xy[:, [1, 3]] = xy[:, [1, 3]].clip(0, height)
|
||||
return xy.astype(np.float32)
|
||||
else:
|
||||
return boxes
|
||||
|
||||
|
||||
def overlay_bbox_cv(img, all_box, class_names):
|
||||
"""Draw result boxes
|
||||
Copy from nanodet/util/visualization.py
|
||||
"""
|
||||
# all_box array of [label, x0, y0, x1, y1, score]
|
||||
all_box.sort(key=lambda v: v[5])
|
||||
for box in all_box:
|
||||
label, x0, y0, x1, y1, score = box
|
||||
# color = self.cmap(i)[:3]
|
||||
color = (_COLORS[label] * 255).astype(np.uint8).tolist()
|
||||
text = "{}:{:.1f}%".format(class_names[label], score * 100)
|
||||
txt_color = (0, 0, 0) if np.mean(_COLORS[label]) > 0.5 else (255, 255, 255)
|
||||
font = cv2.FONT_HERSHEY_SIMPLEX
|
||||
txt_size = cv2.getTextSize(text, font, 0.5, 2)[0]
|
||||
cv2.rectangle(img, (x0, y0), (x1, y1), color, 2)
|
||||
|
||||
cv2.rectangle(
|
||||
img,
|
||||
(x0, y0 - txt_size[1] - 1),
|
||||
(x0 + txt_size[0] + txt_size[1], y0 - 1),
|
||||
color,
|
||||
-1,
|
||||
)
|
||||
cv2.putText(img, text, (x0, y0 - 1), font, 0.5, txt_color, thickness=1)
|
||||
return img
|
||||
|
||||
|
||||
def hard_nms(box_scores, iou_threshold, top_k=-1, candidate_size=200):
|
||||
"""
|
||||
|
||||
Args:
|
||||
box_scores (N, 5): boxes in corner-form and probabilities.
|
||||
iou_threshold: intersection over union threshold.
|
||||
top_k: keep top_k results. If k <= 0, keep all the results.
|
||||
candidate_size: only consider the candidates with the highest scores.
|
||||
Returns:
|
||||
picked: a list of indexes of the kept boxes
|
||||
"""
|
||||
scores = box_scores[:, -1]
|
||||
boxes = box_scores[:, :-1]
|
||||
picked = []
|
||||
# _, indexes = scores.sort(descending=True)
|
||||
indexes = np.argsort(scores)
|
||||
# indexes = indexes[:candidate_size]
|
||||
indexes = indexes[-candidate_size:]
|
||||
while len(indexes) > 0:
|
||||
# current = indexes[0]
|
||||
current = indexes[-1]
|
||||
picked.append(current)
|
||||
if 0 < top_k == len(picked) or len(indexes) == 1:
|
||||
break
|
||||
current_box = boxes[current, :]
|
||||
# indexes = indexes[1:]
|
||||
indexes = indexes[:-1]
|
||||
rest_boxes = boxes[indexes, :]
|
||||
iou = iou_of(
|
||||
rest_boxes,
|
||||
np.expand_dims(current_box, axis=0),
|
||||
)
|
||||
indexes = indexes[iou <= iou_threshold]
|
||||
|
||||
return box_scores[picked, :]
|
||||
|
||||
|
||||
def iou_of(boxes0, boxes1, eps=1e-5):
|
||||
"""Return intersection-over-union (Jaccard index) of boxes.
|
||||
|
||||
Args:
|
||||
boxes0 (N, 4): ground truth boxes.
|
||||
boxes1 (N or 1, 4): predicted boxes.
|
||||
eps: a small number to avoid 0 as denominator.
|
||||
Returns:
|
||||
iou (N): IoU values.
|
||||
"""
|
||||
overlap_left_top = np.maximum(boxes0[..., :2], boxes1[..., :2])
|
||||
overlap_right_bottom = np.minimum(boxes0[..., 2:], boxes1[..., 2:])
|
||||
|
||||
overlap_area = area_of(overlap_left_top, overlap_right_bottom)
|
||||
area0 = area_of(boxes0[..., :2], boxes0[..., 2:])
|
||||
area1 = area_of(boxes1[..., :2], boxes1[..., 2:])
|
||||
return overlap_area / (area0 + area1 - overlap_area + eps)
|
||||
|
||||
|
||||
def area_of(left_top, right_bottom):
|
||||
"""Compute the areas of rectangles given two corners.
|
||||
|
||||
Args:
|
||||
left_top (N, 2): left top corner.
|
||||
right_bottom (N, 2): right bottom corner.
|
||||
|
||||
Returns:
|
||||
area (N): return the area.
|
||||
"""
|
||||
hw = np.clip(right_bottom - left_top, 0.0, None)
|
||||
return hw[..., 0] * hw[..., 1]
|
||||
|
||||
|
||||
class NanoDetABC(metaclass=ABCMeta):
|
||||
def __init__(
|
||||
self,
|
||||
input_shape=[272, 160],
|
||||
reg_max=7,
|
||||
strides=[8, 16, 32],
|
||||
prob_threshold=0.4,
|
||||
iou_threshold=0.3,
|
||||
num_candidate=1000,
|
||||
top_k=-1,
|
||||
class_names=["face"],
|
||||
):
|
||||
self.strides = strides
|
||||
self.input_shape = input_shape
|
||||
self.reg_max = reg_max
|
||||
self.prob_threshold = prob_threshold
|
||||
self.iou_threshold = iou_threshold
|
||||
self.num_candidate = num_candidate
|
||||
self.top_k = top_k
|
||||
self.img_mean = [103.53, 116.28, 123.675]
|
||||
self.img_std = [57.375, 57.12, 58.395]
|
||||
self.input_size = (self.input_shape[1], self.input_shape[0])
|
||||
self.class_names = class_names
|
||||
self.num_classes = len(self.class_names)
|
||||
|
||||
def preprocess(self, img):
|
||||
# resize image
|
||||
ResizeM = get_resize_matrix((img.shape[1], img.shape[0]), self.input_size, True)
|
||||
img_resize = cv2.warpPerspective(img, ResizeM, dsize=self.input_size)
|
||||
|
||||
# normalize image
|
||||
img_input = img_resize.astype(np.float32) / 255
|
||||
img_mean = np.array(self.img_mean, dtype=np.float32).reshape(1, 1, 3) / 255
|
||||
img_std = np.array(self.img_std, dtype=np.float32).reshape(1, 1, 3) / 255
|
||||
img_input = (img_input - img_mean) / img_std
|
||||
|
||||
# expand dims
|
||||
img_input = np.transpose(img_input, [2, 0, 1])
|
||||
img_input = np.expand_dims(img_input, axis=0)
|
||||
return img_input, ResizeM
|
||||
|
||||
def postprocess(self, scores, raw_boxes, ResizeM, raw_shape):
|
||||
# generate centers
|
||||
decode_boxes = []
|
||||
select_scores = []
|
||||
for stride, box_distribute, score in zip(self.strides, raw_boxes, scores):
|
||||
# centers
|
||||
fm_h = self.input_shape[0] / stride
|
||||
fm_w = self.input_shape[1] / stride
|
||||
|
||||
h_range = np.arange(fm_h)
|
||||
w_range = np.arange(fm_w)
|
||||
ww, hh = np.meshgrid(w_range, h_range)
|
||||
|
||||
ct_row = hh.flatten() * stride
|
||||
ct_col = ww.flatten() * stride
|
||||
|
||||
center = np.stack((ct_col, ct_row, ct_col, ct_row), axis=1)
|
||||
|
||||
# box distribution to distance
|
||||
reg_range = np.arange(self.reg_max + 1)
|
||||
box_distance = box_distribute.reshape((-1, self.reg_max + 1))
|
||||
box_distance = softmax(box_distance, axis=1)
|
||||
box_distance = box_distance * np.expand_dims(reg_range, axis=0)
|
||||
box_distance = np.sum(box_distance, axis=1).reshape((-1, 4))
|
||||
box_distance = box_distance * stride
|
||||
|
||||
# top K candidate
|
||||
topk_idx = np.argsort(score.max(axis=1))[::-1]
|
||||
topk_idx = topk_idx[: self.num_candidate]
|
||||
center = center[topk_idx]
|
||||
score = score[topk_idx]
|
||||
box_distance = box_distance[topk_idx]
|
||||
|
||||
# decode box
|
||||
decode_box = center + [-1, -1, 1, 1] * box_distance
|
||||
|
||||
select_scores.append(score)
|
||||
decode_boxes.append(decode_box)
|
||||
|
||||
# nms
|
||||
bboxes = np.concatenate(decode_boxes, axis=0)
|
||||
confidences = np.concatenate(select_scores, axis=0)
|
||||
picked_box_probs = []
|
||||
picked_labels = []
|
||||
for class_index in range(0, confidences.shape[1]):
|
||||
probs = confidences[:, class_index]
|
||||
mask = probs > self.prob_threshold
|
||||
probs = probs[mask]
|
||||
if probs.shape[0] == 0:
|
||||
continue
|
||||
subset_boxes = bboxes[mask, :]
|
||||
box_probs = np.concatenate([subset_boxes, probs.reshape(-1, 1)], axis=1)
|
||||
box_probs = hard_nms(
|
||||
box_probs,
|
||||
iou_threshold=self.iou_threshold,
|
||||
top_k=self.top_k,
|
||||
)
|
||||
picked_box_probs.append(box_probs)
|
||||
picked_labels.extend([class_index] * box_probs.shape[0])
|
||||
if not picked_box_probs:
|
||||
return np.array([]), np.array([]), np.array([])
|
||||
picked_box_probs = np.concatenate(picked_box_probs)
|
||||
|
||||
# resize output boxes
|
||||
picked_box_probs[:, :4] = warp_boxes(
|
||||
picked_box_probs[:, :4], np.linalg.inv(ResizeM), raw_shape[1], raw_shape[0]
|
||||
)
|
||||
return (
|
||||
picked_box_probs[:, :4].astype(np.int32),
|
||||
np.array(picked_labels),
|
||||
picked_box_probs[:, 4],
|
||||
)
|
||||
|
||||
@abstractmethod
|
||||
def infer_image(self, img_input):
|
||||
pass
|
||||
|
||||
def detect(self, img):
|
||||
raw_shape = img.shape
|
||||
img_input, ResizeM = self.preprocess(img)
|
||||
scores, raw_boxes = self.infer_image(img_input)
|
||||
if scores[0].ndim == 1: # handling num_classes=1 case
|
||||
scores = [x[:, None] for x in scores]
|
||||
bbox, label, score = self.postprocess(scores, raw_boxes, ResizeM, raw_shape)
|
||||
|
||||
return bbox, label, score
|
||||
|
||||
|
||||
class FaceDet(NanoDetABC):
|
||||
def __init__(self, model_path="", gpu_id=None, *args, **kwargs):
|
||||
super(FaceDet, self).__init__(*args, **kwargs)
|
||||
|
||||
self.model_path = model_path
|
||||
self.ort_session = create_onnx_session(model_path, gpu_id=gpu_id)
|
||||
self.input_name = self.ort_session.get_inputs()[0].name
|
||||
|
||||
def infer_image(self, img_input):
|
||||
inference_results = self.ort_session.run(None, {self.input_name: img_input})
|
||||
|
||||
scores = [np.squeeze(x) for x in inference_results[:3]]
|
||||
raw_boxes = [np.squeeze(x) for x in inference_results[3:]]
|
||||
return scores, raw_boxes
|
||||
@@ -0,0 +1,149 @@
|
||||
import math
|
||||
import time
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import onnx
|
||||
import onnxruntime
|
||||
|
||||
|
||||
def create_onnx_session(onnx_path, gpu_id=None) -> onnxruntime.InferenceSession:
|
||||
start = time.perf_counter()
|
||||
onnx_model = onnx.load(onnx_path)
|
||||
onnx.checker.check_model(onnx_model)
|
||||
providers = (
|
||||
[
|
||||
(
|
||||
"CUDAExecutionProvider",
|
||||
{
|
||||
"device_id": int(gpu_id),
|
||||
"arena_extend_strategy": "kNextPowerOfTwo",
|
||||
"cudnn_conv_algo_search": "EXHAUSTIVE",
|
||||
"do_copy_in_default_stream": True,
|
||||
},
|
||||
),
|
||||
"CPUExecutionProvider",
|
||||
]
|
||||
if (gpu_id is not None and gpu_id >= 0)
|
||||
else ["CPUExecutionProvider"]
|
||||
)
|
||||
|
||||
sess = onnxruntime.InferenceSession(onnx_path, providers=providers)
|
||||
print(
|
||||
"create onnx session cost: {:.3f}s. {}".format(
|
||||
time.perf_counter() - start, onnx_path
|
||||
)
|
||||
)
|
||||
return sess
|
||||
|
||||
|
||||
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, dx0=0.0, d_cutoff=1.0):
|
||||
"""Initialize the one euro filter."""
|
||||
# self.min_cutoff = float(min_cutoff)
|
||||
# self.beta = float(beta)
|
||||
self.d_cutoff = float(d_cutoff)
|
||||
self.dx_prev = float(dx0)
|
||||
# self.t_e = fcmin
|
||||
|
||||
def __call__(self, x, x_prev, fcmin=1.0, min_cutoff=1.0, beta=0.0):
|
||||
if x_prev is None:
|
||||
return x
|
||||
# t_e = 1
|
||||
a_d = smoothing_factor(fcmin, self.d_cutoff)
|
||||
dx = (x - x_prev) / fcmin
|
||||
dx_hat = exponential_smoothing(a_d, dx, self.dx_prev)
|
||||
cutoff = min_cutoff + beta * abs(dx_hat)
|
||||
a = smoothing_factor(fcmin, cutoff)
|
||||
x_hat = exponential_smoothing(a, x, x_prev)
|
||||
self.dx_prev = dx_hat
|
||||
return x_hat
|
||||
|
||||
|
||||
def get_warp_mat_bbox(
|
||||
face_bbox, base_angle, dst_size=128, expand_ratio=0.15, aug_angle=0.0, aug_scale=1.0
|
||||
):
|
||||
face_x_min, face_y_min, face_x_max, face_y_max = face_bbox
|
||||
face_x_center = (face_x_min + face_x_max) / 2
|
||||
face_y_center = (face_y_min + face_y_max) / 2
|
||||
face_width = face_x_max - face_x_min
|
||||
face_height = face_y_max - face_y_min
|
||||
scale = dst_size / max(face_width, face_height) * (1 - expand_ratio) * aug_scale
|
||||
M = cv2.getRotationMatrix2D(
|
||||
(face_x_center, face_y_center), angle=base_angle + aug_angle, scale=scale
|
||||
)
|
||||
offset = [dst_size / 2 - face_x_center, dst_size / 2 - face_y_center]
|
||||
M[:, 2] += offset
|
||||
return M
|
||||
|
||||
|
||||
def transform_points(points, mat, invert=False):
|
||||
if invert:
|
||||
mat = cv2.invertAffineTransform(mat)
|
||||
points = np.expand_dims(points, axis=1)
|
||||
points = cv2.transform(points, mat, points.shape)
|
||||
points = np.squeeze(points)
|
||||
return points
|
||||
|
||||
|
||||
def get_warp_mat_bbox_by_gt_pts_float(
|
||||
gt_pts, base_angle=0.0, dst_size=128, expand_ratio=0.15, return_info=False
|
||||
):
|
||||
# step 1
|
||||
face_x_min, face_x_max = np.min(gt_pts[:, 0]), np.max(gt_pts[:, 0])
|
||||
face_y_min, face_y_max = np.min(gt_pts[:, 1]), np.max(gt_pts[:, 1])
|
||||
face_x_center = (face_x_min + face_x_max) / 2
|
||||
face_y_center = (face_y_min + face_y_max) / 2
|
||||
M_step_1 = cv2.getRotationMatrix2D(
|
||||
(face_x_center, face_y_center), angle=base_angle, scale=1.0
|
||||
)
|
||||
pts_step_1 = transform_points(gt_pts, M_step_1)
|
||||
face_x_min_step_1, face_x_max_step_1 = np.min(pts_step_1[:, 0]), np.max(
|
||||
pts_step_1[:, 0]
|
||||
)
|
||||
face_y_min_step_1, face_y_max_step_1 = np.min(pts_step_1[:, 1]), np.max(
|
||||
pts_step_1[:, 1]
|
||||
)
|
||||
# step 2
|
||||
face_width = face_x_max_step_1 - face_x_min_step_1
|
||||
face_height = face_y_max_step_1 - face_y_min_step_1
|
||||
scale = dst_size / max(face_width, face_height) * (1 - expand_ratio)
|
||||
M_step_2 = cv2.getRotationMatrix2D(
|
||||
(face_x_center, face_y_center), angle=base_angle, scale=scale
|
||||
)
|
||||
pts_step_2 = transform_points(gt_pts, M_step_2)
|
||||
face_x_min_step_2, face_x_max_step_2 = np.min(pts_step_2[:, 0]), np.max(
|
||||
pts_step_2[:, 0]
|
||||
)
|
||||
face_y_min_step_2, face_y_max_step_2 = np.min(pts_step_2[:, 1]), np.max(
|
||||
pts_step_2[:, 1]
|
||||
)
|
||||
face_x_center_step_2 = (face_x_min_step_2 + face_x_max_step_2) / 2
|
||||
face_y_center_step_2 = (face_y_min_step_2 + face_y_max_step_2) / 2
|
||||
|
||||
M = cv2.getRotationMatrix2D(
|
||||
(face_x_center, face_y_center), angle=base_angle, scale=scale
|
||||
)
|
||||
offset = [dst_size / 2 - face_x_center_step_2, dst_size / 2 - face_y_center_step_2]
|
||||
M[:, 2] += offset
|
||||
|
||||
if not return_info:
|
||||
return M
|
||||
else:
|
||||
transform_info = {
|
||||
"M": M,
|
||||
"center_x": face_x_center,
|
||||
"center_y": face_y_center,
|
||||
"rotate_angle": base_angle,
|
||||
"scale": scale,
|
||||
}
|
||||
return transform_info
|
||||
@@ -0,0 +1,343 @@
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from ..wanvideo.modules.attention import attention
|
||||
|
||||
def FeedForward(dim, mult=4):
|
||||
inner_dim = int(dim * mult)
|
||||
return nn.Sequential(
|
||||
nn.LayerNorm(dim),
|
||||
nn.Linear(dim, inner_dim, bias=False),
|
||||
nn.GELU(),
|
||||
nn.Linear(inner_dim, dim, bias=False),
|
||||
)
|
||||
|
||||
|
||||
def reshape_tensor(x, heads):
|
||||
bs, length, width = x.shape
|
||||
x = x.view(bs, length, heads, -1)
|
||||
x = x.transpose(1, 2)
|
||||
x = x.reshape(bs, heads, length, -1)
|
||||
return x
|
||||
|
||||
|
||||
class MultiProjModel(nn.Module):
|
||||
def __init__(self, adapter_in_dim=1024, cross_attention_dim=1024):
|
||||
super().__init__()
|
||||
|
||||
self.generator = None
|
||||
self.cross_attention_dim = cross_attention_dim
|
||||
self.eye_proj = torch.nn.Linear(6, cross_attention_dim, bias=False)
|
||||
self.emo_proj = torch.nn.Linear(30, cross_attention_dim, bias=False)
|
||||
self.mouth_proj = torch.nn.Linear(512, cross_attention_dim, bias=False)
|
||||
self.headpose_proj = torch.nn.Linear(6, cross_attention_dim, bias=False)
|
||||
|
||||
self.norm = torch.nn.LayerNorm(cross_attention_dim)
|
||||
|
||||
def forward(self, adapter_embeds):
|
||||
B, num_frames, C = adapter_embeds.shape
|
||||
embeds = adapter_embeds
|
||||
split_sizes = [6, 6, 30, 512]
|
||||
headpose, eye, emo, mouth = torch.split(embeds, split_sizes, dim=-1)
|
||||
headpose = self.norm(self.headpose_proj(headpose))
|
||||
eye = self.norm(self.eye_proj(eye))
|
||||
emo = self.norm(self.emo_proj(emo))
|
||||
mouth = self.norm(self.mouth_proj(mouth))
|
||||
|
||||
all_features = torch.stack([headpose, eye, emo, mouth], dim=2)
|
||||
result_final = all_features.view(B, num_frames * 4, self.cross_attention_dim)
|
||||
|
||||
return result_final
|
||||
|
||||
|
||||
class SingleStreamBlockProcessor(nn.Module):
|
||||
def __init__(self, context_dim, hidden_dim):
|
||||
super().__init__()
|
||||
|
||||
self.context_dim = context_dim
|
||||
self.hidden_dim = hidden_dim
|
||||
|
||||
self.ip_adapter_single_stream_k_proj = nn.Linear(
|
||||
context_dim, hidden_dim, bias=False
|
||||
)
|
||||
self.ip_adapter_single_stream_v_proj = nn.Linear(
|
||||
context_dim, hidden_dim, bias=False
|
||||
)
|
||||
|
||||
nn.init.zeros_(self.ip_adapter_single_stream_k_proj.weight)
|
||||
nn.init.zeros_(self.ip_adapter_single_stream_v_proj.weight)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: nn.Module,
|
||||
x: torch.Tensor,
|
||||
context: torch.Tensor,
|
||||
context_lens: torch.Tensor,
|
||||
adapter_proj: torch.Tensor,
|
||||
adapter_context_lens: torch.Tensor,
|
||||
latents_num_frames: int = 21,
|
||||
ip_scale: float = 1.0,
|
||||
adapter_attn_mask: torch.Tensor = None,
|
||||
) -> torch.Tensor:
|
||||
context_img = context[:, :257]
|
||||
context = context[:, 257:]
|
||||
b, n, d = x.size(0), attn.num_heads, attn.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = attn.norm_q(attn.q(x)).view(b, -1, n, d)
|
||||
k = attn.norm_k(attn.k(context)).view(b, -1, n, d)
|
||||
v = attn.v(context).view(b, -1, n, d)
|
||||
k_img = attn.norm_k_img(attn.k_img(context_img)).view(b, -1, n, d)
|
||||
v_img = attn.v_img(context_img).view(b, -1, n, d)
|
||||
img_x = attention(q, k_img, v_img)
|
||||
# compute attention
|
||||
x = attention(q, k, v)
|
||||
|
||||
x = x.flatten(2)
|
||||
img_x = img_x.flatten(2)
|
||||
|
||||
if len(adapter_proj.shape) == 4:
|
||||
adapter_q = q.view(b * latents_num_frames, -1, n, d)
|
||||
ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(
|
||||
b * latents_num_frames, -1, n, d
|
||||
)
|
||||
ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(
|
||||
b * latents_num_frames, -1, n, d
|
||||
)
|
||||
adapter_x = attention(
|
||||
adapter_q, ip_key, ip_value, attn_mask=adapter_attn_mask
|
||||
)
|
||||
adapter_x = adapter_x.view(b, q.size(1), n, d)
|
||||
adapter_x = adapter_x.flatten(2)
|
||||
elif len(adapter_proj.shape) == 3:
|
||||
ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(
|
||||
b, -1, n, d
|
||||
)
|
||||
ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(
|
||||
b, -1, n, d
|
||||
)
|
||||
adapter_x = attention(q, ip_key, ip_value, attn_mask=adapter_attn_mask)
|
||||
adapter_x = adapter_x.flatten(2)
|
||||
|
||||
x = x + img_x + adapter_x * ip_scale
|
||||
x = attn.o(x)
|
||||
return x
|
||||
|
||||
|
||||
class PerceiverAttention(nn.Module):
|
||||
def __init__(self, *, dim, dim_head=64, heads=8):
|
||||
super().__init__()
|
||||
self.scale = dim_head**-0.5
|
||||
self.dim_head = dim_head
|
||||
self.heads = heads
|
||||
inner_dim = dim_head * heads
|
||||
|
||||
self.norm1 = nn.LayerNorm(dim)
|
||||
self.norm2 = nn.LayerNorm(dim)
|
||||
|
||||
self.to_q = nn.Linear(dim, inner_dim, bias=False)
|
||||
self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
|
||||
self.to_out = nn.Linear(inner_dim, dim, bias=False)
|
||||
|
||||
def forward(self, x, latents):
|
||||
"""
|
||||
Args:
|
||||
x (torch.Tensor): image features
|
||||
shape (b, n1, D)
|
||||
latent (torch.Tensor): latent features
|
||||
shape (b, n2, D)
|
||||
"""
|
||||
x = self.norm1(x)
|
||||
latents = self.norm2(latents)
|
||||
|
||||
b, l, _ = latents.shape
|
||||
|
||||
q = self.to_q(latents)
|
||||
kv_input = torch.cat((x, latents), dim=-2)
|
||||
k, v = self.to_kv(kv_input).chunk(2, dim=-1)
|
||||
|
||||
q = reshape_tensor(q, self.heads)
|
||||
k = reshape_tensor(k, self.heads)
|
||||
v = reshape_tensor(v, self.heads)
|
||||
|
||||
# attention
|
||||
scale = 1 / math.sqrt(math.sqrt(self.dim_head))
|
||||
weight = (q * scale) @ (k * scale).transpose(
|
||||
-2, -1
|
||||
) # More stable with f16 than dividing afterwards
|
||||
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
|
||||
out = weight @ v
|
||||
|
||||
out = out.permute(0, 2, 1, 3).reshape(b, l, -1)
|
||||
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class Resampler(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim=1024,
|
||||
depth=8,
|
||||
dim_head=64,
|
||||
heads=16,
|
||||
num_queries=8,
|
||||
embedding_dim=768,
|
||||
output_dim=1024,
|
||||
ff_mult=4,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5)
|
||||
|
||||
self.proj_in = nn.Linear(embedding_dim, dim)
|
||||
|
||||
self.proj_out = nn.Linear(dim, output_dim)
|
||||
self.norm_out = nn.LayerNorm(output_dim)
|
||||
|
||||
self.layers = nn.ModuleList([])
|
||||
for _ in range(depth):
|
||||
self.layers.append(
|
||||
nn.ModuleList(
|
||||
[
|
||||
PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads),
|
||||
FeedForward(dim=dim, mult=ff_mult),
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
def forward(self, x): # x (b, 512, 1)
|
||||
latents = self.latents.repeat(x.size(0), 1, 1)
|
||||
|
||||
x = self.proj_in(x) # (b, 512, 1024)
|
||||
|
||||
for attn, ff in self.layers:
|
||||
latents = attn(x, latents) + latents # b 16 1024
|
||||
latents = ff(latents) + latents
|
||||
|
||||
latents = self.proj_out(latents)
|
||||
return self.norm_out(latents)
|
||||
|
||||
|
||||
class PortraitAdapter(nn.Module):
|
||||
def __init__(self, adapter_in_dim: int, adapter_proj_dim: int, dtype: torch.dtype):
|
||||
super().__init__()
|
||||
|
||||
self.adapter_in_dim = adapter_in_dim
|
||||
self.adapter_proj_dim = adapter_proj_dim
|
||||
self.proj_model = self.init_proj(self.adapter_proj_dim)
|
||||
self.dtype = dtype
|
||||
|
||||
self.mouth_proj_model = Resampler(
|
||||
dim=1280,
|
||||
depth=4,
|
||||
dim_head=64,
|
||||
heads=20,
|
||||
num_queries=16,
|
||||
embedding_dim=512,
|
||||
output_dim=2048,
|
||||
ff_mult=4,
|
||||
)
|
||||
|
||||
self.emo_proj_model = Resampler(
|
||||
dim=1280,
|
||||
depth=4,
|
||||
dim_head=64,
|
||||
heads=20,
|
||||
num_queries=4,
|
||||
embedding_dim=30,
|
||||
output_dim=2048,
|
||||
ff_mult=4,
|
||||
)
|
||||
|
||||
def init_proj(self, cross_attention_dim=5120):
|
||||
proj_model = MultiProjModel(
|
||||
adapter_in_dim=self.adapter_in_dim, cross_attention_dim=cross_attention_dim
|
||||
)
|
||||
return proj_model
|
||||
|
||||
def get_adapter_proj(self, adapter_fea=None):
|
||||
split_sizes = [6, 6, 30, 512]
|
||||
headpose, eye, emo, mouth = torch.split(
|
||||
adapter_fea, split_sizes, dim=-1
|
||||
)
|
||||
B, frames, dim = mouth.shape
|
||||
mouth = mouth.view(B * frames, 1, 512)
|
||||
emo = emo.view(B * frames, 1, 30)
|
||||
|
||||
mouth_fea = self.mouth_proj_model(mouth)
|
||||
emo_fea = self.emo_proj_model(emo)
|
||||
|
||||
mouth_fea = mouth_fea.view(B, frames, 16, 2048)
|
||||
emo_fea = emo_fea.view(B, frames, 4, 2048)
|
||||
|
||||
adapter_fea = self.proj_model(adapter_fea)
|
||||
|
||||
adapter_fea = adapter_fea.view(B, frames, 4, 2048)
|
||||
|
||||
all_fea = torch.cat([adapter_fea, mouth_fea, emo_fea], dim=2)
|
||||
|
||||
result_final = all_fea.view(B, frames * 24, 2048)
|
||||
|
||||
return result_final
|
||||
|
||||
|
||||
def split_audio_adapter_sequence(self, adapter_proj_length, num_frames=80):
|
||||
tokens_pre_frame = adapter_proj_length / num_frames
|
||||
tokens_pre_latents_frame = tokens_pre_frame * 4
|
||||
half_tokens_pre_latents_frame = tokens_pre_latents_frame / 2
|
||||
pos_idx = []
|
||||
for i in range(int((num_frames - 1) / 4) + 1):
|
||||
if i == 0:
|
||||
pos_idx.append(0)
|
||||
else:
|
||||
begin_token_id = tokens_pre_frame * ((i - 1) * 4 + 1)
|
||||
end_token_id = tokens_pre_frame * (i * 4 + 1)
|
||||
pos_idx.append(int((sum([begin_token_id, end_token_id]) / 2)) - 1)
|
||||
pos_idx_range = [
|
||||
[
|
||||
idx - int(half_tokens_pre_latents_frame),
|
||||
idx + int(half_tokens_pre_latents_frame),
|
||||
]
|
||||
for idx in pos_idx
|
||||
]
|
||||
pos_idx_range[0] = [
|
||||
-(int(half_tokens_pre_latents_frame) * 2 - pos_idx_range[1][0]),
|
||||
pos_idx_range[1][0],
|
||||
]
|
||||
return pos_idx_range
|
||||
|
||||
|
||||
def split_tensor_with_padding(self, input_tensor, pos_idx_range, expand_length=0):
|
||||
pos_idx_range = [
|
||||
[idx[0] - expand_length, idx[1] + expand_length] for idx in pos_idx_range
|
||||
]
|
||||
sub_sequences = []
|
||||
seq_len = input_tensor.size(1)
|
||||
max_valid_idx = seq_len - 1
|
||||
k_lens_list = []
|
||||
for start, end in pos_idx_range:
|
||||
pad_front = max(-start, 0)
|
||||
pad_back = max(end - max_valid_idx, 0)
|
||||
|
||||
valid_start = max(start, 0)
|
||||
valid_end = min(end, max_valid_idx)
|
||||
|
||||
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)))
|
||||
|
||||
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
|
||||
)
|
||||
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,203 @@
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
from ..utils import log
|
||||
|
||||
from accelerate import init_empty_weights
|
||||
from accelerate.utils import set_module_tensor_to_device
|
||||
|
||||
import comfy.model_management as mm
|
||||
from comfy.utils import load_torch_file, ProgressBar
|
||||
import folder_paths
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
alignment_model_path = os.path.join(script_directory, "models", "face_landmark.onnx")
|
||||
det_model_path = os.path.join(script_directory, "models", "face_det.onnx")
|
||||
|
||||
from .model import PortraitAdapter
|
||||
from .pdf import get_drive_expression_pd_fgc, det_landmarks, FanEncoder
|
||||
from .camer import CameraDemo
|
||||
from .face_align import FaceAlignment
|
||||
|
||||
def load_pd_fgc_model(state_dict):
|
||||
face_aligner = CameraDemo(
|
||||
face_alignment_module=FaceAlignment(
|
||||
gpu_id=None,
|
||||
alignment_model_path=alignment_model_path,
|
||||
det_model_path=det_model_path,
|
||||
),
|
||||
reset=False,
|
||||
)
|
||||
|
||||
pd_fpg_motion = FanEncoder()
|
||||
m, u = pd_fpg_motion.load_state_dict(state_dict, strict=False)
|
||||
pd_fpg_motion = pd_fpg_motion.eval()
|
||||
|
||||
return face_aligner, pd_fpg_motion
|
||||
|
||||
|
||||
def get_emo_feature(frame_list, face_aligner, pd_fpg_motion, device):
|
||||
|
||||
|
||||
comfy_pbar = ProgressBar(3)
|
||||
landmark_list = det_landmarks(face_aligner, frame_list, comfy_pbar)[1]
|
||||
emo_list = get_drive_expression_pd_fgc(pd_fpg_motion, frame_list, landmark_list, device)
|
||||
comfy_pbar.update(1)
|
||||
|
||||
#emo_feat_list = []
|
||||
head_emo_feat_list = []
|
||||
for emo in emo_list:
|
||||
headpose_emb = emo["headpose_emb"]
|
||||
eye_embed = emo["eye_embed"]
|
||||
emo_embed = emo["emo_embed"]
|
||||
mouth_feat = emo["mouth_feat"]
|
||||
|
||||
emo_feat = torch.cat([eye_embed, emo_embed, mouth_feat], dim=1)
|
||||
head_emo_feat = torch.cat([headpose_emb, emo_feat], dim=1)
|
||||
|
||||
#emo_feat_list.append(emo_feat)
|
||||
head_emo_feat_list.append(head_emo_feat)
|
||||
|
||||
#emo_feat_all = torch.cat(emo_feat_list, dim=0).unsqueeze(0)
|
||||
head_emo_feat_all = torch.cat(head_emo_feat_list, dim=0).unsqueeze(0)
|
||||
|
||||
return head_emo_feat_all
|
||||
|
||||
class FantasyPortraitFaceDetector:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"portrait_model": ("FANTASYPORTRAITMODEL",),
|
||||
"images": ("IMAGE",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("PORTRAIT_EMBEDS",)
|
||||
RETURN_NAMES = ("portrait_embeds", )
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def loadmodel(self, images, portrait_model):
|
||||
B, H, W, C = images.shape
|
||||
num_frames = ((B - 1) // 4) * 4 + 1
|
||||
images = images.clone()[:num_frames]
|
||||
|
||||
def tensor_batch_to_numpy_list(images):
|
||||
images = images.detach().cpu()
|
||||
numpy_list = []
|
||||
for img in images:
|
||||
# img shape: (H, W, C)
|
||||
img = img.numpy()
|
||||
img = img[..., :3]
|
||||
img = (img * 255).clip(0, 255)
|
||||
img = img.astype(np.uint8)
|
||||
numpy_list.append(img)
|
||||
return numpy_list
|
||||
|
||||
|
||||
numpy_list = tensor_batch_to_numpy_list(images)
|
||||
|
||||
pd_fpg_sd = {}
|
||||
for k, v in portrait_model["sd"].items():
|
||||
if k.startswith("pd_fpg."):
|
||||
pd_fpg_sd[k.replace("pd_fpg.", "")] = v
|
||||
|
||||
face_aligner, pd_fpg_motion = load_pd_fgc_model(pd_fpg_sd)
|
||||
|
||||
pd_fpg_motion.to(device)
|
||||
head_emo_feat_all = get_emo_feature(numpy_list, face_aligner, pd_fpg_motion, device=device)
|
||||
pd_fpg_motion.to(offload_device)
|
||||
|
||||
portrait_model = portrait_model["proj_model"]
|
||||
|
||||
portrait_model.to(device)
|
||||
adapter_proj = portrait_model.get_adapter_proj(head_emo_feat_all.to(device, dtype=portrait_model.dtype))
|
||||
portrait_model.to(offload_device)
|
||||
|
||||
pos_idx_range = portrait_model.split_audio_adapter_sequence(adapter_proj.size(1), num_frames=num_frames)
|
||||
proj_split, context_lens = portrait_model.split_tensor_with_padding(adapter_proj, pos_idx_range, expand_length=0)
|
||||
|
||||
return (proj_split,)
|
||||
|
||||
class WanVideoAddFantasyPortrait:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"portrait_embeds": ("PORTRAIT_EMBEDS",),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the portrait embedding"}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, portrait_embeds, strength, start_percent=0.0, end_percent=1.0):
|
||||
new_entry = {
|
||||
"adapter_proj": portrait_embeds,
|
||||
"strength": strength,
|
||||
"start_percent": start_percent,
|
||||
"end_percent": end_percent,
|
||||
}
|
||||
|
||||
updated = dict(embeds)
|
||||
updated["portrait_embeds"] = new_entry
|
||||
return (updated,)
|
||||
|
||||
class FantasyPortraitModelLoader:
|
||||
@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 = ("FANTASYPORTRAITMODEL",)
|
||||
RETURN_NAMES = ("model", )
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def loadmodel(self, model, base_precision):
|
||||
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)
|
||||
adapter_in_dim = sd["proj_model.norm.weight"].shape[0]
|
||||
|
||||
with init_empty_weights():
|
||||
fantasyportrait_proj_adapter = PortraitAdapter(adapter_in_dim=adapter_in_dim, adapter_proj_dim=adapter_in_dim, dtype=base_dtype)
|
||||
|
||||
for name, param in fantasyportrait_proj_adapter.named_parameters():
|
||||
set_module_tensor_to_device(fantasyportrait_proj_adapter, name, device=offload_device, dtype=base_dtype, value=sd[name])
|
||||
|
||||
fantasyportrait = {
|
||||
"proj_model": fantasyportrait_proj_adapter,
|
||||
"sd": sd,
|
||||
}
|
||||
|
||||
return (fantasyportrait,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"FantasyPortraitModelLoader": FantasyPortraitModelLoader,
|
||||
"FantasyPortraitFaceDetector": FantasyPortraitFaceDetector,
|
||||
"WanVideoAddFantasyPortrait": WanVideoAddFantasyPortrait,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"FantasyPortraitModelLoader": "FantasyPortrait Model Loader",
|
||||
"FantasyPortraitFaceDetector": "FantasyPortrait Face Detector",
|
||||
"WanVideoAddFantasyPortrait": "WanVideo Add Fantasy Portrait",
|
||||
}
|
||||
@@ -0,0 +1,406 @@
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
def np_bgr_to_tensor(img_np, dtype):
|
||||
img_rgb = cv2.cvtColor(img_np, cv2.COLOR_BGR2RGB) / 255.0 * 2 - 1
|
||||
return torch.tensor(img_rgb).permute(2, 0, 1).to(dtype=dtype)
|
||||
|
||||
|
||||
def image_preprocess(np_bgr, size, dtype=torch.float32):
|
||||
img_np = cv2.resize(np_bgr, size)
|
||||
return np_bgr_to_tensor(img_np, dtype)
|
||||
|
||||
|
||||
def umeyama(src, dst, estimate_scale):
|
||||
"""Estimate N-D similarity transformation with or without scaling.
|
||||
Parameters
|
||||
----------
|
||||
src : (M, N) array
|
||||
Source coordinates.
|
||||
dst : (M, N) array
|
||||
Destination coordinates.
|
||||
estimate_scale : bool
|
||||
Whether to estimate scaling factor.
|
||||
Returns
|
||||
-------
|
||||
T : (N + 1, N + 1)
|
||||
The homogeneous similarity transformation matrix. The matrix contains
|
||||
NaN values only if the problem is not well-conditioned.
|
||||
References
|
||||
----------
|
||||
.. [1] "Least-squares estimation of transformation parameters between two
|
||||
point patterns", Shinji Umeyama, PAMI 1991, DOI: 10.1109/34.88573
|
||||
"""
|
||||
|
||||
num = src.shape[0]
|
||||
dim = src.shape[1]
|
||||
|
||||
# Compute mean of src and dst.
|
||||
src_mean = src.mean(axis=0)
|
||||
dst_mean = dst.mean(axis=0)
|
||||
|
||||
# Subtract mean from src and dst.
|
||||
src_demean = src - src_mean
|
||||
dst_demean = dst - dst_mean
|
||||
|
||||
# Eq. (38).
|
||||
A = np.dot(dst_demean.T, src_demean) / num
|
||||
|
||||
# Eq. (39).
|
||||
d = np.ones((dim,), dtype=np.double)
|
||||
if np.linalg.det(A) < 0:
|
||||
d[dim - 1] = -1
|
||||
|
||||
T = np.eye(dim + 1, dtype=np.double)
|
||||
|
||||
U, S, V = np.linalg.svd(A)
|
||||
|
||||
# Eq. (40) and (43).
|
||||
rank = np.linalg.matrix_rank(A)
|
||||
if rank == 0:
|
||||
return np.nan * T
|
||||
elif rank == dim - 1:
|
||||
if np.linalg.det(U) * np.linalg.det(V) > 0:
|
||||
T[:dim, :dim] = np.dot(U, V)
|
||||
else:
|
||||
s = d[dim - 1]
|
||||
d[dim - 1] = -1
|
||||
T[:dim, :dim] = np.dot(U, np.dot(np.diag(d), V))
|
||||
d[dim - 1] = s
|
||||
else:
|
||||
T[:dim, :dim] = np.dot(U, np.dot(np.diag(d), V.T))
|
||||
|
||||
if estimate_scale:
|
||||
# Eq. (41) and (42).
|
||||
scale = 1.0 / src_demean.var(axis=0).sum() * np.dot(S, d)
|
||||
else:
|
||||
scale = 1.0
|
||||
|
||||
T[:dim, dim] = dst_mean - scale * np.dot(T[:dim, :dim], src_mean.T)
|
||||
T[:dim, :dim] *= scale
|
||||
|
||||
return T
|
||||
|
||||
|
||||
def warp_face_pd_fgc(image, landmarks222, save_size=224):
|
||||
pt5_idx = [182, 202, 36, 149, 133]
|
||||
dst_pt5 = (
|
||||
np.array(
|
||||
[
|
||||
[0.3843, 0.27],
|
||||
[0.62, 0.2668],
|
||||
[0.503, 0.4185],
|
||||
[0.406, 0.5273],
|
||||
[0.5977, 0.525],
|
||||
]
|
||||
)
|
||||
* save_size
|
||||
)
|
||||
src_pt5 = landmarks222[pt5_idx]
|
||||
|
||||
M = umeyama(src_pt5, dst_pt5, True)[0:2]
|
||||
warped = cv2.warpAffine(image, M, (save_size, save_size), flags=cv2.INTER_CUBIC)
|
||||
|
||||
return warped
|
||||
|
||||
|
||||
def get_drive_expression_pd_fgc(
|
||||
pd_fpg_motion, images, landmarks, device, dtype=torch.float32
|
||||
):
|
||||
emo_list = []
|
||||
|
||||
motion_model = pd_fpg_motion.to(device=device)
|
||||
with tqdm(total=len(images)) as pbar:
|
||||
for frame, landmark in zip(images, landmarks):
|
||||
emo_image = warp_face_pd_fgc(frame, landmark, save_size=224)
|
||||
input_tensor = (
|
||||
image_preprocess(emo_image, (224, 224), dtype)
|
||||
.to(device=device)
|
||||
.unsqueeze(0)
|
||||
)
|
||||
# headpose_emb, eye_embed, emo_embed, mouth_feat
|
||||
# emo_tensor = motion_model(input_tensor)
|
||||
# emo_list.append(emo_tensor)
|
||||
# headpose_emb [1, 6]; eye_embed [1, 6]; emo_embed [1, 30]; mouth_feat [1, 512]
|
||||
headpose_emb, eye_embed, emo_embed, mouth_feat = motion_model(input_tensor)
|
||||
emotion = {
|
||||
"headpose_emb": headpose_emb.cpu(),
|
||||
"eye_embed": eye_embed.cpu(),
|
||||
"emo_embed": emo_embed.cpu(),
|
||||
"mouth_feat": mouth_feat.cpu(),
|
||||
}
|
||||
emo_list.append(emotion)
|
||||
|
||||
pbar.set_description("PD_FPG_MOTION")
|
||||
pbar.update()
|
||||
|
||||
# neg_tensor = motion_model(torch.ones_like(input_tensor)*-1).cpu()
|
||||
|
||||
# ret_tensor = torch.cat(emo_list, dim=0)
|
||||
# pd_fpg_motion.to(device='cpu')
|
||||
# return dict(pd_fpg=ret_tensor.unsqueeze(0), neg_pd_fpg=neg_tensor.unsqueeze(0))
|
||||
return emo_list
|
||||
|
||||
|
||||
def det_landmarks(face_aligner, frame_list, comfy_pbar):
|
||||
rect_list = []
|
||||
new_frame_list = []
|
||||
|
||||
assert len(frame_list) > 0
|
||||
face_aligner.reset_track()
|
||||
|
||||
with tqdm(total=len(frame_list)) as pbar:
|
||||
for frame in frame_list:
|
||||
faces = face_aligner.forward(frame)
|
||||
if len(faces) > 0:
|
||||
face = sorted(
|
||||
faces,
|
||||
key=lambda x: (x["face_rect"][2] - x["face_rect"][0])
|
||||
* (x["face_rect"][3] - x["face_rect"][1]),
|
||||
)[-1]
|
||||
rect_list.append(face["face_rect"])
|
||||
new_frame_list.append(frame)
|
||||
pbar.set_description("DET stage1")
|
||||
pbar.update()
|
||||
comfy_pbar.update(1)
|
||||
|
||||
assert len(new_frame_list) > 0
|
||||
face_aligner.reset_track()
|
||||
save_frame_list = []
|
||||
save_landmark_list = []
|
||||
with tqdm(total=len(new_frame_list)) as pbar:
|
||||
for frame, rect in zip(new_frame_list, rect_list):
|
||||
faces = face_aligner.forward(frame, pre_rect=rect)
|
||||
if len(faces) > 0:
|
||||
face = sorted(
|
||||
faces,
|
||||
key=lambda x: (x["face_rect"][2] - x["face_rect"][0])
|
||||
* (x["face_rect"][3] - x["face_rect"][1]),
|
||||
)[-1]
|
||||
landmarks = face["pre_kpt_222"]
|
||||
save_frame_list.append(frame)
|
||||
save_landmark_list.append(landmarks)
|
||||
pbar.set_description("DET stage2")
|
||||
pbar.update()
|
||||
|
||||
comfy_pbar.update(1)
|
||||
|
||||
assert len(save_frame_list) > 0
|
||||
save_landmark_list = np.stack(save_landmark_list, axis=0)
|
||||
face_aligner.reset_track()
|
||||
return save_frame_list, save_landmark_list, rect_list
|
||||
|
||||
|
||||
def conv3x3(in_planes, out_planes, strd=1, padding=1, bias=False):
|
||||
"3x3 convolution with padding"
|
||||
return nn.Conv2d(
|
||||
in_planes, out_planes, kernel_size=3, stride=strd, padding=padding, bias=bias
|
||||
)
|
||||
|
||||
|
||||
class HourGlass(nn.Module):
|
||||
def __init__(self, num_modules, depth, num_features):
|
||||
super(HourGlass, self).__init__()
|
||||
self.num_modules = num_modules
|
||||
self.depth = depth
|
||||
self.features = num_features
|
||||
self.dropout = nn.Dropout(0.5)
|
||||
|
||||
self._generate_network(self.depth)
|
||||
|
||||
def _generate_network(self, level):
|
||||
self.add_module("b1_" + str(level), ConvBlock(256, 256))
|
||||
|
||||
self.add_module("b2_" + str(level), ConvBlock(256, 256))
|
||||
|
||||
if level > 1:
|
||||
self._generate_network(level - 1)
|
||||
else:
|
||||
self.add_module("b2_plus_" + str(level), ConvBlock(256, 256))
|
||||
|
||||
self.add_module("b3_" + str(level), ConvBlock(256, 256))
|
||||
|
||||
def _forward(self, level, inp):
|
||||
# Upper branch
|
||||
up1 = inp
|
||||
up1 = self._modules["b1_" + str(level)](up1)
|
||||
up1 = self.dropout(up1)
|
||||
# Lower branch
|
||||
low1 = F.max_pool2d(inp, 2, stride=2)
|
||||
low1 = self._modules["b2_" + str(level)](low1)
|
||||
|
||||
if level > 1:
|
||||
low2 = self._forward(level - 1, low1)
|
||||
else:
|
||||
low2 = low1
|
||||
low2 = self._modules["b2_plus_" + str(level)](low2)
|
||||
|
||||
low3 = low2
|
||||
low3 = self._modules["b3_" + str(level)](low3)
|
||||
up1size = up1.size()
|
||||
rescale_size = (up1size[2], up1size[3])
|
||||
up2 = F.interpolate(low3, size=rescale_size, mode="bilinear")
|
||||
|
||||
return up1 + up2
|
||||
|
||||
def forward(self, x):
|
||||
return self._forward(self.depth, x)
|
||||
|
||||
|
||||
class ConvBlock(nn.Module):
|
||||
def __init__(self, in_planes, out_planes):
|
||||
super(ConvBlock, self).__init__()
|
||||
self.bn1 = nn.BatchNorm2d(in_planes)
|
||||
self.conv1 = conv3x3(in_planes, int(out_planes / 2))
|
||||
self.bn2 = nn.BatchNorm2d(int(out_planes / 2))
|
||||
self.conv2 = conv3x3(int(out_planes / 2), int(out_planes / 4))
|
||||
self.bn3 = nn.BatchNorm2d(int(out_planes / 4))
|
||||
self.conv3 = conv3x3(int(out_planes / 4), int(out_planes / 4))
|
||||
|
||||
if in_planes != out_planes:
|
||||
self.downsample = nn.Sequential(
|
||||
nn.BatchNorm2d(in_planes),
|
||||
nn.ReLU(True),
|
||||
nn.Conv2d(in_planes, out_planes, kernel_size=1, stride=1, bias=False),
|
||||
)
|
||||
else:
|
||||
self.downsample = None
|
||||
|
||||
def forward(self, x):
|
||||
residual = x
|
||||
|
||||
out1 = self.bn1(x)
|
||||
out1 = F.relu(out1, True)
|
||||
out1 = self.conv1(out1)
|
||||
|
||||
out2 = self.bn2(out1)
|
||||
out2 = F.relu(out2, True)
|
||||
out2 = self.conv2(out2)
|
||||
|
||||
out3 = self.bn3(out2)
|
||||
out3 = F.relu(out3, True)
|
||||
out3 = self.conv3(out3)
|
||||
|
||||
out3 = torch.cat((out1, out2, out3), 1)
|
||||
|
||||
if self.downsample is not None:
|
||||
residual = self.downsample(residual)
|
||||
|
||||
out3 += residual
|
||||
|
||||
return out3
|
||||
|
||||
|
||||
class FAN_use(nn.Module):
|
||||
def __init__(self):
|
||||
super(FAN_use, self).__init__()
|
||||
self.num_modules = 1
|
||||
|
||||
# Base part
|
||||
self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3)
|
||||
self.bn1 = nn.BatchNorm2d(64)
|
||||
self.conv2 = ConvBlock(64, 128)
|
||||
self.conv3 = ConvBlock(128, 128)
|
||||
self.conv4 = ConvBlock(128, 256)
|
||||
|
||||
# Stacking part
|
||||
hg_module = 0
|
||||
self.add_module("m" + str(hg_module), HourGlass(1, 4, 256))
|
||||
self.add_module("top_m_" + str(hg_module), ConvBlock(256, 256))
|
||||
self.add_module(
|
||||
"conv_last" + str(hg_module),
|
||||
nn.Conv2d(256, 256, kernel_size=1, stride=1, padding=0),
|
||||
)
|
||||
self.add_module(
|
||||
"l" + str(hg_module), nn.Conv2d(256, 68, kernel_size=1, stride=1, padding=0)
|
||||
)
|
||||
self.add_module("bn_end" + str(hg_module), nn.BatchNorm2d(256))
|
||||
|
||||
if hg_module < self.num_modules - 1:
|
||||
self.add_module(
|
||||
"bl" + str(hg_module),
|
||||
nn.Conv2d(256, 256, kernel_size=1, stride=1, padding=0),
|
||||
)
|
||||
self.add_module(
|
||||
"al" + str(hg_module),
|
||||
nn.Conv2d(68, 256, kernel_size=1, stride=1, padding=0),
|
||||
)
|
||||
|
||||
self.avgpool = nn.MaxPool2d((2, 2), 2)
|
||||
self.conv6 = nn.Conv2d(68, 1, 3, 2, 1)
|
||||
self.fc = nn.Linear(28 * 28, 512)
|
||||
self.bn5 = nn.BatchNorm2d(68)
|
||||
self.relu = nn.ReLU(True)
|
||||
|
||||
def forward(self, x):
|
||||
x = F.relu(self.bn1(self.conv1(x)), True)
|
||||
x = F.max_pool2d(self.conv2(x), 2)
|
||||
x = self.conv3(x)
|
||||
x = self.conv4(x)
|
||||
|
||||
previous = x
|
||||
|
||||
i = 0
|
||||
hg = self._modules["m" + str(i)](previous)
|
||||
|
||||
ll = hg
|
||||
ll = self._modules["top_m_" + str(i)](ll)
|
||||
|
||||
ll = self._modules["bn_end" + str(i)](self._modules["conv_last" + str(i)](ll))
|
||||
tmp_out = self._modules["l" + str(i)](F.relu(ll))
|
||||
|
||||
net = self.relu(self.bn5(tmp_out))
|
||||
net = self.conv6(net)
|
||||
net = net.view(-1, net.shape[-2] * net.shape[-1])
|
||||
net = self.relu(net)
|
||||
net = self.fc(net)
|
||||
return net
|
||||
|
||||
|
||||
class FanEncoder(nn.Module):
|
||||
def __init__(self, pose_dim=6, eye_dim=6):
|
||||
super(FanEncoder, self).__init__()
|
||||
self.model = FAN_use()
|
||||
|
||||
self.to_mouth = nn.Sequential(
|
||||
nn.Linear(512, 512), nn.ReLU(), nn.BatchNorm1d(512), nn.Linear(512, 512)
|
||||
)
|
||||
self.mouth_embed = nn.Sequential(
|
||||
nn.ReLU(), nn.Linear(512, 512 - pose_dim - eye_dim)
|
||||
)
|
||||
|
||||
self.to_headpose = nn.Sequential(
|
||||
nn.Linear(512, 512), nn.ReLU(), nn.BatchNorm1d(512), nn.Linear(512, 512)
|
||||
)
|
||||
self.headpose_embed = nn.Sequential(nn.ReLU(), nn.Linear(512, pose_dim))
|
||||
|
||||
self.to_eye = nn.Sequential(
|
||||
nn.Linear(512, 512), nn.ReLU(), nn.BatchNorm1d(512), nn.Linear(512, 512)
|
||||
)
|
||||
self.eye_embed = nn.Sequential(nn.ReLU(), nn.Linear(512, eye_dim))
|
||||
|
||||
self.to_emo = nn.Sequential(
|
||||
nn.Linear(512, 512), nn.ReLU(), nn.BatchNorm1d(512), nn.Linear(512, 512)
|
||||
)
|
||||
self.emo_embed = nn.Sequential(nn.ReLU(), nn.Linear(512, 30))
|
||||
|
||||
def forward_feature(self, x):
|
||||
net = self.model(x)
|
||||
return net
|
||||
|
||||
def forward(self, x):
|
||||
x = self.model(x)
|
||||
mouth_feat = self.to_mouth(x)
|
||||
headpose_feat = self.to_headpose(x)
|
||||
headpose_emb = self.headpose_embed(headpose_feat)
|
||||
eye_feat = self.to_eye(x)
|
||||
eye_embed = self.eye_embed(eye_feat)
|
||||
emo_feat = self.to_emo(x)
|
||||
emo_embed = self.emo_embed(emo_feat)
|
||||
|
||||
return headpose_emb, eye_embed, emo_embed, mouth_feat
|
||||
@@ -1513,7 +1513,7 @@ class WanVideoScheduler: #WIP
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"scheduler": (scheduler_list, {"default": "uni_pc"}),
|
||||
"scheduler": (scheduler_list, {"default": "unipc"}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -1539,7 +1539,7 @@ class WanVideoSampler:
|
||||
"shift": ("FLOAT", {"default": 5.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": (scheduler_list, {"default": "uni_pc",}),
|
||||
"scheduler": (scheduler_list, {"default": "unipc",}),
|
||||
"riflex_freq_index": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1, "tooltip": "Frequency index for RIFLEX, disabled when 0, default 6. Allows for new frames to be generated after without looping"}),
|
||||
},
|
||||
"optional": {
|
||||
@@ -1946,6 +1946,18 @@ class WanVideoSampler:
|
||||
|
||||
shapes = [tuple(e.shape) for e in multitalk_audio_embedding]
|
||||
log.info(f"Multitalk audio features shapes (per speaker): {shapes}")
|
||||
|
||||
# FantasyPortrait
|
||||
fantasy_portrait_input = None
|
||||
fantasy_portrait_embeds = image_embeds.get("portrait_embeds", None)
|
||||
if fantasy_portrait_embeds is not None:
|
||||
print("Using FantasyPortrait embeddings")
|
||||
fantasy_portrait_input = {
|
||||
"adapter_proj": fantasy_portrait_embeds.get("adapter_proj", None),
|
||||
"strength": fantasy_portrait_embeds.get("strength", 1.0),
|
||||
"start_percent": fantasy_portrait_embeds.get("start_percent", 0.0),
|
||||
"end_percent": fantasy_portrait_embeds.get("end_percent", 1.0),
|
||||
}
|
||||
|
||||
# MiniMax Remover
|
||||
minimax_latents = minimax_mask_latents = None
|
||||
@@ -2260,7 +2272,7 @@ class WanVideoSampler:
|
||||
#region model pred
|
||||
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
|
||||
control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None,
|
||||
add_cond=None, cache_state=None, context_window=None, multitalk_audio_embeds=None):
|
||||
add_cond=None, cache_state=None, context_window=None, multitalk_audio_embeds=None, fantasy_portrait_input=None):
|
||||
nonlocal transformer
|
||||
z = z.to(dtype)
|
||||
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])):
|
||||
@@ -2420,7 +2432,8 @@ class WanVideoSampler:
|
||||
"multitalk_audio": multitalk_audio_input if multitalk_audio_embedding is not None else None,
|
||||
"ref_target_masks": ref_target_masks if multitalk_audio_embedding is not None else None,
|
||||
"inner_t": [shot_len] if shot_len else None,
|
||||
"standin_input": standin_input
|
||||
"standin_input": standin_input,
|
||||
"fantasy_portrait_input": fantasy_portrait_input,
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
@@ -2910,6 +2923,10 @@ class WanVideoSampler:
|
||||
if fantasytalking_embeds is not None:
|
||||
partial_audio_proj = audio_proj[:, c]
|
||||
|
||||
if fantasy_portrait_input is not None:
|
||||
partial_fantasy_portrait_input = fantasy_portrait_input.copy()
|
||||
partial_fantasy_portrait_input["adapter_proj"] = fantasy_portrait_input["adapter_proj"][:, c]
|
||||
|
||||
partial_latent_model_input = latent_model_input[:, c]
|
||||
if latents_to_insert is not None and c[0] != 0:
|
||||
partial_latent_model_input[:, :1] = latents_to_insert
|
||||
@@ -2942,7 +2959,7 @@ class WanVideoSampler:
|
||||
cfg[idx], positive,
|
||||
text_embeds["negative_prompt_embeds"],
|
||||
partial_timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data,partial_audio_proj,
|
||||
partial_control_camera_latents, partial_add_cond, current_teacache, context_window=c)
|
||||
partial_control_camera_latents, partial_add_cond, current_teacache, context_window=c, fantasy_portrait_input=partial_fantasy_portrait_input)
|
||||
|
||||
if cache_args is not None:
|
||||
self.window_tracker.cache_states[window_id] = new_teacache
|
||||
@@ -3213,7 +3230,7 @@ class WanVideoSampler:
|
||||
text_embeds["prompt_embeds"],
|
||||
text_embeds["negative_prompt_embeds"],
|
||||
timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
|
||||
cache_state=self.cache_state)
|
||||
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input)
|
||||
|
||||
if latent_shift_loop:
|
||||
#reverse latent shift
|
||||
|
||||
+19
-4
@@ -1,4 +1,5 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import os, gc, uuid
|
||||
from .utils import log, apply_lora
|
||||
import numpy as np
|
||||
@@ -736,6 +737,7 @@ class WanVideoModelLoader:
|
||||
"vace_model": ("VACEPATH", {"default": None, "tooltip": "VACE model to use when not using model that has it included"}),
|
||||
"fantasytalking_model": ("FANTASYTALKINGMODEL", {"default": None, "tooltip": "FantasyTalking model https://github.com/Fantasy-AMAP"}),
|
||||
"multitalk_model": ("MULTITALKMODEL", {"default": None, "tooltip": "Multitalk model"}),
|
||||
"fantasyportrait_model": ("FANTASYPORTRAITMODEL", {"default": None, "tooltip": "FantasyPortrait model"}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -745,7 +747,8 @@ class WanVideoModelLoader:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def loadmodel(self, model, base_precision, load_device, quantization,
|
||||
compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, vram_management_args=None, vace_model=None, fantasytalking_model=None, multitalk_model=None):
|
||||
compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, vram_management_args=None, vace_model=None,
|
||||
fantasytalking_model=None, multitalk_model=None, fantasyportrait_model=None):
|
||||
assert not (vram_management_args is not None and block_swap_args is not None), "Can't use both block_swap_args and vram_management_args at the same time"
|
||||
|
||||
lora_low_mem_load = merge_loras = False
|
||||
@@ -970,7 +973,6 @@ class WanVideoModelLoader:
|
||||
#ReCamMaster
|
||||
if "blocks.0.cam_encoder.weight" in sd:
|
||||
log.info("ReCamMaster model detected, patching model...")
|
||||
import torch.nn as nn
|
||||
for block in transformer.blocks:
|
||||
block.cam_encoder = nn.Linear(12, dim)
|
||||
block.projector = nn.Linear(dim, dim)
|
||||
@@ -983,11 +985,25 @@ class WanVideoModelLoader:
|
||||
if fantasytalking_model is not None:
|
||||
log.info("FantasyTalking model detected, patching model...")
|
||||
context_dim = fantasytalking_model["sd"]["proj_model.proj.weight"].shape[0]
|
||||
import torch.nn as nn
|
||||
for block in transformer.blocks:
|
||||
block.cross_attn.k_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
block.cross_attn.v_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
sd.update(fantasytalking_model["sd"])
|
||||
|
||||
# FantasyPortrait https://github.com/Fantasy-AMAP/fantasy-portrait/
|
||||
if fantasyportrait_model is not None:
|
||||
log.info("FantasyPortrait model detected, patching model...")
|
||||
context_dim = fantasyportrait_model["sd"]["ip_adapter.blocks.0.cross_attn.ip_adapter_single_stream_k_proj.weight"].shape[1]
|
||||
|
||||
for block in transformer.blocks:
|
||||
block.cross_attn.ip_adapter_single_stream_k_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
block.cross_attn.ip_adapter_single_stream_v_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
ip_adapter_sd = {}
|
||||
for k, v in fantasyportrait_model["sd"].items():
|
||||
if k.startswith("ip_adapter."):
|
||||
ip_adapter_sd[k.replace("ip_adapter.", "")] = v
|
||||
sd.update(ip_adapter_sd)
|
||||
|
||||
if multitalk_model is not None:
|
||||
# init audio module
|
||||
from .multitalk.multitalk import SingleStreamMultiAttention
|
||||
@@ -1286,7 +1302,6 @@ class WanVideoModelLoader:
|
||||
for model in mm.current_loaded_models:
|
||||
if model._model() == patcher:
|
||||
mm.current_loaded_models.remove(model)
|
||||
|
||||
return (patcher,)
|
||||
|
||||
# class WanVideoSaveModel:
|
||||
|
||||
+63
-18
@@ -509,7 +509,9 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
self.attention_mode = attention_mode
|
||||
|
||||
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, audio_scale=1.0,
|
||||
num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy", inner_t=None, inner_c=None, cross_freqs=None):
|
||||
num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy",
|
||||
inner_t=None, inner_c=None, cross_freqs=None,
|
||||
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, **kwargs):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
# compute query
|
||||
q = self.norm_q(self.q(x),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d)
|
||||
@@ -536,19 +538,32 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
audio_q = q.view(b * num_latent_frames, -1, n, d)
|
||||
ip_key = self.k_proj(audio_proj).view(b * num_latent_frames, -1, n, d)
|
||||
ip_value = self.v_proj(audio_proj).view(b * num_latent_frames, -1, n, d)
|
||||
audio_x = attention(
|
||||
audio_q, ip_key, ip_value, attention_mode=self.attention_mode
|
||||
)
|
||||
audio_x = attention(audio_q, ip_key, ip_value, attention_mode=self.attention_mode)
|
||||
audio_x = audio_x.view(b, q.size(1), n, d).flatten(2)
|
||||
elif len(audio_proj.shape) == 3:
|
||||
ip_key = self.k_proj(audio_proj).view(b, -1, n, d)
|
||||
ip_value = self.v_proj(audio_proj).view(b, -1, n, d)
|
||||
audio_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode).flatten(2)
|
||||
|
||||
x = x + audio_x * audio_scale
|
||||
|
||||
x = self.o(x)
|
||||
return x
|
||||
# FantasyPortrait adapter attention
|
||||
if adapter_proj is not None:
|
||||
if len(adapter_proj.shape) == 4:
|
||||
adapter_q = q.view(b * num_latent_frames, -1, n, d)
|
||||
ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(b * num_latent_frames, -1, n, d)
|
||||
ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(b * num_latent_frames, -1, n, d)
|
||||
|
||||
adapter_x = attention(adapter_q, ip_key, ip_value, attention_mode=self.attention_mode)
|
||||
adapter_x = adapter_x.view(b, q.size(1), n, d)
|
||||
adapter_x = adapter_x.flatten(2)
|
||||
elif len(adapter_proj.shape) == 3:
|
||||
ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(b, -1, n, d)
|
||||
ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(b, -1, n, d)
|
||||
adapter_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode)
|
||||
adapter_x = adapter_x.flatten(2)
|
||||
x = x + adapter_x * ip_scale
|
||||
|
||||
return self.o(x)
|
||||
|
||||
|
||||
class WanI2VCrossAttention(WanSelfAttention):
|
||||
@@ -562,7 +577,7 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
|
||||
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None,
|
||||
audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy",
|
||||
**kwargs):
|
||||
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, **kwargs):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
@@ -595,19 +610,33 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
audio_q = q.view(b * num_latent_frames, -1, n, d)
|
||||
ip_key = self.k_proj(audio_proj).view(b * num_latent_frames, -1, n, d)
|
||||
ip_value = self.v_proj(audio_proj).view(b * num_latent_frames, -1, n, d)
|
||||
audio_x = attention(
|
||||
audio_q, ip_key, ip_value, attention_mode=self.attention_mode
|
||||
)
|
||||
|
||||
audio_x = attention(audio_q, ip_key, ip_value, attention_mode=self.attention_mode)
|
||||
audio_x = audio_x.view(b, q.size(1), n, d).flatten(2)
|
||||
elif len(audio_proj.shape) == 3:
|
||||
ip_key = self.k_proj(audio_proj).view(b, -1, n, d)
|
||||
ip_value = self.v_proj(audio_proj).view(b, -1, n, d)
|
||||
audio_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode).flatten(2)
|
||||
|
||||
x = x + audio_x * audio_scale
|
||||
|
||||
x = self.o(x)
|
||||
return x
|
||||
# FantasyPortrait adapter attention
|
||||
if adapter_proj is not None:
|
||||
if len(adapter_proj.shape) == 4:
|
||||
adapter_q = q.view(b * num_latent_frames, -1, n, d)
|
||||
ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(b * num_latent_frames, -1, n, d)
|
||||
ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(b * num_latent_frames, -1, n, d)
|
||||
|
||||
adapter_x = attention(adapter_q, ip_key, ip_value, attention_mode=self.attention_mode)
|
||||
adapter_x = adapter_x.view(b, q.size(1), n, d)
|
||||
adapter_x = adapter_x.flatten(2)
|
||||
elif len(adapter_proj.shape) == 3:
|
||||
ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(b, -1, n, d)
|
||||
ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(b, -1, n, d)
|
||||
adapter_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode)
|
||||
adapter_x = adapter_x.flatten(2)
|
||||
x = x + adapter_x * ip_scale
|
||||
|
||||
return self.o(x)
|
||||
|
||||
|
||||
WAN_CROSSATTENTION_CLASSES = {
|
||||
@@ -729,6 +758,8 @@ class WanAttentionBlock(nn.Module):
|
||||
x_ip=None,
|
||||
e_ip=None,
|
||||
freqs_ip=None,
|
||||
adapter_proj=None,
|
||||
ip_scale=1.0,
|
||||
):
|
||||
r"""
|
||||
Args:
|
||||
@@ -869,7 +900,8 @@ class WanAttentionBlock(nn.Module):
|
||||
else:
|
||||
x = self.cross_attn_ffn(x, context, grid_sizes, shift_mlp, scale_mlp, gate_mlp, clip_embed,
|
||||
audio_proj, audio_scale, num_latent_frames, nag_params, nag_context, is_uncond,
|
||||
multitalk_audio_embedding, x_ref_attn_map, human_num, inner_t, inner_c, cross_freqs)
|
||||
multitalk_audio_embedding, x_ref_attn_map, human_num, inner_t, inner_c, cross_freqs,
|
||||
adapter_proj=adapter_proj, ip_scale=ip_scale)
|
||||
else:
|
||||
if self.rope_func == "comfy_chunked":
|
||||
y = self.ffn_chunked(x, shift_mlp, scale_mlp)
|
||||
@@ -887,12 +919,14 @@ class WanAttentionBlock(nn.Module):
|
||||
|
||||
def cross_attn_ffn(self, x, context, grid_sizes, shift_mlp, scale_mlp, gate_mlp, clip_embed,
|
||||
audio_proj, audio_scale, num_latent_frames, nag_params,
|
||||
nag_context, is_uncond, multitalk_audio_embedding, x_ref_attn_map, human_num, inner_t, inner_c, cross_freqs):
|
||||
nag_context, is_uncond, multitalk_audio_embedding, x_ref_attn_map, human_num,
|
||||
inner_t, inner_c, cross_freqs, adapter_proj, ip_scale):
|
||||
|
||||
x = x + self.cross_attn(self.norm3(x), context, grid_sizes, clip_embed=clip_embed,
|
||||
audio_proj=audio_proj, audio_scale=audio_scale,
|
||||
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond,
|
||||
rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs)
|
||||
rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs,
|
||||
adapter_proj=adapter_proj, ip_scale=ip_scale)
|
||||
#multitalk
|
||||
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
|
||||
x_audio = self.audio_cross_attn(self.norm_x(x), encoder_hidden_states=multitalk_audio_embedding,
|
||||
@@ -1475,6 +1509,7 @@ class WanModel(torch.nn.Module):
|
||||
ref_target_masks=None,
|
||||
inner_t=None,
|
||||
standin_input=None,
|
||||
fantasy_portrait_input=None
|
||||
):
|
||||
r"""
|
||||
Forward pass through the diffusion model
|
||||
@@ -1497,9 +1532,17 @@ class WanModel(torch.nn.Module):
|
||||
List[Tensor]:
|
||||
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
|
||||
"""
|
||||
if is_uncond or current_step > 0:
|
||||
# Stand-In only used on first positive pass, then cached in kv_cache
|
||||
if is_uncond or current_step > 0:
|
||||
standin_input = None
|
||||
|
||||
# Fantasy Portrait
|
||||
adapter_proj = ip_scale = None
|
||||
if fantasy_portrait_input is not None:
|
||||
if fantasy_portrait_input['start_percent'] <= current_step_percentage <= fantasy_portrait_input['end_percent']:
|
||||
adapter_proj = fantasy_portrait_input.get("adapter_proj", None)
|
||||
ip_scale = fantasy_portrait_input.get("strength", 1.0)
|
||||
|
||||
if self.lora_scheduling_enabled:
|
||||
for name, submodule in self.named_modules():
|
||||
if isinstance(submodule, nn.Linear):
|
||||
@@ -1918,6 +1961,8 @@ class WanModel(torch.nn.Module):
|
||||
cross_freqs=self.cross_freqs if inner_t is not None else None,
|
||||
freqs_ip=freqs_ip if x_ip is not None else None,
|
||||
e_ip=e0_ip if x_ip is not None else None,
|
||||
adapter_proj=adapter_proj,
|
||||
ip_scale=ip_scale
|
||||
)
|
||||
|
||||
if vace_data is not None:
|
||||
|
||||
Reference in New Issue
Block a user