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:
kijai
2025-08-14 15:41:25 +03:00
parent f3d5f6b3ab
commit 24683f1dad
14 changed files with 4412 additions and 28 deletions
+12
View File
@@ -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
+506
View File
@@ -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
+117
View File
@@ -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
+320
View File
@@ -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
+149
View File
@@ -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
+343
View File
@@ -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.
+203
View File
@@ -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",
}
+406
View File
@@ -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
+23 -6
View File
@@ -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
View File
@@ -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
View File
@@ -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: