diff --git a/facefusion/apis/endpoints/session.py b/facefusion/apis/endpoints/session.py index 39970954..49e97c14 100644 --- a/facefusion/apis/endpoints/session.py +++ b/facefusion/apis/endpoints/session.py @@ -4,7 +4,7 @@ from starlette.requests import Request from starlette.responses import JSONResponse from starlette.status import HTTP_200_OK, HTTP_201_CREATED, HTTP_401_UNAUTHORIZED, HTTP_404_NOT_FOUND -from facefusion import content_store, inference_manager, process_manager, session_context, session_manager, state_manager, translator, video_manager +from facefusion import content_store, face_store, inference_manager, process_manager, session_context, session_manager, state_manager, translator, video_manager from facefusion.apis import asset_store from facefusion.apis.session_helper import validate_api_key from facefusion.apis.stream_manager import destroy_stream @@ -22,6 +22,7 @@ async def create_session(request : Request) -> JSONResponse: state_manager.init() content_store.init() + face_store.init() inference_manager.init() video_manager.init() process_manager.init() @@ -88,6 +89,7 @@ async def destroy_session(request : Request) -> JSONResponse: state_manager.clear() content_store.clear() + face_store.clear() inference_manager.clear() video_manager.clear() process_manager.clear() diff --git a/facefusion/benchmarker.py b/facefusion/benchmarker.py index 8b5bbe0f..15c0d859 100644 --- a/facefusion/benchmarker.py +++ b/facefusion/benchmarker.py @@ -6,10 +6,9 @@ from time import perf_counter from typing import Iterator, List import facefusion.choices -from facefusion import content_analyser, core, state_manager +from facefusion import content_analyser, core, face_store, state_manager from facefusion.cli_helper import render_table from facefusion.download import conditional_download, resolve_download_url -from facefusion.face_store import clear_faces from facefusion.filesystem import get_file_extension from facefusion.types import BenchmarkCycleSet from facefusion.vision import count_video_frame_total, detect_video_fps @@ -64,7 +63,7 @@ def cycle(cycle_count : int) -> BenchmarkCycleSet: if state_manager.get_item('benchmark_mode') == 'cold': content_analyser.analyse_image.cache_clear() content_analyser.analyse_video.cache_clear() - clear_faces() + face_store.clear() start_time = perf_counter() core.conditional_process() diff --git a/facefusion/core.py b/facefusion/core.py index 1578e4a1..bb43d402 100755 --- a/facefusion/core.py +++ b/facefusion/core.py @@ -8,7 +8,7 @@ from time import time import uvicorn import facefusion.apis.core -from facefusion import args_helper, benchmarker, cli_helper, content_analyser, content_store, hash_helper, inference_manager, logger, process_manager, session_manager, state_manager, translator, video_manager +from facefusion import args_helper, benchmarker, cli_helper, content_analyser, content_store, face_store, hash_helper, inference_manager, logger, process_manager, session_manager, state_manager, translator, video_manager from facefusion.args_helper import apply_args from facefusion.download import conditional_download_hashes, conditional_download_sources from facefusion.exit_helper import hard_exit, signal_exit @@ -39,6 +39,7 @@ def cli() -> None: logger.init(state_manager.get_item('log_level')) content_store.init() + face_store.init() inference_manager.init() process_manager.init() video_manager.init() diff --git a/facefusion/face_store.py b/facefusion/face_store.py index 8042735d..246959f9 100644 --- a/facefusion/face_store.py +++ b/facefusion/face_store.py @@ -1,41 +1,58 @@ import threading from typing import List, Optional +from facefusion import store_creator from facefusion.hash_helper import create_hash -from facefusion.types import Face, FaceStore, VisionFrame +from facefusion.session_manager import resolve_owner_id +from facefusion.types import Face, Store, VisionFrame from facefusion.vision import is_vision_frame -FACE_STORE : FaceStore = {} +FACE_STORE : Store = store_creator.create_store({}) + + +def init() -> None: + owner_id = resolve_owner_id() + store_creator.init_content(FACE_STORE, owner_id) def get_faces(vision_frame : VisionFrame) -> Optional[List[Face]]: + owner_id = resolve_owner_id() + face_store = store_creator.get_content(FACE_STORE, owner_id) + if is_vision_frame(vision_frame): vision_hash = create_hash(vision_frame.tobytes()) - if FACE_STORE.get(vision_hash): - return FACE_STORE.get(vision_hash).get('faces') + if face_store.get(vision_hash): + return face_store.get(vision_hash).get('faces') return None def set_faces(vision_frame : VisionFrame, faces : List[Face]) -> None: + owner_id = resolve_owner_id() + face_store = store_creator.get_content(FACE_STORE, owner_id) + if is_vision_frame(vision_frame): vision_hash = create_hash(vision_frame.tobytes()) - FACE_STORE.setdefault(vision_hash, + face_store.setdefault(vision_hash, { 'lock': threading.Lock() })['faces'] = faces def resolve_lock(vision_frame : VisionFrame) -> threading.Lock: + owner_id = resolve_owner_id() + face_store = store_creator.get_content(FACE_STORE, owner_id) + if is_vision_frame(vision_frame): vision_hash = create_hash(vision_frame.tobytes()) - return FACE_STORE.setdefault(vision_hash, + return face_store.setdefault(vision_hash, { 'lock': threading.Lock() }).get('lock') return threading.Lock() -def clear_faces() -> None: - FACE_STORE.clear() +def clear() -> None: + owner_id = resolve_owner_id() + store_creator.init_content(FACE_STORE, owner_id) diff --git a/tests/test_face_creator.py b/tests/test_face_creator.py index b24e6cc2..83f0e9ae 100644 --- a/tests/test_face_creator.py +++ b/tests/test_face_creator.py @@ -2,10 +2,9 @@ import numpy import pytest -from facefusion import face_aligner, face_classifier, face_detector, face_recognizer, ffmpeg, ffmpeg_builder, inference_manager, process_manager, state_manager +from facefusion import face_aligner, face_classifier, face_detector, face_recognizer, face_store, ffmpeg, ffmpeg_builder, inference_manager, process_manager, state_manager from facefusion.download import conditional_download from facefusion.face_creator import average_face_geometry, get_many_faces, get_one_face, refill_faces -from facefusion.face_store import clear_faces from facefusion.vision import read_static_image from .assert_helper import get_test_example_file, get_test_examples_directory @@ -15,6 +14,8 @@ def before_all() -> None: state_manager.init() inference_manager.init() + face_store.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ @@ -56,7 +57,8 @@ def before_each() -> None: face_detector.clear_inference_pool() face_aligner.clear_inference_pool() face_recognizer.clear_inference_pool() - clear_faces() + + face_store.clear() def test_get_one_face() -> None: diff --git a/tests/test_face_store.py b/tests/test_face_store.py new file mode 100644 index 00000000..383dc793 --- /dev/null +++ b/tests/test_face_store.py @@ -0,0 +1,74 @@ +from typing import Iterator + +import numpy +import pytest + +from facefusion import session_context, session_manager +from facefusion.face_store import clear, get_faces, init, resolve_lock, set_faces + + +@pytest.fixture(scope = 'function', autouse = True) +def before_each() -> Iterator[None]: + local_id = session_context.resolve_local_id() + + session_context.set_session_id(local_id) + init() + + yield + + session_context.set_session_id(local_id) + + +def test_init() -> None: + local_id = session_context.resolve_local_id() + vision_frame = numpy.zeros((16, 16, 3), numpy.uint8) + + set_faces(vision_frame, []) + session_context.set_session_id('session-a') + init() + + assert get_faces(vision_frame) is None + + set_faces(vision_frame, []) + session_manager.fork_session() + + assert get_faces(vision_frame) == [] + + session_manager.join_session() + session_context.set_session_id(local_id) + + assert get_faces(vision_frame) == [] + + +def test_get_faces() -> None: + vision_frame = numpy.zeros((16, 16, 3), numpy.uint8) + + assert get_faces(vision_frame) is None + + set_faces(vision_frame, []) + + assert get_faces(vision_frame) == [] + + +def test_set_faces() -> None: + vision_frame = numpy.zeros((16, 16, 3), numpy.uint8) + + set_faces(vision_frame, []) + set_faces(vision_frame, []) + + assert get_faces(vision_frame) == [] + + +def test_resolve_lock() -> None: + vision_frame = numpy.zeros((16, 16, 3), numpy.uint8) + + assert resolve_lock(vision_frame) is resolve_lock(vision_frame) + + +def test_clear() -> None: + vision_frame = numpy.zeros((16, 16, 3), numpy.uint8) + + set_faces(vision_frame, []) + clear() + + assert get_faces(vision_frame) is None diff --git a/tests/test_face_tracker.py b/tests/test_face_tracker.py index 05909712..2c607da6 100644 --- a/tests/test_face_tracker.py +++ b/tests/test_face_tracker.py @@ -1,11 +1,10 @@ import numpy import pytest -from facefusion import face_aligner, face_classifier, face_detector, face_recognizer, inference_manager, state_manager, video_manager +from facefusion import face_aligner, face_classifier, face_detector, face_recognizer, face_store, inference_manager, state_manager, video_manager from facefusion.common_helper import get_first, get_last from facefusion.download import conditional_download from facefusion.face_creator import get_many_faces, get_one_face -from facefusion.face_store import clear_faces from facefusion.face_tracker import create_face_tracks, select_face_track, track_faces from facefusion.vision import read_static_video_frame, select_video_frames from .assert_helper import get_test_example_file, get_test_examples_directory @@ -16,6 +15,8 @@ def before_all() -> None: state_manager.init() inference_manager.init() + face_store.init() + video_manager.init() conditional_download(get_test_examples_directory(), @@ -47,7 +48,8 @@ def before_each() -> None: face_detector.clear_inference_pool() face_aligner.clear_inference_pool() face_recognizer.clear_inference_pool() - clear_faces() + + face_store.clear() def test_track_faces() -> None: