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 .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 .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 .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 .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 .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 .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:
|
try:
|
||||||
from .qwen.qwen import NODE_CLASS_MAPPINGS as QWEN_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as QWEN_NODE_DISPLAY_NAME_MAPPINGS
|
from .qwen.qwen import NODE_CLASS_MAPPINGS as QWEN_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as QWEN_NODE_DISPLAY_NAME_MAPPINGS
|
||||||
except ImportError:
|
except ImportError:
|
||||||
|
QWEN_NODE_CLASS_MAPPINGS = {}
|
||||||
|
QWEN_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||||
print("Qwen not available due to missing dependencies, probably transformers")
|
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:
|
try:
|
||||||
from .unianimate.nodes import NODE_CLASS_MAPPINGS as UNIANIMATE_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNIANIMATE_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
|
||||||
except ImportError:
|
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(UNIANIMATE_NODE_CLASS_MAPPINGS)
|
||||||
NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS)
|
NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS)
|
||||||
NODE_CLASS_MAPPINGS.update(FANTASYTALKING_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(FUN_CAMERA_NODE_CLASS_MAPPINGS)
|
||||||
NODE_CLASS_MAPPINGS.update(UNI3C_NODE_CLASS_MAPPINGS)
|
NODE_CLASS_MAPPINGS.update(UNI3C_NODE_CLASS_MAPPINGS)
|
||||||
NODE_CLASS_MAPPINGS.update(CONTROLNET_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(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
|
||||||
NODE_DISPLAY_NAME_MAPPINGS.update(SKYREELS_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(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(FUN_CAMERA_NODE_DISPLAY_NAME_MAPPINGS)
|
||||||
NODE_DISPLAY_NAME_MAPPINGS.update(UNI3C_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(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
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {"required": {
|
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}),
|
"shift": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 1000.0, "step": 0.01}),
|
||||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||||
"force_offload": ("BOOLEAN", {"default": True, "tooltip": "Moves the model to the offload device after sampling"}),
|
"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"}),
|
"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": {
|
"optional": {
|
||||||
@@ -1946,6 +1946,18 @@ class WanVideoSampler:
|
|||||||
|
|
||||||
shapes = [tuple(e.shape) for e in multitalk_audio_embedding]
|
shapes = [tuple(e.shape) for e in multitalk_audio_embedding]
|
||||||
log.info(f"Multitalk audio features shapes (per speaker): {shapes}")
|
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 Remover
|
||||||
minimax_latents = minimax_mask_latents = None
|
minimax_latents = minimax_mask_latents = None
|
||||||
@@ -2260,7 +2272,7 @@ class WanVideoSampler:
|
|||||||
#region model pred
|
#region model pred
|
||||||
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
|
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,
|
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
|
nonlocal transformer
|
||||||
z = z.to(dtype)
|
z = z.to(dtype)
|
||||||
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])):
|
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,
|
"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,
|
"ref_target_masks": ref_target_masks if multitalk_audio_embedding is not None else None,
|
||||||
"inner_t": [shot_len] if shot_len 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
|
batch_size = 1
|
||||||
@@ -2910,6 +2923,10 @@ class WanVideoSampler:
|
|||||||
if fantasytalking_embeds is not None:
|
if fantasytalking_embeds is not None:
|
||||||
partial_audio_proj = audio_proj[:, c]
|
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]
|
partial_latent_model_input = latent_model_input[:, c]
|
||||||
if latents_to_insert is not None and c[0] != 0:
|
if latents_to_insert is not None and c[0] != 0:
|
||||||
partial_latent_model_input[:, :1] = latents_to_insert
|
partial_latent_model_input[:, :1] = latents_to_insert
|
||||||
@@ -2942,7 +2959,7 @@ class WanVideoSampler:
|
|||||||
cfg[idx], positive,
|
cfg[idx], positive,
|
||||||
text_embeds["negative_prompt_embeds"],
|
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_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:
|
if cache_args is not None:
|
||||||
self.window_tracker.cache_states[window_id] = new_teacache
|
self.window_tracker.cache_states[window_id] = new_teacache
|
||||||
@@ -3213,7 +3230,7 @@ class WanVideoSampler:
|
|||||||
text_embeds["prompt_embeds"],
|
text_embeds["prompt_embeds"],
|
||||||
text_embeds["negative_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,
|
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:
|
if latent_shift_loop:
|
||||||
#reverse latent shift
|
#reverse latent shift
|
||||||
|
|||||||
+19
-4
@@ -1,4 +1,5 @@
|
|||||||
import torch
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
import os, gc, uuid
|
import os, gc, uuid
|
||||||
from .utils import log, apply_lora
|
from .utils import log, apply_lora
|
||||||
import numpy as np
|
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"}),
|
"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"}),
|
"fantasytalking_model": ("FANTASYTALKINGMODEL", {"default": None, "tooltip": "FantasyTalking model https://github.com/Fantasy-AMAP"}),
|
||||||
"multitalk_model": ("MULTITALKMODEL", {"default": None, "tooltip": "Multitalk model"}),
|
"multitalk_model": ("MULTITALKMODEL", {"default": None, "tooltip": "Multitalk model"}),
|
||||||
|
"fantasyportrait_model": ("FANTASYPORTRAITMODEL", {"default": None, "tooltip": "FantasyPortrait model"}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -745,7 +747,8 @@ class WanVideoModelLoader:
|
|||||||
CATEGORY = "WanVideoWrapper"
|
CATEGORY = "WanVideoWrapper"
|
||||||
|
|
||||||
def loadmodel(self, model, base_precision, load_device, quantization,
|
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"
|
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
|
lora_low_mem_load = merge_loras = False
|
||||||
@@ -970,7 +973,6 @@ class WanVideoModelLoader:
|
|||||||
#ReCamMaster
|
#ReCamMaster
|
||||||
if "blocks.0.cam_encoder.weight" in sd:
|
if "blocks.0.cam_encoder.weight" in sd:
|
||||||
log.info("ReCamMaster model detected, patching model...")
|
log.info("ReCamMaster model detected, patching model...")
|
||||||
import torch.nn as nn
|
|
||||||
for block in transformer.blocks:
|
for block in transformer.blocks:
|
||||||
block.cam_encoder = nn.Linear(12, dim)
|
block.cam_encoder = nn.Linear(12, dim)
|
||||||
block.projector = nn.Linear(dim, dim)
|
block.projector = nn.Linear(dim, dim)
|
||||||
@@ -983,11 +985,25 @@ class WanVideoModelLoader:
|
|||||||
if fantasytalking_model is not None:
|
if fantasytalking_model is not None:
|
||||||
log.info("FantasyTalking model detected, patching model...")
|
log.info("FantasyTalking model detected, patching model...")
|
||||||
context_dim = fantasytalking_model["sd"]["proj_model.proj.weight"].shape[0]
|
context_dim = fantasytalking_model["sd"]["proj_model.proj.weight"].shape[0]
|
||||||
import torch.nn as nn
|
|
||||||
for block in transformer.blocks:
|
for block in transformer.blocks:
|
||||||
block.cross_attn.k_proj = nn.Linear(context_dim, dim, bias=False)
|
block.cross_attn.k_proj = nn.Linear(context_dim, dim, bias=False)
|
||||||
block.cross_attn.v_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"])
|
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:
|
if multitalk_model is not None:
|
||||||
# init audio module
|
# init audio module
|
||||||
from .multitalk.multitalk import SingleStreamMultiAttention
|
from .multitalk.multitalk import SingleStreamMultiAttention
|
||||||
@@ -1286,7 +1302,6 @@ class WanVideoModelLoader:
|
|||||||
for model in mm.current_loaded_models:
|
for model in mm.current_loaded_models:
|
||||||
if model._model() == patcher:
|
if model._model() == patcher:
|
||||||
mm.current_loaded_models.remove(model)
|
mm.current_loaded_models.remove(model)
|
||||||
|
|
||||||
return (patcher,)
|
return (patcher,)
|
||||||
|
|
||||||
# class WanVideoSaveModel:
|
# class WanVideoSaveModel:
|
||||||
|
|||||||
+63
-18
@@ -509,7 +509,9 @@ class WanT2VCrossAttention(WanSelfAttention):
|
|||||||
self.attention_mode = attention_mode
|
self.attention_mode = attention_mode
|
||||||
|
|
||||||
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, audio_scale=1.0,
|
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
|
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||||
# compute query
|
# compute query
|
||||||
q = self.norm_q(self.q(x),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d)
|
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)
|
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_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)
|
ip_value = self.v_proj(audio_proj).view(b * num_latent_frames, -1, n, d)
|
||||||
audio_x = attention(
|
audio_x = attention(audio_q, ip_key, ip_value, attention_mode=self.attention_mode)
|
||||||
audio_q, ip_key, ip_value, attention_mode=self.attention_mode
|
|
||||||
)
|
|
||||||
audio_x = audio_x.view(b, q.size(1), n, d).flatten(2)
|
audio_x = audio_x.view(b, q.size(1), n, d).flatten(2)
|
||||||
elif len(audio_proj.shape) == 3:
|
elif len(audio_proj.shape) == 3:
|
||||||
ip_key = self.k_proj(audio_proj).view(b, -1, n, d)
|
ip_key = self.k_proj(audio_proj).view(b, -1, n, d)
|
||||||
ip_value = self.v_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)
|
audio_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode).flatten(2)
|
||||||
|
|
||||||
x = x + audio_x * audio_scale
|
x = x + audio_x * audio_scale
|
||||||
|
|
||||||
x = self.o(x)
|
# FantasyPortrait adapter attention
|
||||||
return x
|
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):
|
class WanI2VCrossAttention(WanSelfAttention):
|
||||||
@@ -562,7 +577,7 @@ class WanI2VCrossAttention(WanSelfAttention):
|
|||||||
|
|
||||||
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None,
|
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",
|
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"""
|
r"""
|
||||||
Args:
|
Args:
|
||||||
x(Tensor): Shape [B, L1, C]
|
x(Tensor): Shape [B, L1, C]
|
||||||
@@ -595,19 +610,33 @@ class WanI2VCrossAttention(WanSelfAttention):
|
|||||||
audio_q = q.view(b * num_latent_frames, -1, n, d)
|
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_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)
|
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)
|
audio_x = audio_x.view(b, q.size(1), n, d).flatten(2)
|
||||||
elif len(audio_proj.shape) == 3:
|
elif len(audio_proj.shape) == 3:
|
||||||
ip_key = self.k_proj(audio_proj).view(b, -1, n, d)
|
ip_key = self.k_proj(audio_proj).view(b, -1, n, d)
|
||||||
ip_value = self.v_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)
|
audio_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode).flatten(2)
|
||||||
|
|
||||||
x = x + audio_x * audio_scale
|
x = x + audio_x * audio_scale
|
||||||
|
|
||||||
x = self.o(x)
|
# FantasyPortrait adapter attention
|
||||||
return x
|
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 = {
|
WAN_CROSSATTENTION_CLASSES = {
|
||||||
@@ -729,6 +758,8 @@ class WanAttentionBlock(nn.Module):
|
|||||||
x_ip=None,
|
x_ip=None,
|
||||||
e_ip=None,
|
e_ip=None,
|
||||||
freqs_ip=None,
|
freqs_ip=None,
|
||||||
|
adapter_proj=None,
|
||||||
|
ip_scale=1.0,
|
||||||
):
|
):
|
||||||
r"""
|
r"""
|
||||||
Args:
|
Args:
|
||||||
@@ -869,7 +900,8 @@ class WanAttentionBlock(nn.Module):
|
|||||||
else:
|
else:
|
||||||
x = self.cross_attn_ffn(x, context, grid_sizes, shift_mlp, scale_mlp, gate_mlp, clip_embed,
|
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,
|
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:
|
else:
|
||||||
if self.rope_func == "comfy_chunked":
|
if self.rope_func == "comfy_chunked":
|
||||||
y = self.ffn_chunked(x, shift_mlp, scale_mlp)
|
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,
|
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,
|
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,
|
x = x + self.cross_attn(self.norm3(x), context, grid_sizes, clip_embed=clip_embed,
|
||||||
audio_proj=audio_proj, audio_scale=audio_scale,
|
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,
|
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
|
#multitalk
|
||||||
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
|
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,
|
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,
|
ref_target_masks=None,
|
||||||
inner_t=None,
|
inner_t=None,
|
||||||
standin_input=None,
|
standin_input=None,
|
||||||
|
fantasy_portrait_input=None
|
||||||
):
|
):
|
||||||
r"""
|
r"""
|
||||||
Forward pass through the diffusion model
|
Forward pass through the diffusion model
|
||||||
@@ -1497,9 +1532,17 @@ class WanModel(torch.nn.Module):
|
|||||||
List[Tensor]:
|
List[Tensor]:
|
||||||
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
|
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
|
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:
|
if self.lora_scheduling_enabled:
|
||||||
for name, submodule in self.named_modules():
|
for name, submodule in self.named_modules():
|
||||||
if isinstance(submodule, nn.Linear):
|
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,
|
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,
|
freqs_ip=freqs_ip if x_ip is not None else None,
|
||||||
e_ip=e0_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:
|
if vace_data is not None:
|
||||||
|
|||||||
Reference in New Issue
Block a user