CoolFace
Apppublic

thaint2901/talking-head-generation-deploy

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
app.py122 linesDownload Raw Back to root
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)