first commit
This commit is contained in:
+18
@@ -0,0 +1,18 @@
|
||||
from .nodes import IP_LAP,LoadVideo,PreViewVideo,CombineAudioVideo
|
||||
WEB_DIRECTORY = "./web"
|
||||
# A dictionary that contains all nodes you want to export with their names
|
||||
# NOTE: names should be globally unique
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"IP_LAP": IP_LAP,
|
||||
"LoadVideo": LoadVideo,
|
||||
"PreViewVideo": PreViewVideo,
|
||||
"CombineAudioVideo": CombineAudioVideo
|
||||
}
|
||||
|
||||
# A dictionary that contains the friendly/humanly readable titles for the nodes
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"IP_LAP": "IP_LAP Node",
|
||||
"LoadVideo": "Video Loader",
|
||||
"PreViewVideo": "PreView Video",
|
||||
"CombineAudioVideo": "Combine Audio Video"
|
||||
}
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,197 @@
|
||||
"""MediaPipe solution drawing utils."""
|
||||
import math
|
||||
from typing import List, Mapping, Optional, Tuple, Union
|
||||
import cv2
|
||||
import dataclasses
|
||||
import numpy as np
|
||||
import tqdm
|
||||
from mediapipe.framework.formats import landmark_pb2
|
||||
_PRESENCE_THRESHOLD = 0.5
|
||||
_VISIBILITY_THRESHOLD = 0.5
|
||||
_BGR_CHANNELS = 3
|
||||
|
||||
WHITE_COLOR = (224, 224, 224)
|
||||
BLACK_COLOR = (0, 0, 0)
|
||||
RED_COLOR = (0, 0, 255)
|
||||
GREEN_COLOR = (0, 128, 0)
|
||||
BLUE_COLOR = (255, 0, 0)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class DrawingSpec:
|
||||
# Color for drawing the annotation. Default to the white color.
|
||||
color: Tuple[int, int, int] = WHITE_COLOR
|
||||
# Thickness for drawing the annotation. Default to 2 pixels.
|
||||
thickness: int = 2
|
||||
# Circle radius. Default to 2 pixels.
|
||||
circle_radius: int = 2
|
||||
|
||||
def _normalized_to_pixel_coordinates(
|
||||
normalized_x: float, normalized_y: float, image_width: int,
|
||||
image_height: int) -> Union[None, Tuple[int, int]]:
|
||||
"""Converts normalized value pair to pixel coordinates."""
|
||||
|
||||
# Checks if the float value is between 0 and 1.
|
||||
def is_valid_normalized_value(value: float) -> bool:
|
||||
return (value > 0 or math.isclose(0, value)) and (value < 1 or
|
||||
math.isclose(1, value))
|
||||
|
||||
if not (is_valid_normalized_value(normalized_x) and
|
||||
is_valid_normalized_value(normalized_y)):
|
||||
# TODO: Draw coordinates even if it's outside of the image bounds.
|
||||
return None
|
||||
x_px = min(math.floor(normalized_x * image_width), image_width - 1)
|
||||
y_px = min(math.floor(normalized_y * image_height), image_height - 1)
|
||||
return x_px, y_px
|
||||
|
||||
|
||||
FACEMESH_LIPS = frozenset([(61, 146), (146, 91), (91, 181), (181, 84), (84, 17),
|
||||
(17, 314), (314, 405), (405, 321), (321, 375),
|
||||
(375, 291), (61, 185), (185, 40), (40, 39), (39, 37),
|
||||
(37, 0), (0, 267),
|
||||
(267, 269), (269, 270), (270, 409), (409, 291),
|
||||
(78, 95), (95, 88), (88, 178), (178, 87), (87, 14),
|
||||
(14, 317), (317, 402), (402, 318), (318, 324),
|
||||
(324, 308), (78, 191), (191, 80), (80, 81), (81, 82),
|
||||
(82, 13), (13, 312), (312, 311), (311, 310),
|
||||
(310, 415), (415, 308)])
|
||||
|
||||
FACEMESH_LEFT_EYE = frozenset([(263, 249), (249, 390), (390, 373), (373, 374),
|
||||
(374, 380), (380, 381), (381, 382), (382, 362),
|
||||
(263, 466), (466, 388), (388, 387), (387, 386),
|
||||
(386, 385), (385, 384), (384, 398), (398, 362)])
|
||||
|
||||
FACEMESH_LEFT_IRIS = frozenset([(474, 475), (475, 476), (476, 477),
|
||||
(477, 474)])
|
||||
|
||||
FACEMESH_LEFT_EYEBROW = frozenset([(276, 283), (283, 282), (282, 295),
|
||||
(295, 285), (300, 293), (293, 334),
|
||||
(334, 296), (296, 336)])
|
||||
|
||||
FACEMESH_RIGHT_EYE = frozenset([(33, 7), (7, 163), (163, 144), (144, 145),
|
||||
(145, 153), (153, 154), (154, 155), (155, 133),
|
||||
(33, 246), (246, 161), (161, 160), (160, 159),
|
||||
(159, 158), (158, 157), (157, 173), (173, 133)])
|
||||
|
||||
FACEMESH_RIGHT_EYEBROW = frozenset([(46, 53), (53, 52), (52, 65), (65, 55),
|
||||
(70, 63), (63, 105), (105, 66), (66, 107)])
|
||||
|
||||
FACEMESH_RIGHT_IRIS = frozenset([(469, 470), (470, 471), (471, 472),
|
||||
(472, 469)])
|
||||
|
||||
FACEMESH_FACE_OVAL = frozenset([(389, 356), (356, 454),
|
||||
(454, 323), (323, 361), (361, 288), (288, 397),
|
||||
(397, 365), (365, 379), (379, 378), (378, 400),
|
||||
(400, 377), (377, 152), (152, 148), (148, 176),
|
||||
(176, 149), (149, 150), (150, 136), (136, 172),
|
||||
(172, 58), (58, 132), (132, 93), (93, 234),
|
||||
(234, 127), (127, 162)])
|
||||
#(10, 338), (338, 297), (297, 332), (332, 284),(284, 251), (251, 389) (162, 21), (21, 54),(54, 103), (103, 67), (67, 109), (109, 10)
|
||||
|
||||
FACEMESH_NOSE= frozenset([(168, 6),(6,197),(197,195),(195,5),(5,4),\
|
||||
(4,45),(45,220),(220,115),(115,48),\
|
||||
(4,275),(275,440),(440,344),(344,278),])
|
||||
FACEMESH_FULL = frozenset().union(*[
|
||||
FACEMESH_LIPS, FACEMESH_LEFT_EYE, FACEMESH_LEFT_EYEBROW, FACEMESH_RIGHT_EYE,
|
||||
FACEMESH_RIGHT_EYEBROW, FACEMESH_FACE_OVAL,FACEMESH_NOSE
|
||||
])
|
||||
connections=FACEMESH_FULL
|
||||
|
||||
def summary_landmark(edge_set):
|
||||
landmarks=set()
|
||||
for a,b in edge_set:
|
||||
landmarks.add(a)
|
||||
landmarks.add(b)
|
||||
return landmarks
|
||||
all_landmark_idx=summary_landmark(FACEMESH_FULL)
|
||||
pose_landmark_idx=\
|
||||
summary_landmark(FACEMESH_NOSE.union(*[FACEMESH_RIGHT_EYEBROW,FACEMESH_RIGHT_EYE,\
|
||||
FACEMESH_LEFT_EYE, FACEMESH_LEFT_EYEBROW,])).union([162,127,234,93,389,356,454,323])
|
||||
content_landmark_idx= all_landmark_idx - pose_landmark_idx
|
||||
|
||||
|
||||
def draw_landmarks(
|
||||
image: np.ndarray,
|
||||
landmark_list: List,
|
||||
connections: Optional[List[Tuple[int, int]]] = None,
|
||||
landmark_drawing_spec: Union[DrawingSpec,
|
||||
Mapping[int, DrawingSpec]] = DrawingSpec(
|
||||
color=RED_COLOR),
|
||||
connection_drawing_spec: Union[DrawingSpec,
|
||||
Mapping[Tuple[int, int],
|
||||
DrawingSpec]] = DrawingSpec()):
|
||||
"""Draws the landmarks and the connections on the image.
|
||||
|
||||
Args:
|
||||
image: A three channel BGR image represented as numpy ndarray.
|
||||
landmark_list: A normalized landmark list proto message to be annotated on
|
||||
the image.
|
||||
connections: A list of landmark index tuples that specifies how landmarks to
|
||||
be connected in the drawing.
|
||||
landmark_drawing_spec: Either a DrawingSpec object or a mapping from
|
||||
hand landmarks to the DrawingSpecs that specifies the landmarks' drawing
|
||||
settings such as color, line thickness, and circle radius.
|
||||
If this argument is explicitly set to None, no landmarks will be drawn.
|
||||
connection_drawing_spec: Either a DrawingSpec object or a mapping from
|
||||
hand connections to the DrawingSpecs that specifies the
|
||||
connections' drawing settings such as color and line thickness.
|
||||
If this argument is explicitly set to None, no landmark connections will
|
||||
be drawn.
|
||||
|
||||
Raises:
|
||||
ValueError: If one of the followings:
|
||||
a) If the input image is not three channel BGR.
|
||||
b) If any connetions contain invalid landmark index.
|
||||
"""
|
||||
if not landmark_list:
|
||||
return
|
||||
if image.shape[2] != _BGR_CHANNELS:
|
||||
raise ValueError('Input image must contain three channel bgr data.')
|
||||
image_rows, image_cols, _ = image.shape
|
||||
idx_to_coordinates = {}
|
||||
for landmark in landmark_list:
|
||||
# if ((landmark.HasField('visibility') and
|
||||
# landmark.visibility < _VISIBILITY_THRESHOLD) or
|
||||
# (landmark.HasField('presence') and
|
||||
# landmark.presence < _PRESENCE_THRESHOLD)):
|
||||
# continue
|
||||
idx=landmark.idx
|
||||
landmark_px = _normalized_to_pixel_coordinates(landmark.x, landmark.y,
|
||||
image_cols, image_rows)
|
||||
if landmark_px:
|
||||
idx_to_coordinates[idx] = landmark_px
|
||||
|
||||
if connections:
|
||||
num_landmarks = len(landmark_list)
|
||||
# Draws the connections if the start and end landmarks are both visible.
|
||||
for connection in connections:
|
||||
start_idx = connection[0]
|
||||
end_idx = connection[1]
|
||||
# if not (0 <= start_idx < num_landmarks and 0 <= end_idx < num_landmarks):
|
||||
# raise ValueError(f'Landmark index is out of range. Invalid connection '
|
||||
# f'from landmark #{start_idx} to landmark #{end_idx}.')
|
||||
if start_idx in idx_to_coordinates and end_idx in idx_to_coordinates:
|
||||
drawing_spec = connection_drawing_spec[connection] if isinstance(
|
||||
connection_drawing_spec, Mapping) else connection_drawing_spec
|
||||
# if start_idx in content_landmark and end_idx in content_landmark:
|
||||
cv2.line(image, idx_to_coordinates[start_idx],
|
||||
idx_to_coordinates[end_idx], drawing_spec.color,
|
||||
drawing_spec.thickness)
|
||||
return image
|
||||
# Draws landmark points after finishing the connection lines, which is
|
||||
# aesthetically better.
|
||||
# if landmark_drawing_spec:
|
||||
# for idx, landmark_px in idx_to_coordinates.items():
|
||||
# drawing_spec = landmark_drawing_spec[idx] if isinstance(
|
||||
# landmark_drawing_spec, Mapping) else landmark_drawing_spec
|
||||
# # White circle border
|
||||
# circle_border_radius = max(drawing_spec.circle_radius + 1,
|
||||
# int(drawing_spec.circle_radius * 1.2))
|
||||
# circle_border_radius=circle_border_radius*0.1
|
||||
# cv2.circle(image, landmark_px, circle_border_radius, WHITE_COLOR,
|
||||
# drawing_spec.thickness)
|
||||
# Fill color into the circle
|
||||
|
||||
# cv2.circle(image, landmark_px, 1,
|
||||
# drawing_spec.color, drawing_spec.thickness)
|
||||
# cv2.putText(image,str(idx),landmark_px,cv2.FONT_HERSHEY_SIMPLEX,0.5,(255,0,0),1,cv2.LINE_AA)
|
||||
@@ -0,0 +1,50 @@
|
||||
import cv2
|
||||
import numpy as np
|
||||
from typing import Any
|
||||
import mediapipe as mp
|
||||
from basicsr.utils.download_util import load_file_from_url
|
||||
|
||||
class FaceMask:
|
||||
def __init__(self) -> None:
|
||||
BaseOptions = mp.tasks.BaseOptions
|
||||
FaceLandmarker = mp.tasks.vision.FaceLandmarker
|
||||
FaceLandmarkerOptions = mp.tasks.vision.FaceLandmarkerOptions
|
||||
VisionRunningMode = mp.tasks.vision.RunningMode
|
||||
|
||||
face_landmarks_detector_path = load_file_from_url(url="https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/latest/face_landmarker.task",
|
||||
model_dir="weights",
|
||||
file_name="face_landmarker.task")
|
||||
options = FaceLandmarkerOptions(
|
||||
base_options=BaseOptions(model_asset_path=face_landmarks_detector_path),
|
||||
running_mode=VisionRunningMode.IMAGE)
|
||||
self.face_landmarks_detector = FaceLandmarker.create_from_options(options)
|
||||
|
||||
def __call__(self,image,*args: Any, **kwds: Any) -> Any:
|
||||
"""
|
||||
Calculate face mask from image. This is done by
|
||||
|
||||
Args:
|
||||
image: numpy array of an image
|
||||
Returns:
|
||||
A uint8 numpy array with the same height and width of the input image, containing a binary mask of the face in the image
|
||||
"""
|
||||
# initialize mask
|
||||
mask = np.zeros((image.shape[0], image.shape[1]), dtype=np.uint8)
|
||||
|
||||
# detect face landmarks
|
||||
mp_image = mp.Image(image_format=mp.ImageFormat.SRGB, data=image)
|
||||
detection = self.face_landmarks_detector.detect(mp_image)
|
||||
|
||||
if len(detection.face_landmarks) == 0:
|
||||
# no face detected - set mask to all of the image
|
||||
mask[:] = 1
|
||||
return mask
|
||||
|
||||
# extract landmarks coordinates
|
||||
face_coords = np.array([[lm.x * image.shape[1], lm.y * image.shape[0]] for lm in detection.face_landmarks[0]])
|
||||
|
||||
# calculate convex hull from face coordinates
|
||||
convex_hull = cv2.convexHull(face_coords.astype(np.float32))
|
||||
|
||||
# apply convex hull to mask
|
||||
return cv2.fillPoly(mask, pts=[convex_hull.squeeze().astype(np.int32)], color=1)
|
||||
@@ -0,0 +1,604 @@
|
||||
import os,cv2,torch,subprocess,platform
|
||||
import mediapipe as mp
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
from .draw_landmark import draw_landmarks
|
||||
import face_alignment
|
||||
from .face_mask import FaceMask
|
||||
from cuda_malloc import cuda_malloc_supported
|
||||
from .models import Landmark_generator as Landmark_transformer,Renderer,audio
|
||||
|
||||
NAME = "IP_LAP"
|
||||
|
||||
# the following is the index sequence for fical landmarks detected by mediapipe
|
||||
ori_sequence_idx = [162, 127, 234, 93, 132, 58, 172, 136, 150, 149, 176, 148, 152, 377, 400, 378, 379, 365, 397, 288,
|
||||
361, 323, 454, 356, 389, #
|
||||
70, 63, 105, 66, 107, 55, 65, 52, 53, 46, #
|
||||
336, 296, 334, 293, 300, 276, 283, 282, 295, 285, #
|
||||
168, 6, 197, 195, 5, #
|
||||
48, 115, 220, 45, 4, 275, 440, 344, 278, #
|
||||
33, 246, 161, 160, 159, 158, 157, 173, 133, 155, 154, 153, 145, 144, 163, 7, #
|
||||
362, 398, 384, 385, 386, 387, 388, 466, 263, 249, 390, 373, 374, 380, 381, 382, #
|
||||
61, 185, 40, 39, 37, 0, 267, 269, 270, 409, 291, 375, 321, 405, 314, 17, 84, 181, 91, 146, #
|
||||
78, 191, 80, 81, 82, 13, 312, 311, 310, 415, 308, 324, 318, 402, 317, 14, 87, 178, 88, 95]
|
||||
|
||||
# the following is the connections of landmarks for drawing sketch image
|
||||
FACEMESH_LIPS = frozenset([(61, 146), (146, 91), (91, 181), (181, 84), (84, 17),
|
||||
(17, 314), (314, 405), (405, 321), (321, 375),
|
||||
(375, 291), (61, 185), (185, 40), (40, 39), (39, 37),
|
||||
(37, 0), (0, 267),
|
||||
(267, 269), (269, 270), (270, 409), (409, 291),
|
||||
(78, 95), (95, 88), (88, 178), (178, 87), (87, 14),
|
||||
(14, 317), (317, 402), (402, 318), (318, 324),
|
||||
(324, 308), (78, 191), (191, 80), (80, 81), (81, 82),
|
||||
(82, 13), (13, 312), (312, 311), (311, 310),
|
||||
(310, 415), (415, 308)])
|
||||
FACEMESH_LEFT_EYE = frozenset([(263, 249), (249, 390), (390, 373), (373, 374),
|
||||
(374, 380), (380, 381), (381, 382), (382, 362),
|
||||
(263, 466), (466, 388), (388, 387), (387, 386),
|
||||
(386, 385), (385, 384), (384, 398), (398, 362)])
|
||||
FACEMESH_LEFT_EYEBROW = frozenset([(276, 283), (283, 282), (282, 295),
|
||||
(295, 285), (300, 293), (293, 334),
|
||||
(334, 296), (296, 336)])
|
||||
FACEMESH_RIGHT_EYE = frozenset([(33, 7), (7, 163), (163, 144), (144, 145),
|
||||
(145, 153), (153, 154), (154, 155), (155, 133),
|
||||
(33, 246), (246, 161), (161, 160), (160, 159),
|
||||
(159, 158), (158, 157), (157, 173), (173, 133)])
|
||||
FACEMESH_RIGHT_EYEBROW = frozenset([(46, 53), (53, 52), (52, 65), (65, 55),
|
||||
(70, 63), (63, 105), (105, 66), (66, 107)])
|
||||
FACEMESH_FACE_OVAL = frozenset([(389, 356), (356, 454),
|
||||
(454, 323), (323, 361), (361, 288), (288, 397),
|
||||
(397, 365), (365, 379), (379, 378), (378, 400),
|
||||
(400, 377), (377, 152), (152, 148), (148, 176),
|
||||
(176, 149), (149, 150), (150, 136), (136, 172),
|
||||
(172, 58), (58, 132), (132, 93), (93, 234),
|
||||
(234, 127), (127, 162)])
|
||||
FACEMESH_NOSE = frozenset([(168, 6), (6, 197), (197, 195), (195, 5), (5, 4),
|
||||
(4, 45), (45, 220), (220, 115), (115, 48),
|
||||
(4, 275), (275, 440), (440, 344), (344, 278), ])
|
||||
FACEMESH_CONNECTION = frozenset().union(*[
|
||||
FACEMESH_LIPS, FACEMESH_LEFT_EYE, FACEMESH_LEFT_EYEBROW, FACEMESH_RIGHT_EYE,
|
||||
FACEMESH_RIGHT_EYEBROW, FACEMESH_FACE_OVAL, FACEMESH_NOSE
|
||||
])
|
||||
|
||||
full_face_landmark_sequence = [*list(range(0, 4)), *list(range(21, 25)), *list(range(25, 91)), #upper-half face
|
||||
*list(range(4, 21)), # jaw
|
||||
*list(range(91, 131))] # mouth
|
||||
|
||||
class LandmarkDict(dict):# Makes a dictionary that behave like an object to represent each landmark
|
||||
def __init__(self, idx, x, y):
|
||||
self['idx'] = idx
|
||||
self['x'] = x
|
||||
self['y'] = y
|
||||
def __getattr__(self, name):
|
||||
try:
|
||||
return self[name]
|
||||
except:
|
||||
raise AttributeError(name)
|
||||
def __setattr__(self, name, value):
|
||||
self[name] = value
|
||||
|
||||
class IP_LAP_infer:
|
||||
|
||||
def __init__(self,T=5,Nl=15,ref_img_N=25,
|
||||
img_size=128,mel_step_size=16,
|
||||
face_det_batch_size=4,
|
||||
checkpoints_path=""):
|
||||
self.T = T
|
||||
self.Nl = Nl
|
||||
self.ref_img_N = ref_img_N
|
||||
self.img_size = img_size
|
||||
self.mel_step_size = mel_step_size
|
||||
self.face_det_batch_size = face_det_batch_size
|
||||
self.pads = [100,100,100,100]
|
||||
self.device = "cuda" if cuda_malloc_supported() else "cpu"
|
||||
|
||||
self.mp_face_mesh = mp.solutions.face_mesh
|
||||
self.drawing_spec = mp.solutions.drawing_utils.DrawingSpec(thickness=1, circle_radius=1)
|
||||
self.lip_index = [0, 17]
|
||||
self.all_landmarks_idx = self.summarize_landmark(FACEMESH_CONNECTION)
|
||||
self.pose_landmark_idx = \
|
||||
self.summarize_landmark(FACEMESH_NOSE.union(*[FACEMESH_RIGHT_EYEBROW, FACEMESH_RIGHT_EYE,
|
||||
FACEMESH_LEFT_EYE, FACEMESH_LEFT_EYEBROW, ])).union(
|
||||
[162, 127, 234, 93, 389, 356, 454, 323])
|
||||
# pose landmarks are landmarks of the upper-half face(eyes,nose,cheek) that represents the pose information
|
||||
|
||||
self.content_landmark_idx = self.all_landmarks_idx - self.pose_landmark_idx
|
||||
# content_landmark include landmarks of lip and jaw which are inferred from audio
|
||||
|
||||
|
||||
landmark_gen_checkpoint_path = os.path.join(checkpoints_path, "landmarkgenerator_checkpoint.pth")
|
||||
renderer_checkpoint_path = os.path.join(checkpoints_path, "renderer_checkpoint.pth")
|
||||
self.landmark_generator_model = self.load_model(
|
||||
model=Landmark_transformer(T=self.T, d_model=512, nlayers=4, nhead=4, dim_feedforward=1024, dropout=0.1),
|
||||
path=landmark_gen_checkpoint_path)
|
||||
self.renderer = self.load_model(model=Renderer(), path=renderer_checkpoint_path)
|
||||
|
||||
self.fa = face_alignment.FaceAlignment(face_alignment.LandmarksType.TWO_D, flip_input=False, device=self.device)
|
||||
|
||||
self.face_mask = FaceMask()
|
||||
|
||||
def __call__(self,video_file, audio_file, outfile):
|
||||
temp_dir = os.path.join(os.path.dirname(outfile), NAME)
|
||||
if not os.path.exists(temp_dir): os.makedirs(temp_dir, exist_ok=True)
|
||||
##(1) Reading input video frames ###
|
||||
print(f'[Step 1]Reading video frames ... from {video_file}', NAME)
|
||||
if not os.path.isfile(video_file):
|
||||
raise ValueError('the input video file does not exist')
|
||||
elif video_file.split('.')[1] in ['jpg', 'png', 'jpeg']: #if input a single image for testing
|
||||
ori_background_frames = [cv2.imread(video_file)]
|
||||
else:
|
||||
video_stream = cv2.VideoCapture(video_file)
|
||||
fps = video_stream.get(cv2.CAP_PROP_FPS)
|
||||
if fps != 25:
|
||||
print(" input video fps:", fps,',converting to 25fps...')
|
||||
tmp_file = '{}/temp_25fps.mp4'.format(temp_dir)
|
||||
if os.path.exists(tmp_file): os.remove(tmp_file)
|
||||
print(tmp_file)
|
||||
command = 'ffmpeg -y -i ' + video_file + f' -r 25 {tmp_file}'
|
||||
subprocess.call(command, shell=platform.system() != 'Windows',stdout=subprocess.PIPE, stderr=subprocess.STDOUT)
|
||||
video_file = '{}/temp_25fps.mp4'.format(temp_dir)
|
||||
video_stream.release()
|
||||
video_stream = cv2.VideoCapture(video_file)
|
||||
fps = video_stream.get(cv2.CAP_PROP_FPS)
|
||||
assert fps == 25
|
||||
|
||||
ori_background_frames = [] #input videos frames (includes background as well as face)
|
||||
frame_idx = 0
|
||||
while 1:
|
||||
still_reading, frame = video_stream.read()
|
||||
if not still_reading:
|
||||
video_stream.release()
|
||||
break
|
||||
ori_background_frames.append(frame)
|
||||
frame_idx = frame_idx + 1
|
||||
input_vid_len = len(ori_background_frames)
|
||||
|
||||
##(2) Extracting audio####
|
||||
print(f'[Step 2]Extracting audio ... from {audio_file}', NAME)
|
||||
if not audio_file.endswith('.wav'):
|
||||
command = 'ffmpeg -y -i {} -strict -2 {}'.format(audio_file, '{}/temp.wav'.format(temp_dir))
|
||||
subprocess.call(command, shell=platform.system() != 'Windows', stdout=subprocess.PIPE, stderr=subprocess.STDOUT)
|
||||
audio_file = '{}/temp.wav'.format(temp_dir)
|
||||
wav = audio.load_wav(audio_file, 16000)
|
||||
mel = audio.melspectrogram(wav) # (H,W) extract mel-spectrum
|
||||
##read audio mel into list###
|
||||
mel_chunks = [] # each mel chunk correspond to 5 video frames, used to generate one video frame
|
||||
mel_idx_multiplier = 80. / fps
|
||||
mel_chunk_idx = 0
|
||||
while 1:
|
||||
start_idx = int(mel_chunk_idx * mel_idx_multiplier)
|
||||
if start_idx + self.mel_step_size > len(mel[0]):
|
||||
break
|
||||
mel_chunks.append(mel[:, start_idx: start_idx + self.mel_step_size]) # mel for generate one video frame
|
||||
mel_chunk_idx += 1
|
||||
|
||||
print('[Step 3]detect facial using face detection tool', NAME)
|
||||
ori_face_frames, ori_face_coords = self.face_detect(ori_background_frames)
|
||||
# print(len(ori_face_frames))
|
||||
import gc; gc.collect(); torch.cuda.empty_cache()
|
||||
|
||||
##(3) detect facial landmarks using mediapipe tool
|
||||
print('[Step 4]detect facial landmarks using mediapipe tool', NAME)
|
||||
boxes = [] #bounding boxes of human face
|
||||
lip_dists = [] #lip dists
|
||||
#we define the lip dist(openness): distance between the midpoints of the upper lip and lower lip
|
||||
face_crop_results = []
|
||||
all_pose_landmarks, all_content_landmarks = [], [] #content landmarks include lip and jaw landmarks
|
||||
with self.mp_face_mesh.FaceMesh(static_image_mode=True, max_num_faces=1, refine_landmarks=True,
|
||||
min_detection_confidence=0) as face_mesh:
|
||||
# (1) get bounding boxes and lip dist
|
||||
for frame_idx, full_frame in tqdm(enumerate(ori_face_frames),total=input_vid_len,
|
||||
desc="get bounding boxes and lip dist"):
|
||||
h, w = full_frame.shape[0], full_frame.shape[1]
|
||||
results = face_mesh.process(cv2.cvtColor(full_frame, cv2.COLOR_BGR2RGB))
|
||||
if not results.multi_face_landmarks:
|
||||
raise NotImplementedError # not detect face
|
||||
face_landmarks = results.multi_face_landmarks[0]
|
||||
|
||||
## calculate the lip dist
|
||||
dx = face_landmarks.landmark[self.lip_index[0]].x - face_landmarks.landmark[self.lip_index[1]].x
|
||||
dy = face_landmarks.landmark[self.lip_index[0]].y - face_landmarks.landmark[self.lip_index[1]].y
|
||||
dist = np.linalg.norm((dx, dy))
|
||||
lip_dists.append((frame_idx, dist))
|
||||
|
||||
# (1)get the marginal landmarks to crop face
|
||||
x_min,x_max,y_min,y_max = 999,-999,999,-999
|
||||
for idx, landmark in enumerate(face_landmarks.landmark):
|
||||
if idx in self.all_landmarks_idx:
|
||||
if landmark.x < x_min:
|
||||
x_min = landmark.x
|
||||
if landmark.x > x_max:
|
||||
x_max = landmark.x
|
||||
if landmark.y < y_min:
|
||||
y_min = landmark.y
|
||||
if landmark.y > y_max:
|
||||
y_max = landmark.y
|
||||
##########plus some pixel to the marginal region##########
|
||||
#note:the landmarks coordinates returned by mediapipe range 0~1
|
||||
plus_pixel = 25
|
||||
x_min = max(x_min - plus_pixel / w, 0)
|
||||
x_max = min(x_max + plus_pixel / w, 1)
|
||||
|
||||
y_min = max(y_min - plus_pixel / h, 0)
|
||||
y_max = min(y_max + plus_pixel / h, 1)
|
||||
y1, y2, x1, x2 = int(y_min * h), int(y_max * h), int(x_min * w), int(x_max * w)
|
||||
boxes.append([y1, y2, x1, x2])
|
||||
boxes = np.array(boxes)
|
||||
|
||||
# (2)croppd face
|
||||
face_crop_results = [[image[y1:y2, x1:x2], (y1, y2, x1, x2)] \
|
||||
for image, (y1, y2, x1, x2) in zip(ori_face_frames, boxes)]
|
||||
|
||||
# (3)detect facial landmarks
|
||||
for frame_idx, full_frame in tqdm(enumerate(ori_face_frames),total=input_vid_len,
|
||||
desc="detect facial landmarks"):
|
||||
h, w = full_frame.shape[0], full_frame.shape[1]
|
||||
results = face_mesh.process(cv2.cvtColor(full_frame, cv2.COLOR_BGR2RGB))
|
||||
if not results.multi_face_landmarks:
|
||||
raise ValueError("not detect face in some frame!") # not detect
|
||||
face_landmarks = results.multi_face_landmarks[0]
|
||||
|
||||
|
||||
|
||||
pose_landmarks, content_landmarks = [], []
|
||||
for idx, landmark in enumerate(face_landmarks.landmark):
|
||||
if idx in self.pose_landmark_idx:
|
||||
pose_landmarks.append((idx, w * landmark.x, h * landmark.y))
|
||||
if idx in self.content_landmark_idx:
|
||||
content_landmarks.append((idx, w * landmark.x, h * landmark.y))
|
||||
|
||||
# normalize landmarks to 0~1
|
||||
y_min, y_max, x_min, x_max = face_crop_results[frame_idx][1] #bounding boxes
|
||||
pose_landmarks = [ \
|
||||
[idx, (x - x_min) / (x_max - x_min), (y - y_min) / (y_max - y_min)] for idx, x, y in pose_landmarks]
|
||||
content_landmarks = [ \
|
||||
[idx, (x - x_min) / (x_max - x_min), (y - y_min) / (y_max - y_min)] for idx, x, y in content_landmarks]
|
||||
all_pose_landmarks.append(pose_landmarks)
|
||||
all_content_landmarks.append(content_landmarks)
|
||||
|
||||
all_pose_landmarks = self.get_smoothened_landmarks(all_pose_landmarks, windows_T=1)
|
||||
all_content_landmarks=self.get_smoothened_landmarks(all_content_landmarks,windows_T=1)
|
||||
|
||||
##randomly select N_l reference landmarks for landmark transformer##
|
||||
print("randomly select N_l reference landmarks for landmark transformer", NAME)
|
||||
dists_sorted = sorted(lip_dists, key=lambda x: x[1])
|
||||
lip_dist_idx = np.asarray([idx for idx, dist in dists_sorted]) #the frame idxs sorted by lip openness
|
||||
|
||||
Nl_idxs = [lip_dist_idx[int(i)] for i in torch.linspace(0, input_vid_len - 1, steps=self.Nl)]
|
||||
Nl_pose_landmarks, Nl_content_landmarks = [], [] #Nl_pose + Nl_content=Nl reference landmarks
|
||||
for reference_idx in Nl_idxs:
|
||||
frame_pose_landmarks = all_pose_landmarks[reference_idx]
|
||||
frame_content_landmarks = all_content_landmarks[reference_idx]
|
||||
Nl_pose_landmarks.append(frame_pose_landmarks)
|
||||
Nl_content_landmarks.append(frame_content_landmarks)
|
||||
|
||||
Nl_pose = torch.zeros((self.Nl, 2, 74)) # 74 landmark
|
||||
Nl_content = torch.zeros((self.Nl, 2, 57)) # 57 landmark
|
||||
for idx in range(self.Nl):
|
||||
#arrange the landmark in a certain order, since the landmark index returned by mediapipe is is chaotic
|
||||
Nl_pose_landmarks[idx] = sorted(Nl_pose_landmarks[idx],
|
||||
key=lambda land_tuple: ori_sequence_idx.index(land_tuple[0]))
|
||||
Nl_content_landmarks[idx] = sorted(Nl_content_landmarks[idx],
|
||||
key=lambda land_tuple: ori_sequence_idx.index(land_tuple[0]))
|
||||
|
||||
Nl_pose[idx, 0, :] = torch.FloatTensor(
|
||||
[Nl_pose_landmarks[idx][i][1] for i in range(len(Nl_pose_landmarks[idx]))]) # x
|
||||
Nl_pose[idx, 1, :] = torch.FloatTensor(
|
||||
[Nl_pose_landmarks[idx][i][2] for i in range(len(Nl_pose_landmarks[idx]))]) # y
|
||||
Nl_content[idx, 0, :] = torch.FloatTensor(
|
||||
[Nl_content_landmarks[idx][i][1] for i in range(len(Nl_content_landmarks[idx]))]) # x
|
||||
Nl_content[idx, 1, :] = torch.FloatTensor(
|
||||
[Nl_content_landmarks[idx][i][2] for i in range(len(Nl_content_landmarks[idx]))]) # y
|
||||
Nl_content = Nl_content.unsqueeze(0) # (1,Nl, 2, 57)
|
||||
Nl_pose = Nl_pose.unsqueeze(0) # (1,Nl,2,74)
|
||||
|
||||
|
||||
##select reference images and draw sketches for rendering according to lip openness##
|
||||
print("select reference images and draw sketches for rendering according to lip openness", NAME)
|
||||
ref_img_idx = [int(lip_dist_idx[int(i)]) for i in torch.linspace(0, input_vid_len - 1, steps=self.ref_img_N)]
|
||||
ref_imgs = [face_crop_results[idx][0] for idx in ref_img_idx]
|
||||
## (N,H,W,3)
|
||||
ref_img_pose_landmarks, ref_img_content_landmarks = [], []
|
||||
for idx in ref_img_idx:
|
||||
ref_img_pose_landmarks.append(all_pose_landmarks[idx])
|
||||
ref_img_content_landmarks.append(all_content_landmarks[idx])
|
||||
|
||||
ref_img_pose = torch.zeros((self.ref_img_N, 2, 74)) # 74 landmark
|
||||
ref_img_content = torch.zeros((self.ref_img_N, 2, 57)) # 57 landmark
|
||||
|
||||
for idx in range(self.ref_img_N):
|
||||
ref_img_pose_landmarks[idx] = sorted(ref_img_pose_landmarks[idx],
|
||||
key=lambda land_tuple: ori_sequence_idx.index(land_tuple[0]))
|
||||
ref_img_content_landmarks[idx] = sorted(ref_img_content_landmarks[idx],
|
||||
key=lambda land_tuple: ori_sequence_idx.index(land_tuple[0]))
|
||||
ref_img_pose[idx, 0, :] = torch.FloatTensor(
|
||||
[ref_img_pose_landmarks[idx][i][1] for i in range(len(ref_img_pose_landmarks[idx]))]) # x
|
||||
ref_img_pose[idx, 1, :] = torch.FloatTensor(
|
||||
[ref_img_pose_landmarks[idx][i][2] for i in range(len(ref_img_pose_landmarks[idx]))]) # y
|
||||
|
||||
ref_img_content[idx, 0, :] = torch.FloatTensor(
|
||||
[ref_img_content_landmarks[idx][i][1] for i in range(len(ref_img_content_landmarks[idx]))]) # x
|
||||
ref_img_content[idx, 1, :] = torch.FloatTensor(
|
||||
[ref_img_content_landmarks[idx][i][2] for i in range(len(ref_img_content_landmarks[idx]))]) # y
|
||||
|
||||
ref_img_full_face_landmarks = torch.cat([ref_img_pose, ref_img_content], dim=2).cpu().numpy() # (N,2,131)
|
||||
ref_img_sketches = []
|
||||
for frame_idx in range(ref_img_full_face_landmarks.shape[0]): # N
|
||||
full_landmarks = ref_img_full_face_landmarks[frame_idx] # (2,131)
|
||||
h, w = ref_imgs[frame_idx].shape[0], ref_imgs[frame_idx].shape[1]
|
||||
drawn_sketech = np.zeros((int(h * self.img_size / min(h, w)), int(w * self.img_size / min(h, w)), 3))
|
||||
mediapipe_format_landmarks = [LandmarkDict(ori_sequence_idx[full_face_landmark_sequence[idx]], full_landmarks[0, idx],
|
||||
full_landmarks[1, idx]) for idx in range(full_landmarks.shape[1])]
|
||||
drawn_sketech = draw_landmarks(drawn_sketech, mediapipe_format_landmarks, connections=FACEMESH_CONNECTION,
|
||||
connection_drawing_spec=self.drawing_spec)
|
||||
drawn_sketech = cv2.resize(drawn_sketech, (self.img_size, self.img_size)) # (128, 128, 3)
|
||||
ref_img_sketches.append(drawn_sketech)
|
||||
ref_img_sketches = torch.FloatTensor(np.asarray(ref_img_sketches) / 255.0).cuda().unsqueeze(0).permute(0, 1, 4, 2, 3)
|
||||
# (1,N, 3, 128, 128)
|
||||
ref_imgs = [cv2.resize(face.copy(), (self.img_size, self.img_size)) for face in ref_imgs]
|
||||
ref_imgs = torch.FloatTensor(np.asarray(ref_imgs) / 255.0).unsqueeze(0).permute(0, 1, 4, 2, 3).cuda()
|
||||
# (1,N,3,H,W)
|
||||
|
||||
##prepare output video strame##
|
||||
frame_h, frame_w = ori_background_frames[0].shape[:-1]
|
||||
'''
|
||||
out_stream = cv2.VideoWriter('{}/result.avi'.format(temp_dir), cv2.VideoWriter_fourcc(*'DIVX'), fps,
|
||||
(frame_w, frame_h)) # +frame_h*3
|
||||
'''
|
||||
out_stream = cv2.VideoWriter(outfile, cv2.VideoWriter_fourcc(*'mp4v'), fps,
|
||||
(frame_w, frame_h)) # +frame_h*3
|
||||
|
||||
##generate final face image and output video##
|
||||
input_mel_chunks_len = len(mel_chunks)
|
||||
input_frame_sequence = torch.arange(input_vid_len).tolist()
|
||||
#the input template video may be shorter than audio
|
||||
#in this case we repeat the input template video as following
|
||||
num_of_repeat=input_mel_chunks_len//input_vid_len+1
|
||||
input_frame_sequence = input_frame_sequence + list(reversed(input_frame_sequence))
|
||||
input_frame_sequence=input_frame_sequence*((num_of_repeat+1)//2)
|
||||
file_num = 0
|
||||
for batch_idx, batch_start_idx in tqdm(enumerate(range(0, input_mel_chunks_len-2, 1)),
|
||||
total=len(range(0, input_mel_chunks_len-2, 1)), desc="[IP_LAP] [Step 5]Lipsync..."):
|
||||
T_input_frame, T_ori_face_coordinates = [], []
|
||||
#note: input_frame include background as well as face
|
||||
T_mel_batch, T_crop_face,T_pose_landmarks = [], [],[]
|
||||
|
||||
A_input_frame, A_ori_face_coordinates = [], []
|
||||
|
||||
# (1) for each batch of T frame, generate corresponding landmarks using landmark generator
|
||||
for mel_chunk_idx in range(batch_start_idx, batch_start_idx + self.T): # for each T frame
|
||||
# 1 input audio
|
||||
T_mel_batch.append(mel_chunks[max(0, mel_chunk_idx - 2)])
|
||||
|
||||
# 2.input face
|
||||
input_frame_idx = int(input_frame_sequence[mel_chunk_idx])
|
||||
face, coords = face_crop_results[input_frame_idx]
|
||||
T_crop_face.append(face)
|
||||
T_ori_face_coordinates.append((face, coords)) ##input face
|
||||
# 3.pose landmarks
|
||||
T_pose_landmarks.append(all_pose_landmarks[input_frame_idx])
|
||||
# 3.face background
|
||||
T_input_frame.append(ori_face_frames[input_frame_idx].copy())
|
||||
# 4.frame background
|
||||
A_ori_face_coordinates.append(ori_face_coords[input_frame_idx])
|
||||
A_input_frame.append(ori_background_frames[input_frame_idx].copy())
|
||||
|
||||
T_mels = torch.FloatTensor(np.asarray(T_mel_batch)).unsqueeze(1).unsqueeze(0) # 1,T,1,h,w
|
||||
#prepare pose landmarks
|
||||
T_pose = torch.zeros((self.T, 2, 74)) # 74 landmark
|
||||
for idx in range(self.T):
|
||||
T_pose_landmarks[idx] = sorted(T_pose_landmarks[idx],
|
||||
key=lambda land_tuple: ori_sequence_idx.index(land_tuple[0]))
|
||||
T_pose[idx, 0, :] = torch.FloatTensor(
|
||||
[T_pose_landmarks[idx][i][1] for i in range(len(T_pose_landmarks[idx]))]) # x
|
||||
T_pose[idx, 1, :] = torch.FloatTensor(
|
||||
[T_pose_landmarks[idx][i][2] for i in range(len(T_pose_landmarks[idx]))]) # y
|
||||
T_pose = T_pose.unsqueeze(0) # (1,T, 2,74)
|
||||
|
||||
#landmark generator inference
|
||||
Nl_pose, Nl_content = Nl_pose.cuda(), Nl_content.cuda() # (Nl,2,74) (Nl,2,57)
|
||||
T_mels, T_pose = T_mels.cuda(), T_pose.cuda()
|
||||
with torch.no_grad(): # require (1,T,1,hv,wv)(1,T,2,74)(1,T,2,57)
|
||||
predict_content = self.landmark_generator_model(T_mels, T_pose, Nl_pose, Nl_content) # (1*T,2,57)
|
||||
T_pose = torch.cat([T_pose[i] for i in range(T_pose.size(0))], dim=0) # (1*T,2,74)
|
||||
T_predict_full_landmarks = torch.cat([T_pose, predict_content], dim=2).cpu().numpy() # (1*T,2,131)
|
||||
|
||||
#1.draw target sketch
|
||||
T_target_sketches = []
|
||||
for frame_idx in range(self.T):
|
||||
full_landmarks = T_predict_full_landmarks[frame_idx] # (2,131)
|
||||
h, w = T_crop_face[frame_idx].shape[0], T_crop_face[frame_idx].shape[1]
|
||||
drawn_sketech = np.zeros((int(h * self.img_size / min(h, w)), int(w * self.img_size / min(h, w)), 3))
|
||||
mediapipe_format_landmarks = [LandmarkDict(ori_sequence_idx[full_face_landmark_sequence[idx]]
|
||||
, full_landmarks[0, idx], full_landmarks[1, idx]) for idx in
|
||||
range(full_landmarks.shape[1])]
|
||||
drawn_sketech = draw_landmarks(drawn_sketech, mediapipe_format_landmarks, connections=FACEMESH_CONNECTION,
|
||||
connection_drawing_spec=self.drawing_spec)
|
||||
drawn_sketech = cv2.resize(drawn_sketech, (self.img_size, self.img_size)) # (128, 128, 3)
|
||||
if frame_idx == 2:
|
||||
show_sketch = cv2.resize(drawn_sketech, (frame_w, frame_h)).astype(np.uint8)
|
||||
T_target_sketches.append(torch.FloatTensor(drawn_sketech) / 255)
|
||||
T_target_sketches = torch.stack(T_target_sketches, dim=0).permute(0, 3, 1, 2) # (T,3,128, 128)
|
||||
target_sketches = T_target_sketches.unsqueeze(0).cuda() # (1,T,3,128, 128)
|
||||
|
||||
# 2.lower-half masked face
|
||||
ori_face_img = torch.FloatTensor(cv2.resize(T_crop_face[2], (self.img_size, self.img_size)) / 255).permute(2, 0, 1).unsqueeze(
|
||||
0).unsqueeze(0).cuda() #(1,1,3,H, W)
|
||||
|
||||
# 3. render the full face
|
||||
# require (1,1,3,H,W) (1,T,3,H,W) (1,N,3,H,W) (1,N,3,H,W) (1,1,1,h,w)
|
||||
# return (1,3,H,W)
|
||||
with torch.no_grad():
|
||||
generated_face, _, _, _ = self.renderer(ori_face_img, target_sketches, ref_imgs, ref_img_sketches,
|
||||
T_mels[:, 2].unsqueeze(0)) # T=1
|
||||
gen_face = (generated_face.squeeze(0).permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8) # (H,W,3)
|
||||
|
||||
# 4. paste each generated face
|
||||
y1, y2, x1, x2 = T_ori_face_coordinates[2][1] # coordinates of face bounding box
|
||||
original_background = T_input_frame[2].copy()
|
||||
T_input_frame[2][y1:y2, x1:x2] = cv2.resize(gen_face,(x2 - x1, y2 - y1)) #resize and paste generated face
|
||||
# 5. post-process
|
||||
full_face = self.merge_face_contour_only(original_background, T_input_frame[2], T_ori_face_coordinates[2][1],self.fa) #(H,W,3)
|
||||
# 6.output
|
||||
# full = np.concatenate([show_sketch, full], axis=1)
|
||||
# print(f"full_face.shape{full_face.shape}")
|
||||
ori_x1, ori_y1, ori_x2, ori_y2 = A_ori_face_coordinates[2]
|
||||
# print(ori_face_coords[file_num])
|
||||
full_frame = A_input_frame[2]
|
||||
#full_mask = np.zeros_like(full_frame)
|
||||
# print(full_frame.shape)
|
||||
if ori_x1 != -1:
|
||||
p = cv2.resize(full_face.astype(np.uint8), (ori_x2 - ori_x1, ori_y2 - ori_y1))
|
||||
# print(p.shape)
|
||||
full_frame[ori_y1:ori_y2, ori_x1:ori_x2] = p
|
||||
# height, width = full_frame.shape[:2]
|
||||
# img = self.Laplacian_Pyramid_Blending_with_mask(full_frame, ori_background_frames[file_num], full_mask[:, :, 0], 6)
|
||||
# pp = np.uint8(cv2.resize(np.clip(img, 0 ,255), (width, height)))
|
||||
mask = self.face_mask(p)
|
||||
full_frame[ori_y1:ori_y2, ori_x1:ori_x2] = full_frame[ori_y1:ori_y2, ori_x1:ori_x2] * (1 - mask[..., None]) + p * mask[..., None]
|
||||
full = full_frame.copy()
|
||||
|
||||
out_stream.write(full)
|
||||
|
||||
try:
|
||||
# cv2.imwrite(temp_frame_paths[batch_idx+2],full)
|
||||
file_num += 1
|
||||
except:
|
||||
pass
|
||||
|
||||
if batch_idx == 0:
|
||||
out_stream.write(full)
|
||||
out_stream.write(full)
|
||||
# cv2.imwrite(temp_frame_paths[batch_idx],full)
|
||||
# cv2.imwrite(temp_frame_paths[batch_idx+1],full)
|
||||
|
||||
out_stream.release()
|
||||
# command = 'ffmpeg -y -i {} -i {} -strict -2 -q:v 1 {}'.format(voice_file, '{}/result.avi'.format(temp_dir), outfile)
|
||||
# subprocess.call(command, shell=True, stdout=subprocess.PIPE, stderr=subprocess.STDOUT)
|
||||
print(f"succeed output results to:{outfile}", NAME)
|
||||
|
||||
|
||||
def face_detect(self, images):
|
||||
|
||||
batch_size = self.face_det_batch_size
|
||||
|
||||
while 1:
|
||||
predictions = []
|
||||
try:
|
||||
for i in tqdm(range(0, len(images), batch_size)):
|
||||
imgs = np.array(images[i:i + batch_size])
|
||||
imgs_numpy = imgs.transpose(0, 3, 1, 2)
|
||||
image_batch = torch.from_numpy(imgs_numpy.copy())
|
||||
_, _, bboxes =self.fa.get_landmarks_from_batch(image_batch,return_bboxes=True)
|
||||
predictions.extend(bboxes)
|
||||
except RuntimeError:
|
||||
if batch_size == 1:
|
||||
raise RuntimeError('Image too big to run face detection on GPU. Please use the --resize_factor argument')
|
||||
batch_size //= 2
|
||||
print('Recovering from OOM error; New batch size: {}'.format(batch_size))
|
||||
continue
|
||||
break
|
||||
|
||||
results = []
|
||||
pady1, pady2, padx1, padx2 = self.pads
|
||||
for rect, image in zip(predictions, images):
|
||||
|
||||
if rect is None:
|
||||
# cv2.imwrite('temp/faulty_frame.jpg', image) # check this frame where the face was not detected.
|
||||
# results.append([-1,-1,-1,-1])
|
||||
raise ValueError('Face not detected! Ensure the video contains a face in all the frames.')
|
||||
else:
|
||||
rect = rect[0]
|
||||
rect = np.clip(rect, 0, None)
|
||||
x1_0, y1_0, x2_0, y2_0 = map(int, rect[:-1])
|
||||
|
||||
y1 = max(0, y1_0 - pady1)
|
||||
y2 = min(image.shape[0], y2_0 + pady2)
|
||||
x1 = max(0, x1_0 - padx1)
|
||||
x2 = min(image.shape[1], x2_0 + padx2)
|
||||
|
||||
results.append([x1, y1, x2, y2])
|
||||
|
||||
boxes = np.array(results)
|
||||
|
||||
faces = [image[y1: y2, x1:x2] for image, (x1, y1, x2, y2) in zip(images, boxes)]
|
||||
return faces, boxes
|
||||
|
||||
|
||||
def merge_face_contour_only(self,src_frame, generated_frame, face_region_coord, fa): #function used in post-process
|
||||
"""Merge the face from generated_frame into src_frame
|
||||
"""
|
||||
input_img = src_frame
|
||||
y1, y2, x1, x2 = 0, 0, 0, 0
|
||||
if face_region_coord is not None:
|
||||
y1, y2, x1, x2 = face_region_coord
|
||||
input_img = src_frame[y1:y2, x1:x2]
|
||||
### 1) Detect the facial landmarks
|
||||
try:
|
||||
preds = fa.get_landmarks(input_img)[0] # 68x2
|
||||
except:
|
||||
preds = np.int64(-1 * np.ones((68,2)))
|
||||
if face_region_coord is not None:
|
||||
preds += np.array([x1, y1])
|
||||
lm_pts = preds.astype(int)
|
||||
contour_idx = list(range(0, 17)) + list(range(17, 27))[::-1]
|
||||
contour_pts = lm_pts[contour_idx]
|
||||
### 2) Make the landmark region mark image
|
||||
mask_img = np.zeros((src_frame.shape[0], src_frame.shape[1], 1), np.uint8)
|
||||
cv2.fillConvexPoly(mask_img, contour_pts, 255)
|
||||
### 3) Do swap
|
||||
img = self.swap_masked_region(src_frame, generated_frame, mask=mask_img)
|
||||
return img
|
||||
|
||||
def swap_masked_region(self,target_img, src_img, mask): #function used in post-process
|
||||
"""From src_img crop masked region to replace corresponding masked region
|
||||
in target_img
|
||||
""" # swap_masked_region(src_frame, generated_frame, mask=mask_img)
|
||||
mask_img = cv2.GaussianBlur(mask, (21, 21), 11)
|
||||
mask1 = mask_img / 255
|
||||
mask1 = np.tile(np.expand_dims(mask1, axis=2), (1, 1, 3))
|
||||
img = src_img * mask1 + target_img * (1 - mask1)
|
||||
return img.astype(np.uint8)
|
||||
|
||||
# smooth landmarks
|
||||
def get_smoothened_landmarks(self,all_landmarks, windows_T=1):
|
||||
for i in range(len(all_landmarks)): # frame i
|
||||
if i + windows_T > len(all_landmarks):
|
||||
window = all_landmarks[len(all_landmarks) - windows_T:]
|
||||
else:
|
||||
window = all_landmarks[i: i + windows_T]
|
||||
#####
|
||||
for j in range(len(all_landmarks[i])): # landmark j
|
||||
all_landmarks[i][j][1] = np.mean([frame_landmarks[j][1] for frame_landmarks in window]) # x
|
||||
all_landmarks[i][j][2] = np.mean([frame_landmarks[j][2] for frame_landmarks in window]) # y
|
||||
return all_landmarks
|
||||
|
||||
def load_model(self, model, path):
|
||||
print("Load checkpoint from: {}".format(path))
|
||||
checkpoint = self._load(path)
|
||||
s = checkpoint["state_dict"]
|
||||
new_s = {}
|
||||
for k, v in s.items():
|
||||
if k[:6] == 'module':
|
||||
new_k=k.replace('module.', '', 1)
|
||||
else:
|
||||
new_k =k
|
||||
new_s[new_k] = v
|
||||
model.load_state_dict(new_s)
|
||||
model = model.to(self.device)
|
||||
return model.eval()
|
||||
|
||||
def _load(self,checkpoint_path):
|
||||
if self.device == 'cuda':
|
||||
checkpoint = torch.load(checkpoint_path)
|
||||
else:
|
||||
checkpoint = torch.load(checkpoint_path, map_location=lambda storage, loc: storage)
|
||||
return checkpoint
|
||||
|
||||
def summarize_landmark(self, edge_set): # summarize all ficial landmarks used to construct edge
|
||||
landmarks = set()
|
||||
for a, b in edge_set:
|
||||
landmarks.add(a)
|
||||
landmarks.add(b)
|
||||
return landmarks
|
||||
@@ -0,0 +1,4 @@
|
||||
from . import audio
|
||||
from .landmark_generator import Landmark_generator
|
||||
from .video_renderer import Renderer
|
||||
from .pix2pixHD_disc import define_D
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,237 @@
|
||||
import librosa
|
||||
import librosa.filters
|
||||
import numpy as np
|
||||
from scipy import signal
|
||||
from scipy.io import wavfile
|
||||
import lws
|
||||
|
||||
class HParams:
|
||||
def __init__(self, **kwargs):
|
||||
self.data = {}
|
||||
|
||||
for key, value in kwargs.items():
|
||||
self.data[key] = value
|
||||
|
||||
def __getattr__(self, key):
|
||||
if key not in self.data:
|
||||
raise AttributeError("'HParams' object has no attribute %s" % key)
|
||||
return self.data[key]
|
||||
|
||||
def set_hparam(self, key, value):
|
||||
self.data[key] = value
|
||||
|
||||
|
||||
# Default hyperparameters
|
||||
hp = HParams(
|
||||
num_mels=80, # Number of mel-spectrogram channels and local conditioning dimensionality
|
||||
# network
|
||||
rescale=True, # Whether to rescale audio prior to preprocessing
|
||||
rescaling_max=0.9, # Rescaling value
|
||||
|
||||
# Use LWS (https://github.com/Jonathan-LeRoux/lws) for STFT and phase reconstruction
|
||||
# It"s preferred to set True to use with https://github.com/r9y9/wavenet_vocoder
|
||||
# Does not work if n_ffit is not multiple of hop_size!!
|
||||
use_lws=False,
|
||||
|
||||
n_fft=800, # Extra window size is filled with 0 paddings to match this parameter
|
||||
hop_size=200, # For 16000Hz, 200 = 12.5 ms (0.0125 * sample_rate)
|
||||
win_size=800, # For 16000Hz, 800 = 50 ms (If None, win_size = n_fft) (0.05 * sample_rate)
|
||||
sample_rate=16000, # 16000Hz (corresponding to librispeech) (sox --i <filename>)
|
||||
|
||||
frame_shift_ms=None, # Can replace hop_size parameter. (Recommended: 12.5)
|
||||
|
||||
# Mel and Linear spectrograms normalization/scaling and clipping
|
||||
signal_normalization=True,
|
||||
# Whether to normalize mel spectrograms to some predefined range (following below parameters)
|
||||
allow_clipping_in_normalization=True, # Only relevant if mel_normalization = True
|
||||
symmetric_mels=True,
|
||||
# Whether to scale the data to be symmetric around 0. (Also multiplies the output range by 2,
|
||||
# faster and cleaner convergence)
|
||||
max_abs_value=4.,
|
||||
# max absolute value of data. If symmetric, data will be [-max, max] else [0, max] (Must not
|
||||
# be too big to avoid gradient explosion,
|
||||
# not too small for fast convergence)
|
||||
# Contribution by @begeekmyfriend
|
||||
# Spectrogram Pre-Emphasis (Lfilter: Reduce spectrogram noise and helps model certitude
|
||||
# levels. Also allows for better G&L phase reconstruction)
|
||||
preemphasize=True, # whether to apply filter
|
||||
preemphasis=0.97, # filter coefficient.
|
||||
|
||||
# Limits
|
||||
min_level_db=-100,
|
||||
ref_level_db=20,
|
||||
fmin=55,
|
||||
# Set this to 55 if your speaker is male! if female, 95 should help taking off noise. (To
|
||||
# test depending on dataset. Pitch info: male~[65, 260], female~[100, 525])
|
||||
fmax=7600, # To be increased/reduced depending on data.
|
||||
|
||||
###################### Our training parameters #################################
|
||||
img_size=288,
|
||||
fps=25,
|
||||
|
||||
batch_size=8,
|
||||
initial_learning_rate=1e-4,
|
||||
nepochs=200000000000000000,
|
||||
### ctrl + c, stop whenever eval loss is consistently greater than train loss for ~10 epochs
|
||||
num_workers=4,
|
||||
checkpoint_interval=6000,
|
||||
eval_interval=6000,
|
||||
save_optimizer_state=True,
|
||||
|
||||
syncnet_wt=0.0, # is initially zero, will be set automatically to 0.03 later. Leads to faster convergence.
|
||||
syncnet_batch_size=128,
|
||||
syncnet_lr=1e-4,
|
||||
syncnet_eval_interval=4500,
|
||||
syncnet_checkpoint_interval=4500,
|
||||
|
||||
disc_wt=0.07,
|
||||
disc_initial_learning_rate=1e-4,
|
||||
)
|
||||
|
||||
|
||||
def load_wav(path, sr):
|
||||
return librosa.core.load(path, sr=sr)[0]
|
||||
|
||||
|
||||
def save_wav(wav, path, sr):
|
||||
wav *= 32767 / max(0.01, np.max(np.abs(wav)))
|
||||
# proposed by @dsmiller
|
||||
wavfile.write(path, sr, wav.astype(np.int16))
|
||||
|
||||
|
||||
def save_wavenet_wav(wav, path, sr):
|
||||
librosa.output.write_wav(path, wav, sr=sr)
|
||||
|
||||
|
||||
def preemphasis(wav, k, preemphasize=True):
|
||||
if preemphasize:
|
||||
return signal.lfilter([1, -k], [1], wav)
|
||||
return wav
|
||||
|
||||
|
||||
def inv_preemphasis(wav, k, inv_preemphasize=True):
|
||||
if inv_preemphasize:
|
||||
return signal.lfilter([1], [1, -k], wav)
|
||||
return wav
|
||||
|
||||
|
||||
def get_hop_size():
|
||||
hop_size = hp.hop_size
|
||||
if hop_size is None:
|
||||
assert hp.frame_shift_ms is not None
|
||||
hop_size = int(hp.frame_shift_ms / 1000 * hp.sample_rate)
|
||||
return hop_size
|
||||
|
||||
|
||||
def linearspectrogram(wav):
|
||||
D = _stft(preemphasis(wav, hp.preemphasis, hp.preemphasize))
|
||||
S = _amp_to_db(np.abs(D)) - hp.ref_level_db
|
||||
|
||||
if hp.signal_normalization:
|
||||
return _normalize(S)
|
||||
return S
|
||||
|
||||
|
||||
def melspectrogram(wav):
|
||||
D = _stft(preemphasis(wav, hp.preemphasis, hp.preemphasize))
|
||||
S = _amp_to_db(_linear_to_mel(np.abs(D))) - hp.ref_level_db
|
||||
|
||||
if hp.signal_normalization:
|
||||
return _normalize(S)
|
||||
return S
|
||||
|
||||
|
||||
def _lws_processor():
|
||||
return lws.lws(hp.n_fft, get_hop_size(), fftsize=hp.win_size, mode="speech")
|
||||
|
||||
|
||||
def _stft(y):
|
||||
if hp.use_lws:
|
||||
return _lws_processor(hp).stft(y).T
|
||||
else:
|
||||
return librosa.stft(y=y, n_fft=hp.n_fft, hop_length=get_hop_size(), win_length=hp.win_size)
|
||||
|
||||
|
||||
##########################################################
|
||||
# Those are only correct when using lws!!! (This was messing with Wavenet quality for a long time!)
|
||||
def num_frames(length, fsize, fshift):
|
||||
"""Compute number of time frames of spectrogram
|
||||
"""
|
||||
pad = (fsize - fshift)
|
||||
if length % fshift == 0:
|
||||
M = (length + pad * 2 - fsize) // fshift + 1
|
||||
else:
|
||||
M = (length + pad * 2 - fsize) // fshift + 2
|
||||
return M
|
||||
|
||||
|
||||
def pad_lr(x, fsize, fshift):
|
||||
"""Compute left and right padding
|
||||
"""
|
||||
M = num_frames(len(x), fsize, fshift)
|
||||
pad = (fsize - fshift)
|
||||
T = len(x) + 2 * pad
|
||||
r = (M - 1) * fshift + fsize - T
|
||||
return pad, pad + r
|
||||
|
||||
|
||||
##########################################################
|
||||
# Librosa correct padding
|
||||
def librosa_pad_lr(x, fsize, fshift):
|
||||
return 0, (x.shape[0] // fshift + 1) * fshift - x.shape[0]
|
||||
|
||||
|
||||
# Conversions
|
||||
_mel_basis = None
|
||||
|
||||
|
||||
def _linear_to_mel(spectogram):
|
||||
global _mel_basis
|
||||
if _mel_basis is None:
|
||||
_mel_basis = _build_mel_basis()
|
||||
return np.dot(_mel_basis, spectogram)
|
||||
|
||||
|
||||
def _build_mel_basis():
|
||||
assert hp.fmax <= hp.sample_rate // 2
|
||||
return librosa.filters.mel(sr=hp.sample_rate, n_fft=hp.n_fft, n_mels=hp.num_mels,
|
||||
fmin=hp.fmin, fmax=hp.fmax)
|
||||
|
||||
|
||||
def _amp_to_db(x):
|
||||
min_level = np.exp(hp.min_level_db / 20 * np.log(10))
|
||||
return 20 * np.log10(np.maximum(min_level, x))
|
||||
|
||||
|
||||
def _db_to_amp(x):
|
||||
return np.power(10.0, (x) * 0.05)
|
||||
|
||||
|
||||
def _normalize(S):
|
||||
if hp.allow_clipping_in_normalization:
|
||||
if hp.symmetric_mels:
|
||||
return np.clip((2 * hp.max_abs_value) * ((S - hp.min_level_db) / (-hp.min_level_db)) - hp.max_abs_value,
|
||||
-hp.max_abs_value, hp.max_abs_value)
|
||||
else:
|
||||
return np.clip(hp.max_abs_value * ((S - hp.min_level_db) / (-hp.min_level_db)), 0, hp.max_abs_value)
|
||||
|
||||
assert S.max() <= 0 and S.min() - hp.min_level_db >= 0
|
||||
if hp.symmetric_mels:
|
||||
return (2 * hp.max_abs_value) * ((S - hp.min_level_db) / (-hp.min_level_db)) - hp.max_abs_value
|
||||
else:
|
||||
return hp.max_abs_value * ((S - hp.min_level_db) / (-hp.min_level_db))
|
||||
|
||||
|
||||
def _denormalize(D):
|
||||
if hp.allow_clipping_in_normalization:
|
||||
if hp.symmetric_mels:
|
||||
return (((np.clip(D, -hp.max_abs_value,
|
||||
hp.max_abs_value) + hp.max_abs_value) * -hp.min_level_db / (2 * hp.max_abs_value))
|
||||
+ hp.min_level_db)
|
||||
else:
|
||||
return ((np.clip(D, 0, hp.max_abs_value) * -hp.min_level_db / hp.max_abs_value) + hp.min_level_db)
|
||||
|
||||
if hp.symmetric_mels:
|
||||
return (((D + hp.max_abs_value) * -hp.min_level_db / (2 * hp.max_abs_value)) + hp.min_level_db)
|
||||
else:
|
||||
return ((D * -hp.min_level_db / hp.max_abs_value) + hp.min_level_db)
|
||||
@@ -0,0 +1,239 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn import TransformerEncoder, TransformerEncoderLayer
|
||||
import math
|
||||
|
||||
class PositionalEmbedding(nn.Module):
|
||||
def __init__(self, d_model=512, max_len=512):
|
||||
super().__init__()
|
||||
|
||||
# Compute the positional encodings once in log space.
|
||||
pe = torch.zeros(max_len, d_model).float()
|
||||
pe.require_grad = False
|
||||
|
||||
position = torch.arange(0, max_len).float().unsqueeze(1)
|
||||
div_term = (torch.arange(0, d_model, 2).float() * -(math.log(10000.0) / d_model)).exp()
|
||||
|
||||
pe[:, 0::2] = torch.sin(position * div_term)
|
||||
pe[:, 1::2] = torch.cos(position * div_term)
|
||||
|
||||
pe = pe.unsqueeze(0)
|
||||
self.register_buffer('pe', pe)
|
||||
|
||||
def forward(self, x):
|
||||
return self.pe[:, :x.size(1)]
|
||||
|
||||
class Conv1d(nn.Module):
|
||||
def __init__(self, cin, cout, kernel_size, stride, padding, residual=False,act='ReLU', *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.conv_block = nn.Sequential(
|
||||
nn.Conv1d(cin, cout, kernel_size, stride, padding),
|
||||
nn.BatchNorm1d(cout)
|
||||
)
|
||||
if act=='ReLU':
|
||||
self.act = nn.ReLU()
|
||||
elif act=='Tanh':
|
||||
self.act =nn.Tanh()
|
||||
self.residual = residual
|
||||
|
||||
def forward(self, x):
|
||||
out = self.conv_block(x)
|
||||
if self.residual:
|
||||
out += x
|
||||
return self.act(out)
|
||||
|
||||
class Conv2d(nn.Module):
|
||||
def __init__(self, cin, cout, kernel_size, stride, padding, residual=False, act='ReLU',*args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.conv_block = nn.Sequential(
|
||||
nn.Conv2d(cin, cout, kernel_size, stride, padding),
|
||||
nn.BatchNorm2d(cout)
|
||||
)
|
||||
if act == 'ReLU':
|
||||
self.act = nn.ReLU()
|
||||
elif act == 'Tanh':
|
||||
self.act = nn.Tanh()
|
||||
|
||||
self.residual = residual
|
||||
|
||||
def forward(self, x):
|
||||
out = self.conv_block(x)
|
||||
if self.residual:
|
||||
out += x
|
||||
return self.act(out)
|
||||
|
||||
|
||||
|
||||
def weight_init(m):
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.xavier_normal_(m.weight)
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.BatchNorm1d):
|
||||
nn.init.constant_(m.weight, 1)
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
|
||||
class Fusion_transformer_encoder(nn.Module):
|
||||
def __init__(self,T, d_model, nlayers, nhead, dim_feedforward, # 1024 128
|
||||
dropout=0.1):
|
||||
super().__init__()
|
||||
self.T=T
|
||||
self.position_v = PositionalEmbedding(d_model=512) #for visual landmarks
|
||||
self.position_a = PositionalEmbedding(d_model=512) #for audio embedding
|
||||
self.modality = nn.Embedding(4, 512, padding_idx=0) # 1 for pose, 2 for audio, 3 for reference landmarks
|
||||
self.dropout = nn.Dropout(p=dropout)
|
||||
encoder_layers = TransformerEncoderLayer(d_model, nhead, dim_feedforward, dropout, batch_first=True)
|
||||
self.transformer_encoder = TransformerEncoder(encoder_layers, nlayers)
|
||||
|
||||
def forward(self,ref_embedding,mel_embedding,pose_embedding):#(B,Nl,512) (B,T,512) (B,T,512)
|
||||
|
||||
# (1). positional(temporal) encoding
|
||||
position_v_encoding = self.position_v(pose_embedding) # (1, T, 512)
|
||||
position_a_encoding = self.position_a(mel_embedding)
|
||||
|
||||
#(2) modality encoding
|
||||
modality_v = self.modality(1 * torch.ones((pose_embedding.size(0), self.T), dtype=torch.int).cuda())
|
||||
modality_a = self.modality(2 * torch.ones((mel_embedding.size(0), self.T), dtype=torch.int).cuda())
|
||||
|
||||
pose_tokens = pose_embedding + position_v_encoding + modality_v #(B , T, 512 )
|
||||
audio_tokens = mel_embedding + position_a_encoding + modality_a #(B , T, 512 )
|
||||
ref_tokens = ref_embedding + self.modality(
|
||||
3 * torch.ones((ref_embedding.size(0), ref_embedding.size(1)), dtype=torch.int).cuda())
|
||||
|
||||
#(3) concat tokens
|
||||
input_tokens = torch.cat((ref_tokens, audio_tokens, pose_tokens), dim=1) # (B, 1+T+T, 512 )
|
||||
input_tokens = self.dropout(input_tokens)
|
||||
|
||||
#(4) input to transformer
|
||||
output = self.transformer_encoder(input_tokens)
|
||||
return output
|
||||
|
||||
|
||||
class Landmark_generator(nn.Module):
|
||||
def __init__(self,T,d_model,nlayers,nhead,dim_feedforward,dropout=0.1):
|
||||
super(Landmark_generator, self).__init__()
|
||||
self.mel_encoder=nn.Sequential( # (B*T,1,hv,wv)
|
||||
Conv2d(1, 32, kernel_size=3, stride=1, padding=1),
|
||||
Conv2d(32, 32, kernel_size=3, stride=1, padding=1, residual=True),
|
||||
Conv2d(32, 32, kernel_size=3, stride=1, padding=1, residual=True),
|
||||
|
||||
Conv2d(32, 64, kernel_size=3, stride=(3, 1), padding=1),
|
||||
Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True),
|
||||
Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True),
|
||||
|
||||
Conv2d(64, 128, kernel_size=3, stride=3, padding=1),
|
||||
Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True),
|
||||
Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True),
|
||||
|
||||
Conv2d(128, 256, kernel_size=3, stride=(3, 2), padding=1),
|
||||
Conv2d(256, 256, kernel_size=3, stride=1, padding=1, residual=True),
|
||||
|
||||
Conv2d(256, 512, kernel_size=3, stride=1, padding=0),
|
||||
Conv2d(512, 512, kernel_size=1, stride=1, padding=0,act='Tanh'),
|
||||
)
|
||||
|
||||
self.ref_encoder=nn.Sequential( # (B*Nl,2,131)
|
||||
Conv1d(2, 4, 3, 1, 1), #131
|
||||
|
||||
Conv1d(4, 8, 3, 2,1), #66
|
||||
Conv1d(8, 8, 3, 1, 1,residual=True),
|
||||
Conv1d(8, 8, 3, 1, 1,residual=True),
|
||||
|
||||
Conv1d(8, 16, 3, 2, 1), # 33
|
||||
Conv1d(16, 16, 3, 1, 1, residual=True),
|
||||
Conv1d(16, 16, 3, 1, 1, residual=True),
|
||||
|
||||
Conv1d(16, 32, 3, 2,1),# 17
|
||||
Conv1d(32, 32, 3, 1, 1,residual=True),
|
||||
Conv1d(32, 32, 3, 1, 1,residual=True),
|
||||
|
||||
Conv1d(32, 64, 3, 2,1), # 9
|
||||
Conv1d(64, 64, 3, 1, 1,residual=True),
|
||||
Conv1d(64, 64, 3, 1, 1,residual=True),
|
||||
|
||||
Conv1d(64, 128, 3, 2,1), # 5
|
||||
Conv1d(128, 128, 3, 1, 1,residual=True),
|
||||
Conv1d(128, 128, 3, 1, 1,residual=True),
|
||||
|
||||
Conv1d(128, 256, 3, 2,1), #3
|
||||
Conv1d(256, 256, 3, 1, 1,residual=True),
|
||||
|
||||
Conv1d(256, 512, 3, 1,0), #1
|
||||
Conv1d(512, 512, 1, 1,0,act='Tanh'), #1
|
||||
)
|
||||
self.pose_encoder=nn.Sequential( # (B*T,2,74)
|
||||
Conv1d(2, 4, 3, 1, 1),
|
||||
|
||||
Conv1d(4, 8, 3, 1, 1), #74
|
||||
Conv1d(8, 8, 3, 1, 1,residual=True),
|
||||
Conv1d(8, 8, 3, 1, 1, residual=True),
|
||||
|
||||
Conv1d(8, 16, 3, 2, 1), # 37
|
||||
Conv1d(16, 16, 3, 1, 1, residual=True),
|
||||
Conv1d(16, 16, 3, 1, 1, residual=True),
|
||||
|
||||
Conv1d(16, 32, 3, 2, 1), # 19
|
||||
Conv1d(32, 32, 3, 1, 1, residual=True),
|
||||
Conv1d(32, 32, 3, 1, 1, residual=True),
|
||||
|
||||
Conv1d(32, 64, 3, 2, 1), #10
|
||||
Conv1d(64, 64, 3, 1, 1, residual=True),
|
||||
Conv1d(64, 64, 3, 1, 1, residual=True),
|
||||
|
||||
Conv1d(64, 128, 3, 2, 1), # 5
|
||||
Conv1d(128, 128, 3, 1, 1, residual=True),
|
||||
Conv1d(128, 128, 3, 1, 1, residual=True),
|
||||
|
||||
Conv1d(128, 256, 3, 2, 1), # 3
|
||||
Conv1d(256, 256, 3, 1, 1, residual=True),
|
||||
Conv1d(256, 256, 3, 1, 1, residual=True),
|
||||
|
||||
Conv1d(256, 512, 3, 1, 0), # 1
|
||||
Conv1d(512, 512, 1, 1, 0, residual=True,act='Tanh'),
|
||||
)
|
||||
|
||||
self.fusion_transformer = Fusion_transformer_encoder(T,d_model,nlayers,nhead,dim_feedforward,dropout)
|
||||
|
||||
self.mouse_keypoint_map = nn.Linear(d_model, 40 * 2)
|
||||
self.jaw_keypoint_map = nn.Linear(d_model, 17 * 2)
|
||||
|
||||
self.apply(weight_init)
|
||||
self.Norm=nn.LayerNorm(512)
|
||||
|
||||
def forward(self, T_mels, T_pose, Nl_pose, Nl_content):
|
||||
# (B,T,1,hv,wv) (B,T,2,74) (B,N_l,2,74) (B,N_l,2,57)
|
||||
B,T,N_l= T_mels.size(0),T_mels.size(1),Nl_content.size(1)
|
||||
|
||||
#1. obtain full reference landmarks
|
||||
Nl_ref = torch.cat([Nl_pose, Nl_content], dim=3) #(B,Nl,2,131=74+57)
|
||||
Nl_ref = torch.cat([Nl_ref[i] for i in range(Nl_ref.size(0))], dim=0) # (B*Nl,2,131)
|
||||
|
||||
T_mels=torch.cat([T_mels[i] for i in range(T_mels.size(0))],dim=0) #(B*T,1,hv,wv)
|
||||
T_pose = torch.cat([T_pose[i] for i in range(T_pose.size(0))],dim=0) # (B*T,2,74)
|
||||
|
||||
# 2. get embedding
|
||||
mel_embedding=self.mel_encoder(T_mels).squeeze(-1).squeeze(-1)#(B*T,512)
|
||||
pose_embedding=self.pose_encoder(T_pose).squeeze(-1) # (B*T,512)
|
||||
ref_embedding = self.ref_encoder(Nl_ref).squeeze(-1) # (B*Nl,512)
|
||||
#normalization
|
||||
mel_embedding = self.Norm(mel_embedding) # (B*T,512)
|
||||
pose_embedding =self.Norm(pose_embedding) # (B*T,512)
|
||||
ref_embedding = self.Norm(ref_embedding) # (B*Nl,512)
|
||||
|
||||
mel_embedding = torch.stack(torch.split(mel_embedding,T),dim=0) #(B,T,512)
|
||||
pose_embedding = torch.stack(torch.split(pose_embedding, T), dim=0) # (B,T,512)
|
||||
ref_embedding=torch.stack(torch.split(ref_embedding,N_l,dim=0),dim=0) #(B,N_l,512)
|
||||
|
||||
#3. fuse embedding
|
||||
output_tokens=self.fusion_transformer(ref_embedding,mel_embedding,pose_embedding)
|
||||
|
||||
#4.output landmark
|
||||
lip_embedding=output_tokens[:,N_l:N_l+T,:] #(B,T,dim)
|
||||
jaw_embedding=output_tokens[:,N_l+T:,:] #(B,T,dim)
|
||||
output_mouse_landmark=self.mouse_keypoint_map(lip_embedding) ##(B,T,40*2)
|
||||
output_jaw_landmark=self.jaw_keypoint_map(jaw_embedding) ##(B,T,17*2)
|
||||
|
||||
predict_content=torch.reshape(torch.cat([output_jaw_landmark,output_mouse_landmark],dim=2),(B,T,-1,2)) #(B,T,57,2)
|
||||
predict_content=torch.cat([predict_content[i] for i in range(predict_content.size(0))],dim=0).permute(0,2,1)#(B*T,2,57)
|
||||
return predict_content #(B*T,2,57)
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import functools
|
||||
from torch.autograd import Variable
|
||||
import numpy as np
|
||||
|
||||
|
||||
def weights_init(m):
|
||||
classname = m.__class__.__name__
|
||||
if classname.find('Conv') != -1:
|
||||
m.weight.data.normal_(0.0, 0.02)
|
||||
elif classname.find('BatchNorm2d') != -1:
|
||||
m.weight.data.normal_(1.0, 0.02)
|
||||
m.bias.data.fill_(0)
|
||||
|
||||
|
||||
def define_D(input_nc=3, ndf=64, n_layers_D=3, norm='instance', use_sigmoid=False, num_D=2, getIntermFeat=True):
|
||||
#('--ndf', type=int, default=64, help='# of discrim filters in first conv layer')
|
||||
#('--input_nc', type=int, default=3, help='# of input image channels')
|
||||
#('--n_layers_D', type=int, default=3, help='only used if which_model_netD==n_layers')
|
||||
# ('--num_D', type=int, default=2, help='number of discriminators to use')
|
||||
|
||||
norm_layer = get_norm_layer(norm_type=norm)
|
||||
netD = MultiscaleDiscriminator(input_nc, ndf, n_layers_D, norm_layer, use_sigmoid, num_D, getIntermFeat)
|
||||
#print(netD)
|
||||
netD.apply(weights_init)
|
||||
return netD
|
||||
|
||||
|
||||
class NLayerDiscriminator(nn.Module):
|
||||
def __init__(self, input_nc, ndf=64, n_layers=3, norm_layer=nn.BatchNorm2d, use_sigmoid=False, getIntermFeat=False):
|
||||
super(NLayerDiscriminator, self).__init__()
|
||||
self.getIntermFeat = getIntermFeat
|
||||
self.n_layers = n_layers
|
||||
|
||||
kw = 4
|
||||
padw = int(np.ceil((kw-1.0)/2))
|
||||
sequence = [[nn.Conv2d(input_nc, ndf, kernel_size=kw, stride=2, padding=padw), nn.LeakyReLU(0.2, True)]]
|
||||
|
||||
nf = ndf
|
||||
for n in range(1, n_layers):
|
||||
nf_prev = nf
|
||||
nf = min(nf * 2, 512)
|
||||
sequence += [[
|
||||
nn.Conv2d(nf_prev, nf, kernel_size=kw, stride=2, padding=padw),
|
||||
norm_layer(nf), nn.LeakyReLU(0.2, True)
|
||||
]]
|
||||
|
||||
nf_prev = nf
|
||||
nf = min(nf * 2, 512)
|
||||
sequence += [[
|
||||
nn.Conv2d(nf_prev, nf, kernel_size=kw, stride=1, padding=padw),
|
||||
norm_layer(nf),
|
||||
nn.LeakyReLU(0.2, True)
|
||||
]]
|
||||
|
||||
sequence += [[nn.Conv2d(nf, 1, kernel_size=kw, stride=1, padding=padw)]]
|
||||
|
||||
if use_sigmoid:
|
||||
sequence += [[nn.Sigmoid()]]
|
||||
|
||||
if getIntermFeat:
|
||||
for n in range(len(sequence)):
|
||||
setattr(self, 'model'+str(n), nn.Sequential(*sequence[n]))
|
||||
else:
|
||||
sequence_stream = []
|
||||
for n in range(len(sequence)):
|
||||
sequence_stream += sequence[n]
|
||||
self.model = nn.Sequential(*sequence_stream)
|
||||
|
||||
def forward(self, input):
|
||||
|
||||
if self.getIntermFeat:
|
||||
res = [input]
|
||||
for n in range(self.n_layers+2):
|
||||
model = getattr(self, 'model'+str(n))
|
||||
res.append(model(res[-1]))
|
||||
return res[1:]
|
||||
else:
|
||||
return self.model(input)
|
||||
|
||||
|
||||
def get_norm_layer(norm_type='instance'):
|
||||
if norm_type == 'batch':
|
||||
norm_layer = functools.partial(nn.BatchNorm2d, affine=True)
|
||||
elif norm_type == 'instance':
|
||||
norm_layer = functools.partial(nn.InstanceNorm2d, affine=False)
|
||||
else:
|
||||
raise NotImplementedError('normalization layer [%s] is not found' % norm_type)
|
||||
return norm_layer
|
||||
|
||||
|
||||
|
||||
|
||||
class MultiscaleDiscriminator(nn.Module):
|
||||
def __init__(self, input_nc, ndf=64, n_layers=3, norm_layer=nn.BatchNorm2d,
|
||||
use_sigmoid=False, num_D=3, getIntermFeat=False):
|
||||
super(MultiscaleDiscriminator, self).__init__()
|
||||
self.num_D = num_D
|
||||
self.n_layers = n_layers
|
||||
self.getIntermFeat = getIntermFeat
|
||||
|
||||
for i in range(num_D):
|
||||
netD = NLayerDiscriminator(input_nc, ndf, n_layers, norm_layer, use_sigmoid, getIntermFeat)
|
||||
if getIntermFeat:
|
||||
for j in range(n_layers + 2):
|
||||
setattr(self, 'scale' + str(i) + '_layer' + str(j), getattr(netD, 'model' + str(j)))
|
||||
else:
|
||||
setattr(self, 'layer' + str(i), netD.model)
|
||||
|
||||
self.downsample = nn.AvgPool2d(3, stride=2, padding=[1, 1], count_include_pad=False)
|
||||
|
||||
def singleD_forward(self, model, input):
|
||||
if self.getIntermFeat:
|
||||
result = [input]
|
||||
for i in range(len(model)):
|
||||
result.append(model[i](result[-1]))
|
||||
return result[1:]
|
||||
else:
|
||||
return [model(input)]
|
||||
|
||||
def forward(self, input): #: (B,T,C,H,W)
|
||||
# input = torch.cat([input[i,:] for i in range(input.size(0))], dim=0)# : (B*T,C,H,W)
|
||||
num_D = self.num_D
|
||||
result = []
|
||||
input_downsampled = input
|
||||
for i in range(num_D):
|
||||
if self.getIntermFeat:
|
||||
model = [getattr(self, 'scale' + str(num_D - 1 - i) + '_layer' + str(j)) for j in
|
||||
range(self.n_layers + 2)]
|
||||
else:
|
||||
model = getattr(self, 'layer' + str(num_D - 1 - i))
|
||||
result.append(self.singleD_forward(model, input_downsampled))
|
||||
if i != (num_D - 1):
|
||||
input_downsampled = self.downsample(input_downsampled)
|
||||
return result
|
||||
|
||||
@@ -0,0 +1,571 @@
|
||||
from torch.nn import functional as F
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torchvision
|
||||
|
||||
|
||||
|
||||
class AdaINLayer(nn.Module):
|
||||
def __init__(self, input_nc, modulation_nc):
|
||||
super().__init__()
|
||||
|
||||
self.InstanceNorm2d = nn.InstanceNorm2d(input_nc, affine=False)
|
||||
|
||||
nhidden = 128
|
||||
use_bias=True
|
||||
|
||||
self.mlp_shared = nn.Sequential(
|
||||
nn.Linear(modulation_nc, nhidden, bias=use_bias),
|
||||
nn.ReLU()
|
||||
)
|
||||
self.mlp_gamma = nn.Linear(nhidden, input_nc, bias=use_bias)
|
||||
self.mlp_beta = nn.Linear(nhidden, input_nc, bias=use_bias)
|
||||
|
||||
def forward(self, input, modulation_input):
|
||||
|
||||
# Part 1. generate parameter-free normalized activations
|
||||
normalized = self.InstanceNorm2d(input)
|
||||
|
||||
# Part 2. produce scaling and bias conditioned on feature
|
||||
modulation_input = modulation_input.view(modulation_input.size(0), -1)
|
||||
actv = self.mlp_shared(modulation_input)
|
||||
gamma = self.mlp_gamma(actv)
|
||||
beta = self.mlp_beta(actv)
|
||||
|
||||
# apply scale and bias
|
||||
gamma = gamma.view(*gamma.size()[:2], 1,1)
|
||||
beta = beta.view(*beta.size()[:2], 1,1)
|
||||
out = normalized * (1 + gamma) + beta
|
||||
return out
|
||||
|
||||
class AdaIN(torch.nn.Module):
|
||||
|
||||
def __init__(self, input_channel, modulation_channel,kernel_size=3, stride=1, padding=1):
|
||||
super(AdaIN, self).__init__()
|
||||
self.conv_1 = torch.nn.Conv2d(input_channel, input_channel, kernel_size=kernel_size, stride=stride, padding=padding)
|
||||
self.conv_2 = torch.nn.Conv2d(input_channel, input_channel, kernel_size=kernel_size, stride=stride, padding=padding)
|
||||
self.leaky_relu = torch.nn.LeakyReLU(0.2)
|
||||
self.adain_layer_1 = AdaINLayer(input_channel, modulation_channel)
|
||||
self.adain_layer_2 = AdaINLayer(input_channel, modulation_channel)
|
||||
|
||||
def forward(self, x, modulation):
|
||||
|
||||
x = self.adain_layer_1(x, modulation)
|
||||
x = self.leaky_relu(x)
|
||||
x = self.conv_1(x)
|
||||
x = self.adain_layer_2(x, modulation)
|
||||
x = self.leaky_relu(x)
|
||||
x = self.conv_2(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
|
||||
|
||||
class SPADELayer(torch.nn.Module):
|
||||
def __init__(self, input_channel, modulation_channel, hidden_size=256, kernel_size=3, stride=1, padding=1):
|
||||
super(SPADELayer, self).__init__()
|
||||
self.instance_norm = torch.nn.InstanceNorm2d(input_channel)
|
||||
|
||||
self.conv1 = torch.nn.Conv2d(modulation_channel, hidden_size, kernel_size=kernel_size, stride=stride,
|
||||
padding=padding)
|
||||
self.gamma = torch.nn.Conv2d(hidden_size, input_channel, kernel_size=kernel_size, stride=stride, padding=padding)
|
||||
self.beta = torch.nn.Conv2d(hidden_size, input_channel, kernel_size=kernel_size, stride=stride, padding=padding)
|
||||
|
||||
def forward(self, input, modulation):
|
||||
norm = self.instance_norm(input)
|
||||
|
||||
conv_out = self.conv1(modulation)
|
||||
|
||||
gamma = self.gamma(conv_out)
|
||||
beta = self.beta(conv_out)
|
||||
|
||||
return norm + norm * gamma + beta
|
||||
|
||||
|
||||
class SPADE(torch.nn.Module):
|
||||
def __init__(self, num_channel, num_channel_modulation, hidden_size=256, kernel_size=3, stride=1, padding=1):
|
||||
super(SPADE, self).__init__()
|
||||
self.conv_1 = torch.nn.Conv2d(num_channel, num_channel, kernel_size=kernel_size, stride=stride, padding=padding)
|
||||
self.conv_2 = torch.nn.Conv2d(num_channel, num_channel, kernel_size=kernel_size, stride=stride, padding=padding)
|
||||
self.leaky_relu = torch.nn.LeakyReLU(0.2)
|
||||
self.spade_layer_1 = SPADELayer(num_channel, num_channel_modulation, hidden_size, kernel_size=kernel_size,
|
||||
stride=stride, padding=padding)
|
||||
self.spade_layer_2 = SPADELayer(num_channel, num_channel_modulation, hidden_size, kernel_size=kernel_size,
|
||||
stride=stride, padding=padding)
|
||||
|
||||
def forward(self, input, modulations):
|
||||
input = self.spade_layer_1(input, modulations)
|
||||
input = self.leaky_relu(input)
|
||||
input = self.conv_1(input)
|
||||
input = self.spade_layer_2(input, modulations)
|
||||
input = self.leaky_relu(input)
|
||||
input = self.conv_2(input)
|
||||
return input
|
||||
|
||||
class Conv2d(nn.Module):
|
||||
def __init__(self, cin, cout, kernel_size, stride, padding, residual=False, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.conv_block = nn.Sequential(
|
||||
nn.Conv2d(cin, cout, kernel_size, stride, padding),
|
||||
nn.BatchNorm2d(cout)
|
||||
)
|
||||
self.act = nn.ReLU()
|
||||
self.residual = residual
|
||||
|
||||
def forward(self, x):
|
||||
out = self.conv_block(x)
|
||||
if self.residual:
|
||||
out += x
|
||||
return self.act(out)
|
||||
|
||||
def downsample(x, size):
|
||||
if len(x.size()) == 5:
|
||||
size = (x.size(2), size[0], size[1])
|
||||
return torch.nn.functional.interpolate(x, size=size, mode='nearest')
|
||||
return torch.nn.functional.interpolate(x, size=size, mode='nearest')
|
||||
|
||||
|
||||
def convert_flow_to_deformation(flow):
|
||||
r"""convert flow fields to deformations.
|
||||
Args:
|
||||
flow (tensor): Flow field obtained by the model
|
||||
Returns:
|
||||
deformation (tensor): The deformation used for warpping
|
||||
"""
|
||||
b, c, h, w = flow.shape
|
||||
flow_norm = 2 * torch.cat([flow[:, :1, ...] / (w - 1), flow[:, 1:, ...] / (h - 1)], 1)
|
||||
grid = make_coordinate_grid(flow)
|
||||
deformation = grid + flow_norm.permute(0, 2, 3, 1)
|
||||
return deformation
|
||||
|
||||
|
||||
def make_coordinate_grid(flow):
|
||||
r"""obtain coordinate grid with the same size as the flow filed.
|
||||
Args:
|
||||
flow (tensor): Flow field obtained by the model
|
||||
Returns:
|
||||
grid (tensor): The grid with the same size as the input flow
|
||||
"""
|
||||
b, c, h, w = flow.shape
|
||||
|
||||
x = torch.arange(w).to(flow)
|
||||
y = torch.arange(h).to(flow)
|
||||
|
||||
x = (2 * (x / (w - 1)) - 1)
|
||||
y = (2 * (y / (h - 1)) - 1)
|
||||
|
||||
yy = y.view(-1, 1).repeat(1, w)
|
||||
xx = x.view(1, -1).repeat(h, 1)
|
||||
|
||||
meshed = torch.cat([xx.unsqueeze_(2), yy.unsqueeze_(2)], 2)
|
||||
meshed = meshed.expand(b, -1, -1, -1)
|
||||
return meshed
|
||||
|
||||
|
||||
def warping(source_image, deformation):
|
||||
r"""warp the input image according to the deformation
|
||||
Args:
|
||||
source_image (tensor): source images to be warpped
|
||||
deformation (tensor): deformations used to warp the images; value in range (-1, 1)
|
||||
Returns:
|
||||
output (tensor): the warpped images
|
||||
"""
|
||||
_, h_old, w_old, _ = deformation.shape
|
||||
_, _, h, w = source_image.shape
|
||||
if h_old != h or w_old != w:
|
||||
deformation = deformation.permute(0, 3, 1, 2)
|
||||
deformation = torch.nn.functional.interpolate(deformation, size=(h, w), mode='bilinear')
|
||||
deformation = deformation.permute(0, 2, 3, 1)
|
||||
return torch.nn.functional.grid_sample(source_image, deformation)
|
||||
|
||||
|
||||
class DenseFlowNetwork(torch.nn.Module):
|
||||
def __init__(self, num_channel=6, num_channel_modulation=3*5, hidden_size=256):
|
||||
super(DenseFlowNetwork, self).__init__()
|
||||
|
||||
# Convolutional Layers
|
||||
self.conv1 = torch.nn.Conv2d(num_channel, 32, kernel_size=7, stride=1, padding=3)
|
||||
self.conv1_bn = torch.nn.BatchNorm2d(num_features=32, affine=True)
|
||||
self.conv1_relu = torch.nn.ReLU()
|
||||
|
||||
self.conv2 = torch.nn.Conv2d(32, 256, kernel_size=3, stride=2, padding=1)
|
||||
self.conv2_bn = torch.nn.BatchNorm2d(num_features=256, affine=True)
|
||||
self.conv2_relu = torch.nn.ReLU()
|
||||
|
||||
|
||||
# SPADE Blocks
|
||||
self.spade_layer_1 = SPADE(256, num_channel_modulation, hidden_size)
|
||||
self.spade_layer_2 = SPADE(256, num_channel_modulation, hidden_size)
|
||||
self.pixel_shuffle_1 = torch.nn.PixelShuffle(2)
|
||||
self.spade_layer_4 = SPADE(64, num_channel_modulation, hidden_size)
|
||||
|
||||
# Final Convolutional Layer
|
||||
self.conv_4 = torch.nn.Conv2d(64, 2, kernel_size=7, stride=1, padding=3)
|
||||
self.conv_5= nn.Sequential(torch.nn.Conv2d(64, 32, kernel_size=7, stride=1, padding=3),
|
||||
torch.nn.ReLU(),
|
||||
torch.nn.Conv2d(32, 1, kernel_size=7, stride=1, padding=3),
|
||||
torch.nn.Sigmoid(),
|
||||
)#predict weight
|
||||
|
||||
def forward(self, ref_N_frame_img, ref_N_frame_sketch, T_driving_sketch): #to output: (B*T,3,H,W)
|
||||
# (B, N, 3, H, W)(B, N, 3, H, W) (B, 5, 3, H, W) #
|
||||
ref_N = ref_N_frame_img.size(1)
|
||||
|
||||
driving_sketch=torch.cat([T_driving_sketch[:,i] for i in range(T_driving_sketch.size(1))], dim=1) #(B, 3*5, H, W)
|
||||
|
||||
wrapped_h1_sum, wrapped_h2_sum, wrapped_ref_sum=0.,0.,0.
|
||||
softmax_denominator=0.
|
||||
T = 1 # during rendering, generate T=1 image at a time
|
||||
for ref_idx in range(ref_N): # each ref img provide information for each B*T frame
|
||||
ref_img= ref_N_frame_img[:, ref_idx] #(B, 3, H, W)
|
||||
ref_img = ref_img.unsqueeze(1).expand(-1, T, -1, -1, -1) # (B,T, 3, H, W)
|
||||
ref_img = torch.cat([ref_img[i] for i in range(ref_img.size(0))], dim=0) # (B*T, 3, H, W)
|
||||
|
||||
ref_sketch = ref_N_frame_sketch[:, ref_idx] #(B, 3, H, W)
|
||||
ref_sketch = ref_sketch.unsqueeze(1).expand(-1, T, -1, -1, -1) # (B,T, 3, H, W)
|
||||
ref_sketch = torch.cat([ref_sketch[i] for i in range(ref_sketch.size(0))], dim=0) # (B*T, 3, H, W)
|
||||
|
||||
#predict flow and weight
|
||||
flow_module_input = torch.cat((ref_img, ref_sketch), dim=1) #(B*T, 3+3, H, W)
|
||||
# Convolutional Layers
|
||||
h1 = self.conv1_relu(self.conv1_bn(self.conv1(flow_module_input))) #(32,128,128)
|
||||
h2 = self.conv2_relu(self.conv2_bn(self.conv2(h1))) #(256,64,64)
|
||||
# SPADE Blocks
|
||||
downsample_64 = downsample(driving_sketch, (64, 64)) # driving_sketch:(B*T, 3, H, W)
|
||||
|
||||
spade_layer = self.spade_layer_1(h2, downsample_64) #(256,64,64)
|
||||
spade_layer = self.spade_layer_2(spade_layer, downsample_64) #(256,64,64)
|
||||
|
||||
spade_layer = self.pixel_shuffle_1(spade_layer) #(64,128,128)
|
||||
|
||||
spade_layer = self.spade_layer_4(spade_layer, driving_sketch) #(64,128,128)
|
||||
|
||||
# Final Convolutional Layer
|
||||
output_flow = self.conv_4(spade_layer) # (B*T,2,128,128)
|
||||
output_weight=self.conv_5(spade_layer) # (B*T,1,128,128)
|
||||
|
||||
deformation=convert_flow_to_deformation(output_flow)
|
||||
wrapped_h1 = warping(h1, deformation) #(32,128,128)
|
||||
wrapped_h2 = warping(h2, deformation) #(256,64,64)
|
||||
wrapped_ref = warping(ref_img, deformation) #(3,128,128)
|
||||
|
||||
softmax_denominator+=output_weight
|
||||
wrapped_h1_sum+=wrapped_h1*output_weight
|
||||
wrapped_h2_sum+=wrapped_h2*downsample(output_weight, (64,64))
|
||||
wrapped_ref_sum+=wrapped_ref*output_weight
|
||||
#return weighted warped feataure and images
|
||||
softmax_denominator+=0.00001
|
||||
wrapped_h1_sum=wrapped_h1_sum/softmax_denominator
|
||||
wrapped_h2_sum = wrapped_h2_sum / downsample(softmax_denominator, (64,64))
|
||||
wrapped_ref_sum = wrapped_ref_sum / softmax_denominator
|
||||
return wrapped_h1_sum, wrapped_h2_sum, wrapped_ref_sum
|
||||
|
||||
|
||||
class TranslationNetwork(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super(TranslationNetwork, self).__init__()
|
||||
self.audio_encoder = nn.Sequential(
|
||||
Conv2d(1, 32, kernel_size=3, stride=1, padding=1),
|
||||
Conv2d(32, 32, kernel_size=3, stride=1, padding=1, residual=True),
|
||||
Conv2d(32, 32, kernel_size=3, stride=1, padding=1, residual=True),
|
||||
|
||||
Conv2d(32, 64, kernel_size=3, stride=(3, 1), padding=1),
|
||||
Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True),
|
||||
Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True),
|
||||
|
||||
Conv2d(64, 128, kernel_size=3, stride=3, padding=1),
|
||||
Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True),
|
||||
Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True),
|
||||
|
||||
Conv2d(128, 256, kernel_size=3, stride=(3, 2), padding=1),
|
||||
Conv2d(256, 256, kernel_size=3, stride=1, padding=1, residual=True),
|
||||
|
||||
Conv2d(256, 512, kernel_size=3, stride=1, padding=0),
|
||||
Conv2d(512, 512, kernel_size=1, stride=1, padding=0), )
|
||||
|
||||
# Encoder
|
||||
self.conv1 = torch.nn.Conv2d(in_channels=3+3*5, out_channels=32, kernel_size=7, stride=1, padding=3, bias=False)
|
||||
self.conv1_bn = torch.nn.BatchNorm2d(num_features=32, affine=True)
|
||||
self.conv1_relu = torch.nn.ReLU()
|
||||
|
||||
self.conv2 = torch.nn.Conv2d(in_channels=32, out_channels=256, kernel_size=3, stride=2, padding=1, bias=False)
|
||||
self.conv2_bn = torch.nn.BatchNorm2d(num_features=256, affine=True)
|
||||
self.conv2_relu = torch.nn.ReLU()
|
||||
|
||||
# Decoder
|
||||
self.spade_1 = SPADE(num_channel=256, num_channel_modulation=256)
|
||||
self.adain_1 = AdaIN(256,512)
|
||||
self.pixel_suffle_1 = nn.PixelShuffle(upscale_factor=2)
|
||||
|
||||
self.spade_2 = SPADE(num_channel=64, num_channel_modulation=32)
|
||||
self.adain_2 = AdaIN(input_channel=64,modulation_channel=512)
|
||||
|
||||
self.spade_4 = SPADE(num_channel=64, num_channel_modulation=3)
|
||||
|
||||
# Final layer
|
||||
self.leaky_relu = torch.nn.LeakyReLU()
|
||||
self.conv_last = torch.nn.Conv2d(in_channels=64, out_channels=3, kernel_size=7, stride=1, padding=3, bias=False)
|
||||
self.Sigmoid=torch.nn.Sigmoid()
|
||||
def forward(self, translation_input, wrapped_ref, wrapped_h1, wrapped_h2, T_mels):
|
||||
# (B,3+3,H,W) (B,3,128,128) (B,32,128,128) (B,256,64,64) (B,T,1,h,w) #T=1
|
||||
# Encoder
|
||||
T_mels=torch.cat([T_mels[i] for i in range(T_mels.size(0))],dim=0)# B*T,1,h,w
|
||||
x = self.conv1_relu(self.conv1_bn(self.conv1(translation_input))) #32,128,128
|
||||
x = self.conv2_relu(self.conv2_bn(self.conv2(x))) #256,64,64
|
||||
|
||||
audio_feature = self.audio_encoder(T_mels).squeeze(-1).permute(0,2,1) #(B*T,1,512)
|
||||
|
||||
# Decoder
|
||||
x = self.spade_1(x, wrapped_h2) # (C=256,64,64)
|
||||
x = self.adain_1(x, audio_feature) # (C=256,64,64)
|
||||
x = self.pixel_suffle_1(x) # (C=64,128,128)
|
||||
|
||||
x = self.spade_2(x, wrapped_h1) # (64,128,128)
|
||||
x = self.adain_2(x, audio_feature) # (64,128,128)
|
||||
x = self.spade_4(x, wrapped_ref) # (64,128,128)
|
||||
|
||||
# output layer
|
||||
x = self.leaky_relu(x)
|
||||
x = self.conv_last(x)
|
||||
x = self.Sigmoid(x)
|
||||
return x
|
||||
|
||||
class Renderer(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super(Renderer, self).__init__()
|
||||
|
||||
# 1.flow Network
|
||||
self.flow_module = DenseFlowNetwork()
|
||||
#2. translation Network
|
||||
self.translation = TranslationNetwork()
|
||||
#3.return loss
|
||||
self.perceptual = PerceptualLoss(network='vgg19',
|
||||
layers=['relu_1_1', 'relu_2_1', 'relu_3_1', 'relu_4_1', 'relu_5_1'],
|
||||
num_scales=2)
|
||||
|
||||
def forward(self, face_frame_img, target_sketches, ref_N_frame_img, ref_N_frame_sketch, audio_mels): #T=1
|
||||
# (B,1,3,H,W) (B,5,3,H,W) (B,N,3,H,W) (B,N,3,H,W) (B,T,1,hv,wv)T=1
|
||||
# (1)warping reference images and their feature
|
||||
wrapped_h1, wrapped_h2, wrapped_ref = self.flow_module(ref_N_frame_img, ref_N_frame_sketch, target_sketches)
|
||||
#(B,C,H,W)
|
||||
|
||||
# (2)translation module
|
||||
target_sketches = torch.cat([target_sketches[:, i] for i in range(target_sketches.size(1))], dim=1)
|
||||
# (B,3*T,H,W)
|
||||
gt_face = torch.cat([face_frame_img[i] for i in range(face_frame_img.size(0))], dim=0)
|
||||
# (B,3,H,W)
|
||||
gt_mask_face = gt_face.clone()
|
||||
gt_mask_face[:, :, gt_mask_face.size(2) // 2:, :] = 0 # (B,3,H,W)
|
||||
#
|
||||
translation_input=torch.cat([gt_mask_face, target_sketches], dim=1) # (B*T,3+3,H,W)
|
||||
generated_face = self.translation(translation_input, wrapped_ref, wrapped_h1, wrapped_h2, audio_mels) #translation_input
|
||||
|
||||
perceptual_gen_loss = self.perceptual(generated_face, gt_face, use_style_loss=True,
|
||||
weight_style_to_perceptual=250).mean()
|
||||
perceptual_warp_loss = self.perceptual(wrapped_ref, gt_face, use_style_loss=False,
|
||||
weight_style_to_perceptual=0.).mean()
|
||||
return generated_face, wrapped_ref, torch.unsqueeze(perceptual_warp_loss, 0), torch.unsqueeze(
|
||||
perceptual_gen_loss, 0)
|
||||
# (B,3,H,W) and losses
|
||||
|
||||
#the following is the code for Perceptual(VGG) loss
|
||||
|
||||
def apply_imagenet_normalization(input):
|
||||
r"""Normalize using ImageNet mean and std.
|
||||
|
||||
Args:
|
||||
input (4D tensor NxCxHxW): The input images, assuming to be [-1, 1].
|
||||
|
||||
Returns:
|
||||
Normalized inputs using the ImageNet normalization.
|
||||
"""
|
||||
# normalize the input back to [0, 1]
|
||||
normalized_input = (input + 1) / 2
|
||||
# normalize the input using the ImageNet mean and std
|
||||
mean = normalized_input.new_tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)
|
||||
std = normalized_input.new_tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)
|
||||
output = (normalized_input - mean) / std
|
||||
return output
|
||||
|
||||
class _PerceptualNetwork(nn.Module):
|
||||
r"""The network that extracts features to compute the perceptual loss.
|
||||
|
||||
Args:
|
||||
network (nn.Sequential) : The network that extracts features.
|
||||
layer_name_mapping (dict) : The dictionary that
|
||||
maps a layer's index to its name.
|
||||
layers (list of str): The list of layer names that we are using.
|
||||
"""
|
||||
|
||||
def __init__(self, network, layer_name_mapping, layers):
|
||||
super().__init__()
|
||||
assert isinstance(network, nn.Sequential), \
|
||||
'The network needs to be of type "nn.Sequential".'
|
||||
self.network = network
|
||||
self.layer_name_mapping = layer_name_mapping
|
||||
self.layers = layers
|
||||
for param in self.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
def forward(self, x):
|
||||
r"""Extract perceptual features."""
|
||||
output = {}
|
||||
for i, layer in enumerate(self.network):
|
||||
x = layer(x)
|
||||
layer_name = self.layer_name_mapping.get(i, None)
|
||||
if layer_name in self.layers:
|
||||
# If the current layer is used by the perceptual loss.
|
||||
output[layer_name] = x
|
||||
return output
|
||||
|
||||
def _vgg19(layers):
|
||||
r"""Get vgg19 layers"""
|
||||
network = torchvision.models.vgg19(pretrained=True).features
|
||||
layer_name_mapping = {1: 'relu_1_1',
|
||||
3: 'relu_1_2',
|
||||
6: 'relu_2_1',
|
||||
8: 'relu_2_2',
|
||||
11: 'relu_3_1',
|
||||
13: 'relu_3_2',
|
||||
15: 'relu_3_3',
|
||||
17: 'relu_3_4',
|
||||
20: 'relu_4_1',
|
||||
22: 'relu_4_2',
|
||||
24: 'relu_4_3',
|
||||
26: 'relu_4_4',
|
||||
29: 'relu_5_1'}
|
||||
return _PerceptualNetwork(network, layer_name_mapping, layers)
|
||||
|
||||
class PerceptualLoss(nn.Module):
|
||||
r"""Perceptual loss initialization.
|
||||
|
||||
Args:
|
||||
network (str) : The name of the loss network: 'vgg16' | 'vgg19'.
|
||||
layers (str or list of str) : The layers used to compute the loss.
|
||||
weights (float or list of float : The loss weights of each layer.
|
||||
criterion (str): The type of distance function: 'l1' | 'l2'.
|
||||
resize (bool) : If ``True``, resize the input images to 224x224.
|
||||
resize_mode (str): Algorithm used for resizing.
|
||||
instance_normalized (bool): If ``True``, applies instance normalization
|
||||
to the feature maps before computing the distance.
|
||||
num_scales (int): The loss will be evaluated at original size and
|
||||
this many times downsampled sizes.
|
||||
"""
|
||||
|
||||
def __init__(self, network='vgg19', layers='relu_4_1', weights=None,
|
||||
criterion='l1', resize=False, resize_mode='bilinear',
|
||||
instance_normalized=False, num_scales=1,):
|
||||
super().__init__()
|
||||
if isinstance(layers, str):
|
||||
layers = [layers]
|
||||
if weights is None:
|
||||
weights = [1.] * len(layers)
|
||||
elif isinstance(layers, float) or isinstance(layers, int):
|
||||
weights = [weights]
|
||||
|
||||
assert len(layers) == len(weights), \
|
||||
'The number of layers (%s) must be equal to ' \
|
||||
'the number of weights (%s).' % (len(layers), len(weights))
|
||||
if network == 'vgg19':
|
||||
self.model = _vgg19(layers)
|
||||
else:
|
||||
raise ValueError('Network %s is not recognized' % network)
|
||||
|
||||
self.num_scales = num_scales
|
||||
self.layers = layers
|
||||
self.weights = weights
|
||||
if criterion == 'l1':
|
||||
self.criterion = nn.L1Loss()
|
||||
elif criterion == 'l2' or criterion == 'mse':
|
||||
self.criterion = nn.MSELoss()
|
||||
else:
|
||||
raise ValueError('Criterion %s is not recognized' % criterion)
|
||||
self.resize = resize
|
||||
self.resize_mode = resize_mode
|
||||
self.instance_normalized = instance_normalized
|
||||
|
||||
|
||||
print('Perceptual loss:')
|
||||
print('\tMode: {}'.format(network))
|
||||
|
||||
def forward(self, inp, target, mask=None,use_style_loss=False,weight_style_to_perceptual=0.):
|
||||
r"""Perceptual loss forward.
|
||||
|
||||
Args:
|
||||
inp (4D tensor) : Input tensor.
|
||||
target (4D tensor) : Ground truth tensor, same shape as the input.
|
||||
|
||||
Returns:
|
||||
(scalar tensor) : The perceptual loss.
|
||||
"""
|
||||
# Perceptual loss should operate in eval mode by default.
|
||||
self.model.eval()
|
||||
inp, target = \
|
||||
apply_imagenet_normalization(inp), \
|
||||
apply_imagenet_normalization(target)
|
||||
if self.resize:
|
||||
inp = F.interpolate(
|
||||
inp, mode=self.resize_mode, size=(256, 256),
|
||||
align_corners=False)
|
||||
target = F.interpolate(
|
||||
target, mode=self.resize_mode, size=(256, 256),
|
||||
align_corners=False)
|
||||
|
||||
# Evaluate perceptual loss at each scale.
|
||||
loss = 0
|
||||
style_loss=0
|
||||
for scale in range(self.num_scales):
|
||||
input_features, target_features = \
|
||||
self.model(inp), self.model(target)
|
||||
for layer, weight in zip(self.layers, self.weights):
|
||||
# Example per-layer VGG19 loss values after applying
|
||||
# [0.03125, 0.0625, 0.125, 0.25, 1.0] weighting.
|
||||
# relu_1_1, 0.014698
|
||||
# relu_2_1, 0.085817
|
||||
# relu_3_1, 0.349977
|
||||
# relu_4_1, 0.544188
|
||||
# relu_5_1, 0.906261
|
||||
input_feature = input_features[layer]
|
||||
target_feature = target_features[layer].detach()
|
||||
if self.instance_normalized:
|
||||
input_feature = F.instance_norm(input_feature)
|
||||
target_feature = F.instance_norm(target_feature)
|
||||
|
||||
if mask is not None:
|
||||
mask_ = F.interpolate(mask, input_feature.shape[2:],
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
input_feature = input_feature * mask_
|
||||
target_feature = target_feature * mask_
|
||||
# print('mask',mask_.shape)
|
||||
|
||||
|
||||
loss += weight * self.criterion(input_feature,
|
||||
target_feature)
|
||||
if use_style_loss and scale==0:
|
||||
style_loss += self.criterion(self.compute_gram(input_feature),
|
||||
self.compute_gram(target_feature))
|
||||
|
||||
# Downsample the input and target.
|
||||
if scale != self.num_scales - 1:
|
||||
inp = F.interpolate(
|
||||
inp, mode=self.resize_mode, scale_factor=0.5,
|
||||
align_corners=False, recompute_scale_factor=True)
|
||||
target = F.interpolate(
|
||||
target, mode=self.resize_mode, scale_factor=0.5,
|
||||
align_corners=False, recompute_scale_factor=True)
|
||||
|
||||
if use_style_loss:
|
||||
return loss + style_loss*weight_style_to_perceptual
|
||||
else:
|
||||
return loss
|
||||
|
||||
|
||||
def compute_gram(self, x):
|
||||
b, ch, h, w = x.size()
|
||||
f = x.view(b, ch, w * h)
|
||||
f_T = f.transpose(1, 2)
|
||||
G = f.bmm(f_T) / (h * w * ch)
|
||||
return G
|
||||
|
||||
@@ -0,0 +1,143 @@
|
||||
import os
|
||||
import platform
|
||||
import subprocess
|
||||
import folder_paths
|
||||
from pydub import AudioSegment
|
||||
from moviepy.editor import VideoFileClip,AudioFileClip
|
||||
|
||||
parent_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
from .ip_lap.inference import IP_LAP_infer
|
||||
|
||||
input_path = folder_paths.get_input_directory()
|
||||
out_path = folder_paths.get_output_directory()
|
||||
|
||||
class CombineAudioVideo:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{"vocal_AUDIO": ("AUDIO",),
|
||||
"bgm_AUDIO": ("AUDIO",),
|
||||
"video": ("VIDEO",)
|
||||
}
|
||||
}
|
||||
|
||||
CATEGORY = "AIFSH_IP_LAP"
|
||||
DESCRIPTION = "hello world!"
|
||||
|
||||
RETURN_TYPES = ("VIDEO",)
|
||||
|
||||
OUTPUT_NODE = False
|
||||
|
||||
FUNCTION = "combine"
|
||||
|
||||
def combine(self, vocal_AUDIO,bgm_AUDIO,video):
|
||||
vocal = AudioSegment.from_file(vocal_AUDIO)
|
||||
bgm = AudioSegment.from_file(bgm_AUDIO)
|
||||
audio = vocal.overlay(bgm)
|
||||
audio_file = os.path.join(out_path,"ip_lap_voice.wav")
|
||||
audio.export(audio_file, format="wav")
|
||||
cm_video_file = os.path.join(out_path,"voice_"+os.path.basename(video))
|
||||
video_clip = VideoFileClip(video)
|
||||
audio_clip = AudioFileClip(audio_file)
|
||||
new_video_clip = video_clip.set_audio(audio_clip)
|
||||
new_video_clip.write_videofile(cm_video_file)
|
||||
return (cm_video_file,)
|
||||
|
||||
|
||||
class PreViewVideo:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":{
|
||||
"video":("VIDEO",),
|
||||
}}
|
||||
|
||||
CATEGORY = "AIFSH_IP_LAP"
|
||||
DESCRIPTION = "hello world!"
|
||||
|
||||
RETURN_TYPES = ()
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
FUNCTION = "load_video"
|
||||
|
||||
def load_video(self, video):
|
||||
video_name = os.path.basename(video)
|
||||
video_path_name = os.path.basename(os.path.dirname(video))
|
||||
return {"ui":{"video":[video_name,video_path_name]}}
|
||||
|
||||
class LoadVideo:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
files = [f for f in os.listdir(input_path) if os.path.isfile(os.path.join(input_path, f)) and f.split('.')[-1] in ["mp4", "webm","mkv","avi"]]
|
||||
return {"required":{
|
||||
"video":(files,),
|
||||
}}
|
||||
|
||||
CATEGORY = "AIFSH_IP_LAP"
|
||||
DESCRIPTION = "hello world!"
|
||||
|
||||
RETURN_TYPES = ("VIDEO","AUDIO")
|
||||
|
||||
OUTPUT_NODE = False
|
||||
|
||||
FUNCTION = "load_video"
|
||||
|
||||
def load_video(self, video):
|
||||
video_path = os.path.join(input_path,video)
|
||||
video_clip = VideoFileClip(video_path)
|
||||
audio_path = os.path.join(input_path,video+".wav")
|
||||
video_clip.audio.write_audiofile(audio_path)
|
||||
return (video_path,audio_path,)
|
||||
|
||||
class IP_LAP:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required":
|
||||
{
|
||||
"audio": ("AUDIO",),
|
||||
"video": ("VIDEO",),
|
||||
"T":("INT",{
|
||||
"default": 5,
|
||||
}),
|
||||
"Nl":("INT",{
|
||||
"default": 15,
|
||||
}),
|
||||
"ref_img_N":("INT",{
|
||||
"default": 25,
|
||||
}),
|
||||
"img_size":("INT",{
|
||||
"default": 128,
|
||||
}),
|
||||
"mel_step_size":("INT",{
|
||||
"default": 16,
|
||||
}),
|
||||
"face_det_batch_size":("INT",{
|
||||
"default": 4,
|
||||
}),
|
||||
"checkpoints_path":("STRING",{
|
||||
"default": os.path.join(parent_directory,"weights")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
CATEGORY = "AIFSH_IP_LAP"
|
||||
DESCRIPTION = "hello world!"
|
||||
|
||||
RETURN_TYPES = ("VIDEO",)
|
||||
|
||||
OUTPUT_NODE = False
|
||||
|
||||
FUNCTION = "process"
|
||||
|
||||
def process(self, audio, video, T=5,Nl=15,ref_img_N=25,img_size=128,
|
||||
mel_step_size=16,face_det_batch_size=4,checkpoints_path=""):
|
||||
ip_lap = IP_LAP_infer(T,Nl,ref_img_N,img_size,mel_step_size,face_det_batch_size,checkpoints_path)
|
||||
video_name = os.path.basename(video)
|
||||
out_video_file = os.path.join(out_path, f"ip_lap_{video_name}")
|
||||
ip_lap(video,audio,out_video_file)
|
||||
# res_video_file = os.path.join(out_path, f"result_ip_lap_{video_name}")
|
||||
# command = f'ffmpeg -y -i {out_video_file} -i {audio} -map 0:0 -map 1:0 -c:a libmp3lame -q:a 1 -q:v 1 -shortest {res_video_file}'
|
||||
# subprocess.call(command, shell=platform.system() != 'Windows')
|
||||
return (out_video_file,)
|
||||
@@ -0,0 +1,10 @@
|
||||
1. ModuleNotFoundError: No module named 'torchvision.transforms.functional_tensor'
|
||||
|
||||
from
|
||||
from torchvision.transforms.functional_tensor import rgb_to_grayscale
|
||||
to
|
||||
from torchvision.transforms.functional import rgb_to_grayscale
|
||||
|
||||
2.ImportError: libGL.so.1: cannot open shared object file: No such file or directory
|
||||
apt update
|
||||
apt install ffmpeg -y
|
||||
@@ -0,0 +1,7 @@
|
||||
face_alignment
|
||||
mediapipe
|
||||
basicsr
|
||||
lws
|
||||
librosa
|
||||
moviepy
|
||||
pydub
|
||||
@@ -0,0 +1,155 @@
|
||||
import { app } from "../../../scripts/app.js";
|
||||
import { api } from '../../../scripts/api.js'
|
||||
|
||||
function fitHeight(node) {
|
||||
node.setSize([node.size[0], node.computeSize([node.size[0], node.size[1]])[1]])
|
||||
node?.graph?.setDirtyCanvas(true);
|
||||
}
|
||||
function chainCallback(object, property, callback) {
|
||||
if (object == undefined) {
|
||||
//This should not happen.
|
||||
console.error("Tried to add callback to non-existant object")
|
||||
return;
|
||||
}
|
||||
if (property in object) {
|
||||
const callback_orig = object[property]
|
||||
object[property] = function () {
|
||||
const r = callback_orig.apply(this, arguments);
|
||||
callback.apply(this, arguments);
|
||||
return r
|
||||
};
|
||||
} else {
|
||||
object[property] = callback;
|
||||
}
|
||||
}
|
||||
|
||||
function addPreviewOptions(nodeType) {
|
||||
chainCallback(nodeType.prototype, "getExtraMenuOptions", function(_, options) {
|
||||
// The intended way of appending options is returning a list of extra options,
|
||||
// but this isn't used in widgetInputs.js and would require
|
||||
// less generalization of chainCallback
|
||||
let optNew = []
|
||||
try {
|
||||
const previewWidget = this.widgets.find((w) => w.name === "videopreview");
|
||||
|
||||
let url = null
|
||||
if (previewWidget.videoEl?.hidden == false && previewWidget.videoEl.src) {
|
||||
//Use full quality video
|
||||
//url = api.apiURL('/view?' + new URLSearchParams(previewWidget.value.params));
|
||||
url = previewWidget.videoEl.src
|
||||
}
|
||||
if (url) {
|
||||
optNew.push(
|
||||
{
|
||||
content: "Open preview",
|
||||
callback: () => {
|
||||
window.open(url, "_blank")
|
||||
},
|
||||
},
|
||||
{
|
||||
content: "Save preview",
|
||||
callback: () => {
|
||||
const a = document.createElement("a");
|
||||
a.href = url;
|
||||
a.setAttribute("download", new URLSearchParams(previewWidget.value.params).get("filename"));
|
||||
document.body.append(a);
|
||||
a.click();
|
||||
requestAnimationFrame(() => a.remove());
|
||||
},
|
||||
}
|
||||
);
|
||||
}
|
||||
if(options.length > 0 && options[0] != null && optNew.length > 0) {
|
||||
optNew.push(null);
|
||||
}
|
||||
options.unshift(...optNew);
|
||||
|
||||
} catch (error) {
|
||||
console.log(error);
|
||||
}
|
||||
|
||||
});
|
||||
}
|
||||
function previewVideo(node,file,type){
|
||||
var element = document.createElement("div");
|
||||
const previewNode = node;
|
||||
var previewWidget = node.addDOMWidget("videopreview", "preview", element, {
|
||||
serialize: false,
|
||||
hideOnZoom: false,
|
||||
getValue() {
|
||||
return element.value;
|
||||
},
|
||||
setValue(v) {
|
||||
element.value = v;
|
||||
},
|
||||
});
|
||||
previewWidget.computeSize = function(width) {
|
||||
if (this.aspectRatio && !this.parentEl.hidden) {
|
||||
let height = (previewNode.size[0]-20)/ this.aspectRatio + 10;
|
||||
if (!(height > 0)) {
|
||||
height = 0;
|
||||
}
|
||||
this.computedHeight = height + 10;
|
||||
return [width, height];
|
||||
}
|
||||
return [width, -4];//no loaded src, widget should not display
|
||||
}
|
||||
// element.style['pointer-events'] = "none"
|
||||
previewWidget.value = {hidden: false, paused: false, params: {}}
|
||||
previewWidget.parentEl = document.createElement("div");
|
||||
previewWidget.parentEl.className = "video_preview";
|
||||
previewWidget.parentEl.style['width'] = "100%"
|
||||
element.appendChild(previewWidget.parentEl);
|
||||
previewWidget.videoEl = document.createElement("video");
|
||||
previewWidget.videoEl.controls = true;
|
||||
previewWidget.videoEl.loop = false;
|
||||
previewWidget.videoEl.muted = false;
|
||||
previewWidget.videoEl.style['width'] = "100%"
|
||||
previewWidget.videoEl.addEventListener("loadedmetadata", () => {
|
||||
|
||||
previewWidget.aspectRatio = previewWidget.videoEl.videoWidth / previewWidget.videoEl.videoHeight;
|
||||
fitHeight(this);
|
||||
});
|
||||
previewWidget.videoEl.addEventListener("error", () => {
|
||||
//TODO: consider a way to properly notify the user why a preview isn't shown.
|
||||
previewWidget.parentEl.hidden = true;
|
||||
fitHeight(this);
|
||||
});
|
||||
|
||||
let params = {
|
||||
"filename": file,
|
||||
"type": type,
|
||||
}
|
||||
|
||||
previewWidget.parentEl.hidden = previewWidget.value.hidden;
|
||||
previewWidget.videoEl.autoplay = !previewWidget.value.paused && !previewWidget.value.hidden;
|
||||
let target_width = 256
|
||||
if (element.style?.width) {
|
||||
//overscale to allow scrolling. Endpoint won't return higher than native
|
||||
target_width = element.style.width.slice(0,-2)*2;
|
||||
}
|
||||
if (!params.force_size || params.force_size.includes("?") || params.force_size == "Disabled") {
|
||||
params.force_size = target_width+"x?"
|
||||
} else {
|
||||
let size = params.force_size.split("x")
|
||||
let ar = parseInt(size[0])/parseInt(size[1])
|
||||
params.force_size = target_width+"x"+(target_width/ar)
|
||||
}
|
||||
|
||||
previewWidget.videoEl.src = api.apiURL('/view?' + new URLSearchParams(params));
|
||||
|
||||
previewWidget.videoEl.hidden = false;
|
||||
previewWidget.parentEl.appendChild(previewWidget.videoEl)
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: "IP_LAP.VideoPreviewer",
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||
if (nodeData?.name == "PreViewVideo") {
|
||||
nodeType.prototype.onExecuted = function (data) {
|
||||
previewVideo(this, data.video[0], data.video[1]);
|
||||
}
|
||||
addPreviewOptions(nodeType)
|
||||
}
|
||||
}
|
||||
});
|
||||
@@ -0,0 +1,203 @@
|
||||
import { app } from "../../../scripts/app.js";
|
||||
import { api } from '../../../scripts/api.js'
|
||||
import { ComfyWidgets } from "../../../scripts/widgets.js"
|
||||
|
||||
function fitHeight(node) {
|
||||
node.setSize([node.size[0], node.computeSize([node.size[0], node.size[1]])[1]])
|
||||
node?.graph?.setDirtyCanvas(true);
|
||||
}
|
||||
|
||||
function previewVideo(node,file){
|
||||
while (node.widgets.length > 2){
|
||||
node.widgets.pop()
|
||||
}
|
||||
try {
|
||||
var el = document.getElementById("uploadVideo");
|
||||
el.remove();
|
||||
} catch (error) {
|
||||
console.log(error);
|
||||
}
|
||||
var element = document.createElement("div");
|
||||
element.id = "uploadVideo";
|
||||
const previewNode = node;
|
||||
var previewWidget = node.addDOMWidget("videopreview", "preview", element, {
|
||||
serialize: false,
|
||||
hideOnZoom: false,
|
||||
getValue() {
|
||||
return element.value;
|
||||
},
|
||||
setValue(v) {
|
||||
element.value = v;
|
||||
},
|
||||
});
|
||||
previewWidget.computeSize = function(width) {
|
||||
if (this.aspectRatio && !this.parentEl.hidden) {
|
||||
let height = (previewNode.size[0]-20)/ this.aspectRatio + 10;
|
||||
if (!(height > 0)) {
|
||||
height = 0;
|
||||
}
|
||||
this.computedHeight = height + 10;
|
||||
return [width, height];
|
||||
}
|
||||
return [width, -4];//no loaded src, widget should not display
|
||||
}
|
||||
// element.style['pointer-events'] = "none"
|
||||
previewWidget.value = {hidden: false, paused: false, params: {}}
|
||||
previewWidget.parentEl = document.createElement("div");
|
||||
previewWidget.parentEl.className = "video_preview";
|
||||
previewWidget.parentEl.style['width'] = "100%"
|
||||
element.appendChild(previewWidget.parentEl);
|
||||
previewWidget.videoEl = document.createElement("video");
|
||||
previewWidget.videoEl.controls = true;
|
||||
previewWidget.videoEl.loop = false;
|
||||
previewWidget.videoEl.muted = false;
|
||||
previewWidget.videoEl.style['width'] = "100%"
|
||||
previewWidget.videoEl.addEventListener("loadedmetadata", () => {
|
||||
|
||||
previewWidget.aspectRatio = previewWidget.videoEl.videoWidth / previewWidget.videoEl.videoHeight;
|
||||
fitHeight(this);
|
||||
});
|
||||
previewWidget.videoEl.addEventListener("error", () => {
|
||||
//TODO: consider a way to properly notify the user why a preview isn't shown.
|
||||
previewWidget.parentEl.hidden = true;
|
||||
fitHeight(this);
|
||||
});
|
||||
|
||||
let params = {
|
||||
"filename": file,
|
||||
"type": "input",
|
||||
}
|
||||
|
||||
previewWidget.parentEl.hidden = previewWidget.value.hidden;
|
||||
previewWidget.videoEl.autoplay = !previewWidget.value.paused && !previewWidget.value.hidden;
|
||||
let target_width = 256
|
||||
if (element.style?.width) {
|
||||
//overscale to allow scrolling. Endpoint won't return higher than native
|
||||
target_width = element.style.width.slice(0,-2)*2;
|
||||
}
|
||||
if (!params.force_size || params.force_size.includes("?") || params.force_size == "Disabled") {
|
||||
params.force_size = target_width+"x?"
|
||||
} else {
|
||||
let size = params.force_size.split("x")
|
||||
let ar = parseInt(size[0])/parseInt(size[1])
|
||||
params.force_size = target_width+"x"+(target_width/ar)
|
||||
}
|
||||
|
||||
previewWidget.videoEl.src = api.apiURL('/view?' + new URLSearchParams(params));
|
||||
|
||||
previewWidget.videoEl.hidden = false;
|
||||
previewWidget.parentEl.appendChild(previewWidget.videoEl)
|
||||
}
|
||||
|
||||
function videoUpload(node, inputName, inputData, app) {
|
||||
const videoWidget = node.widgets.find((w) => w.name === "video");
|
||||
let uploadWidget;
|
||||
/*
|
||||
A method that returns the required style for the html
|
||||
*/
|
||||
var default_value = videoWidget.value;
|
||||
Object.defineProperty(videoWidget, "value", {
|
||||
set : function(value) {
|
||||
this._real_value = value;
|
||||
},
|
||||
|
||||
get : function() {
|
||||
let value = "";
|
||||
if (this._real_value) {
|
||||
value = this._real_value;
|
||||
} else {
|
||||
return default_value;
|
||||
}
|
||||
|
||||
if (value.filename) {
|
||||
let real_value = value;
|
||||
value = "";
|
||||
if (real_value.subfolder) {
|
||||
value = real_value.subfolder + "/";
|
||||
}
|
||||
|
||||
value += real_value.filename;
|
||||
|
||||
if(real_value.type && real_value.type !== "input")
|
||||
value += ` [${real_value.type}]`;
|
||||
}
|
||||
return value;
|
||||
}
|
||||
});
|
||||
async function uploadFile(file, updateNode, pasted = false) {
|
||||
try {
|
||||
// Wrap file in formdata so it includes filename
|
||||
const body = new FormData();
|
||||
body.append("image", file);
|
||||
if (pasted) body.append("subfolder", "pasted");
|
||||
const resp = await api.fetchApi("/upload/image", {
|
||||
method: "POST",
|
||||
body,
|
||||
});
|
||||
|
||||
if (resp.status === 200) {
|
||||
const data = await resp.json();
|
||||
// Add the file to the dropdown list and update the widget value
|
||||
let path = data.name;
|
||||
if (data.subfolder) path = data.subfolder + "/" + path;
|
||||
|
||||
if (!videoWidget.options.values.includes(path)) {
|
||||
videoWidget.options.values.push(path);
|
||||
}
|
||||
|
||||
if (updateNode) {
|
||||
videoWidget.value = path;
|
||||
previewVideo(node,path)
|
||||
|
||||
}
|
||||
} else {
|
||||
alert(resp.status + " - " + resp.statusText);
|
||||
}
|
||||
} catch (error) {
|
||||
alert(error);
|
||||
}
|
||||
}
|
||||
|
||||
const fileInput = document.createElement("input");
|
||||
Object.assign(fileInput, {
|
||||
type: "file",
|
||||
accept: "video/webm,video/mp4,video/mkv,video/avi",
|
||||
style: "display: none",
|
||||
onchange: async () => {
|
||||
if (fileInput.files.length) {
|
||||
await uploadFile(fileInput.files[0], true);
|
||||
}
|
||||
},
|
||||
});
|
||||
document.body.append(fileInput);
|
||||
|
||||
// Create the button widget for selecting the files
|
||||
uploadWidget = node.addWidget("button", "choose video file to upload", "Video", () => {
|
||||
fileInput.click();
|
||||
});
|
||||
|
||||
uploadWidget.serialize = false;
|
||||
|
||||
previewVideo(node, videoWidget.value);
|
||||
const cb = node.callback;
|
||||
videoWidget.callback = function () {
|
||||
previewVideo(node,videoWidget.value);
|
||||
if (cb) {
|
||||
return cb.apply(this, arguments);
|
||||
}
|
||||
};
|
||||
|
||||
return { widget: uploadWidget };
|
||||
}
|
||||
|
||||
ComfyWidgets.VIDEOPLOAD = videoUpload;
|
||||
|
||||
app.registerExtension({
|
||||
name: "IP_LAP.UploadVideo",
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||
if (nodeData?.name == "LoadVideo") {
|
||||
nodeData.input.required.upload = ["VIDEOPLOAD"];
|
||||
}
|
||||
},
|
||||
});
|
||||
|
||||
Reference in New Issue
Block a user