Files
facefusion-labs/face_swapper/src/inferencing.py
T

41 lines
1.6 KiB
Python

import configparser
import cv2
import torch
from .helper import calc_embedding, convert_to_vision_frame, convert_to_vision_tensor
from .models.generator import Generator
from .types import EmbedderModule, GeneratorModule, VisionFrame
CONFIG = configparser.ConfigParser()
CONFIG.read('config.ini')
def run_swap(generator : GeneratorModule, embedder : EmbedderModule, source_vision_frame : VisionFrame, target_vision_frame : VisionFrame) -> VisionFrame:
source_vision_tensor = convert_to_vision_tensor(source_vision_frame)
target_vision_tensor = convert_to_vision_tensor(target_vision_frame)
source_embedding = calc_embedding(embedder, source_vision_tensor, (0, 0, 0, 0))
output_vision_tensor = generator(source_embedding, target_vision_tensor)[0]
output_vision_frame = convert_to_vision_frame(output_vision_tensor)
return output_vision_frame
def infer() -> None:
generator_path = CONFIG.get('inferencing', 'generator_path')
embedder_path = CONFIG.get('inferencing', 'embedder_path')
source_path = CONFIG.get('inferencing', 'source_path')
target_path = CONFIG.get('inferencing', 'target_path')
output_path = CONFIG.get('inferencing', 'output_path')
state_dict = torch.load(generator_path).get('state_dict').get('generator')
generator = Generator()
generator.load_state_dict(state_dict)
generator.eval()
embedder = torch.jit.load(embedder_path, map_location = 'cpu') # type:ignore[no-untyped-call]
embedder.eval()
source_vision_frame = cv2.imread(source_path)
target_vision_frame = cv2.imread(target_path)
output_vision_frame = run_swap(generator, embedder, source_vision_frame, target_vision_frame)
cv2.imwrite(output_path, output_vision_frame)