diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..a77e58e --- /dev/null +++ b/__init__.py @@ -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" +} diff --git a/ip_lap/__pycache__/draw_landmark.cpython-310.pyc b/ip_lap/__pycache__/draw_landmark.cpython-310.pyc new file mode 100644 index 0000000..24d7dc7 Binary files /dev/null and b/ip_lap/__pycache__/draw_landmark.cpython-310.pyc differ diff --git a/ip_lap/__pycache__/face_mask.cpython-310.pyc b/ip_lap/__pycache__/face_mask.cpython-310.pyc new file mode 100644 index 0000000..29e8dd0 Binary files /dev/null and b/ip_lap/__pycache__/face_mask.cpython-310.pyc differ diff --git a/ip_lap/__pycache__/inference.cpython-310.pyc b/ip_lap/__pycache__/inference.cpython-310.pyc new file mode 100644 index 0000000..5e82e83 Binary files /dev/null and b/ip_lap/__pycache__/inference.cpython-310.pyc differ diff --git a/ip_lap/draw_landmark.py b/ip_lap/draw_landmark.py new file mode 100644 index 0000000..5e191a1 --- /dev/null +++ b/ip_lap/draw_landmark.py @@ -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) diff --git a/ip_lap/face_mask.py b/ip_lap/face_mask.py new file mode 100644 index 0000000..1b1b4de --- /dev/null +++ b/ip_lap/face_mask.py @@ -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) diff --git a/ip_lap/inference.py b/ip_lap/inference.py new file mode 100644 index 0000000..88f4cc4 --- /dev/null +++ b/ip_lap/inference.py @@ -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 \ No newline at end of file diff --git a/ip_lap/models/__init__.py b/ip_lap/models/__init__.py new file mode 100644 index 0000000..6e4582c --- /dev/null +++ b/ip_lap/models/__init__.py @@ -0,0 +1,4 @@ +from . import audio +from .landmark_generator import Landmark_generator +from .video_renderer import Renderer +from .pix2pixHD_disc import define_D diff --git a/ip_lap/models/__pycache__/__init__.cpython-310.pyc b/ip_lap/models/__pycache__/__init__.cpython-310.pyc new file mode 100644 index 0000000..3d31575 Binary files /dev/null and b/ip_lap/models/__pycache__/__init__.cpython-310.pyc differ diff --git a/ip_lap/models/__pycache__/audio.cpython-310.pyc b/ip_lap/models/__pycache__/audio.cpython-310.pyc new file mode 100644 index 0000000..a46f8eb Binary files /dev/null and b/ip_lap/models/__pycache__/audio.cpython-310.pyc differ diff --git a/ip_lap/models/__pycache__/landmark_generator.cpython-310.pyc b/ip_lap/models/__pycache__/landmark_generator.cpython-310.pyc new file mode 100644 index 0000000..26647da Binary files /dev/null and b/ip_lap/models/__pycache__/landmark_generator.cpython-310.pyc differ diff --git a/ip_lap/models/__pycache__/pix2pixHD_disc.cpython-310.pyc b/ip_lap/models/__pycache__/pix2pixHD_disc.cpython-310.pyc new file mode 100644 index 0000000..4dafe5b Binary files /dev/null and b/ip_lap/models/__pycache__/pix2pixHD_disc.cpython-310.pyc differ diff --git a/ip_lap/models/__pycache__/video_renderer.cpython-310.pyc b/ip_lap/models/__pycache__/video_renderer.cpython-310.pyc new file mode 100644 index 0000000..aff043a Binary files /dev/null and b/ip_lap/models/__pycache__/video_renderer.cpython-310.pyc differ diff --git a/ip_lap/models/audio.py b/ip_lap/models/audio.py new file mode 100644 index 0000000..615f38d --- /dev/null +++ b/ip_lap/models/audio.py @@ -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 ) + + 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) diff --git a/ip_lap/models/landmark_generator.py b/ip_lap/models/landmark_generator.py new file mode 100644 index 0000000..7d41bae --- /dev/null +++ b/ip_lap/models/landmark_generator.py @@ -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) + diff --git a/ip_lap/models/pix2pixHD_disc.py b/ip_lap/models/pix2pixHD_disc.py new file mode 100644 index 0000000..6a13ea5 --- /dev/null +++ b/ip_lap/models/pix2pixHD_disc.py @@ -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 + diff --git a/ip_lap/models/video_renderer.py b/ip_lap/models/video_renderer.py new file mode 100644 index 0000000..9fb2443 --- /dev/null +++ b/ip_lap/models/video_renderer.py @@ -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 + diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..fe81d67 --- /dev/null +++ b/nodes.py @@ -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,) \ No newline at end of file diff --git a/note.txt b/note.txt new file mode 100644 index 0000000..738f613 --- /dev/null +++ b/note.txt @@ -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 \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..ace7f0f --- /dev/null +++ b/requirements.txt @@ -0,0 +1,7 @@ +face_alignment +mediapipe +basicsr +lws +librosa +moviepy +pydub \ No newline at end of file diff --git a/web/js/previewVideo.js b/web/js/previewVideo.js new file mode 100644 index 0000000..5c8f1ca --- /dev/null +++ b/web/js/previewVideo.js @@ -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) + } + } +}); \ No newline at end of file diff --git a/web/js/uploadVideo.js b/web/js/uploadVideo.js new file mode 100644 index 0000000..1c92ce4 --- /dev/null +++ b/web/js/uploadVideo.js @@ -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"]; + } + }, +}); +