thaint2901/talking-head-generation-deploy
0
1import av2import sys3import numpy as np4import cv25import streamlit as st6from PIL import Image7from streamlit_webrtc import WebRtcMode, webrtc_streamer8 9sys.path.insert(1, "./retinaface")10sys.path.insert(1, "./TPSMM/pkgs")11from tpsmm import TPSMM12from detect import Detect13from turn import get_ice_servers14 15 16def parse_roi_box_from_bbox(bbox, shape):17 img_h, img_w = shape[:2]18 left, top, right, bottom = bbox[:4]19 old_size = (right - left + bottom - top) / 220 center_x = right - (right - left) / 2.021 center_y = bottom - (bottom - top) / 2.0 + old_size * 0.1422 23 size = int(min((old_size * 2.0) / 2, center_x, img_w-center_x, center_y, img_h-center_y) * 2.0)24 25 roi_box = [0] * 426 roi_box[0] = center_x - size / 227 roi_box[1] = center_y - size / 228 roi_box[2] = roi_box[0] + size29 roi_box[3] = roi_box[1] + size30 31 return roi_box32 33cache_key = "retinaface"34if cache_key in st.session_state:35 detector = st.session_state[cache_key]36else:37 detector = Detect("./retinaface/weights/mobilenet0.25_epoch_842.pth", net_inshape=(486, 864))38 st.session_state[cache_key] = detector39 40cache_key = "tpsmm"41if cache_key in st.session_state:42 generator = st.session_state[cache_key]43else:44 generator = TPSMM()45 st.session_state[cache_key] = generator46 47 48@st.cache_resource # type: ignore49def get_images():50 images = [51 cv2.imread("assets/0.jpg"),52 cv2.imread("assets/1.jpg"),53 cv2.imread("assets/2.jpg"),54 cv2.imread("assets/3.jpg"),55 ]56 item_list = [str(i) for i in range(len(images))]57 images = [generator.process_source(src_img) for src_img in images]58 59 return dict(zip(item_list, images))60images = get_images()61user_option = st.selectbox("Choose an item", list(images.keys()))62 63uploaded_file = st.file_uploader("Or upload your file here...", type=['png', 'jpeg', 'jpg'])64@st.cache_resource65def process_file(uploaded_file):66 img = Image.open(uploaded_file)67 img = cv2.cvtColor(np.array(img), cv2.COLOR_RGB2BGR)68 dets = detector(img)69 for i, b in enumerate(dets):70 bbox = parse_roi_box_from_bbox(b[:4], img.shape)71 bbox = [int(i) for i in bbox]72 73 face_img = img[bbox[1]:bbox[3], bbox[0]:bbox[2]].copy()74 # cv2.imwrite("./tmp.jpg", face_img)75 return generator.process_source(face_img)76 77 return None78if uploaded_file is not None:79 uploaded_file = process_file(uploaded_file)80 81def callback(frame: av.VideoFrame) -> av.VideoFrame:82 img = frame.to_ndarray(format="bgr24")83 84 try:85 dets = detector(img)86 output = None87 for i, b in enumerate(dets):88 text = "{:.4f}".format(b[4])89 b = b.astype(np.int32)90 cv2.rectangle(img, (b[0], b[1]), (b[2], b[3]), (0, 0, 255), 2)91 bbox = parse_roi_box_from_bbox(b[:4], img.shape)92 bbox = [int(i) for i in bbox]93 cv2.rectangle(img, (bbox[0], bbox[1]), (bbox[2], bbox[3]), (255, 0, 0), 2)94 95 face_img = img[bbox[1]:bbox[3], bbox[0]:bbox[2]].copy()96 if uploaded_file is None:97 source_tensor, kp_source = images[user_option]98 else:99 source_tensor, kp_source = uploaded_file100 output = generator.gen_image(face_img, source_tensor, kp_source)101 102 landm = b[5:15]103 landm = landm.reshape((5, 2))104 cv2.circle(img, tuple(landm[0]), 1, (0, 0, 255), 2)105 cv2.circle(img, tuple(landm[1]), 1, (0, 255, 255), 2)106 cv2.circle(img, tuple(landm[2]), 1, (255, 0, 255), 2)107 cv2.circle(img, tuple(landm[3]), 1, (0, 255, 0), 2)108 cv2.circle(img, tuple(landm[4]), 1, (255, 0, 0), 2)109 110 if output is not None:111 img[:256, :256] = output112 except Exception as e:113 print(e)114 115 return av.VideoFrame.from_ndarray(img, format="bgr24")116 117webrtc_streamer(118 key="sample",119 rtc_configuration={"iceServers": get_ice_servers()},120 video_frame_callback=callback,121 media_stream_constraints={"video": True, "audio": False},122)