diff --git a/facefusion/apis/asset_helper.py b/facefusion/apis/asset_helper.py index 2cb4f300..c4766a7b 100644 --- a/facefusion/apis/asset_helper.py +++ b/facefusion/apis/asset_helper.py @@ -8,8 +8,8 @@ from starlette.datastructures import UploadFile import facefusion.choices from facefusion import ffmpeg, process_manager, state_manager from facefusion.filesystem import create_directory, get_file_extension, get_file_format, is_audio, is_image, is_video -from facefusion.types import ImageMetadata, MediaType -from facefusion.vision import detect_image_resolution +from facefusion.types import ImageAsset, ImageMetadata, MediaType, VideoAsset, VisionFrame +from facefusion.vision import detect_image_resolution, fit_contain_frame, is_vision_frame, read_static_image, read_static_video_frame, unpack_resolution def extract_image_metadata(file_path : str) -> ImageMetadata: @@ -89,3 +89,34 @@ async def save_asset_files(upload_files : List[UploadFile]) -> List[str]: process_manager.end() return asset_paths + + +def read_asset_frames(asset : ImageAsset | VideoAsset, frame_numbers : List[str]) -> List[VisionFrame]: + vision_frames = [] + + if asset.get('media') == 'image': + vision_frame = read_static_image(asset.get('path')) + + if is_vision_frame(vision_frame): + vision_frames.append(vision_frame) + + if asset.get('media') == 'video': + for frame_number in frame_numbers: + if frame_number.isdigit(): + vision_frame = read_static_video_frame(asset.get('path'), int(frame_number)) + + if is_vision_frame(vision_frame): + vision_frames.append(vision_frame) + + return vision_frames + + +def capture_asset_frames(asset : ImageAsset | VideoAsset, frame_numbers : List[str], resolution : str) -> List[VisionFrame]: + capture_vision_frames = [] + temp_vision_frames = read_asset_frames(asset, frame_numbers) + + for temp_vision_frame in temp_vision_frames: + capture_vision_frame = fit_contain_frame(temp_vision_frame, unpack_resolution(resolution)) + capture_vision_frames.append(capture_vision_frame) + + return capture_vision_frames diff --git a/facefusion/apis/endpoints/assets.py b/facefusion/apis/endpoints/assets.py index 621b0e55..0e267efd 100644 --- a/facefusion/apis/endpoints/assets.py +++ b/facefusion/apis/endpoints/assets.py @@ -7,9 +7,10 @@ from starlette.status import HTTP_200_OK, HTTP_201_CREATED, HTTP_400_BAD_REQUEST from facefusion import session_context, session_manager from facefusion.apis import asset_store -from facefusion.apis.asset_helper import save_asset_files, validate_asset_files -from facefusion.apis.endpoints.session import extract_access_token +from facefusion.apis.asset_helper import capture_asset_frames, save_asset_files, validate_asset_files +from facefusion.apis.session_helper import extract_access_token from facefusion.filesystem import remove_file +from facefusion.vision import is_vision_frames, to_strip_buffer async def upload_asset(request : Request) -> Response: @@ -89,6 +90,16 @@ async def get_asset(request : Request) -> Response: asset = asset_store.get_asset(session_id, asset_id) if asset: + if asset.get('media') in [ 'image', 'video' ] and request.query_params.get('action') == 'capture' and request.query_params.get('subject') == 'frame': + resolution = request.query_params.get('resolution') + frame_numbers = request.query_params.getlist('frame_number') + vision_frames = capture_asset_frames(asset, frame_numbers, resolution) #type:ignore[arg-type] + + if is_vision_frames(vision_frames): + return Response(content = to_strip_buffer(vision_frames), media_type = 'image/jpeg') + + return Response(status_code = HTTP_400_BAD_REQUEST) + if request.query_params.get('action') == 'download': asset_path = asset.get('path') diff --git a/facefusion/apis/endpoints/state.py b/facefusion/apis/endpoints/state.py index ba121c3c..6ac5463f 100644 --- a/facefusion/apis/endpoints/state.py +++ b/facefusion/apis/endpoints/state.py @@ -4,7 +4,7 @@ from starlette.status import HTTP_200_OK, HTTP_400_BAD_REQUEST, HTTP_404_NOT_FOU from facefusion import args_helper, capability_store, session_manager, state_manager, translator from facefusion.apis import asset_store -from facefusion.apis.endpoints.session import extract_access_token +from facefusion.apis.session_helper import extract_access_token async def get_state(request : Request) -> JSONResponse: diff --git a/facefusion/apis/stream_event.py b/facefusion/apis/stream_event.py index 5da36040..64655b2e 100644 --- a/facefusion/apis/stream_event.py +++ b/facefusion/apis/stream_event.py @@ -21,9 +21,7 @@ def create_receive_event(track : int, frame_handler : FrameHandler) -> threading def dispatch_frame(frame_handler : FrameHandler, track : int, data : ctypes.c_void_p, size : int, info : ctypes.c_void_p, pointer : ctypes.c_void_p) -> None: - frame_buffer = ctypes.string_at(data, size) - frame_timestamp = ctypes.cast(info, ctypes.POINTER(ctypes.c_uint32)).contents.value - frame_handler(frame_buffer, frame_timestamp) + frame_handler(ctypes.string_at(data, size), ctypes.cast(info, ctypes.POINTER(ctypes.c_uint32)).contents.value) def dispatch_event(event : threading.Event, track : int, pointer : ctypes.c_void_p) -> None: diff --git a/facefusion/apis/stream_manager.py b/facefusion/apis/stream_manager.py index 62268ee7..d83f6319 100644 --- a/facefusion/apis/stream_manager.py +++ b/facefusion/apis/stream_manager.py @@ -14,29 +14,28 @@ from facefusion.apis.stream_audio import receive_audio_frames, run_audio_encode_ from facefusion.apis.stream_video import receive_video_frames, run_video_encode_loop from facefusion.libraries import datachannel as datachannel_module from facefusion.types import AudioCodec, AudioFrame, BufferPack, PeerConnection, RtcPeer, RtcPeerAudio, SdpAnswer, SdpOffer, SessionId, Time, VideoCodec, VisionFrame -from facefusion.vision import read_static_images +from facefusion.vision import is_vision_frame, read_static_images, to_buffer async def process_image(websocket : WebSocket) -> None: capture_vision_frame = await anext(receive_vision_frames(websocket), None) - if numpy.any(capture_vision_frame): + if is_vision_frame(capture_vision_frame): source_vision_frames = read_static_images(state_manager.get_item('source_paths')) output_vision_frame = streamer.process_stream_frame(source_vision_frames, capture_vision_frame) - is_success, output_frame_buffer = cv2.imencode('.jpg', output_vision_frame) + output_vision_buffer = to_buffer(output_vision_frame) - if is_success: - await websocket.send_bytes(output_frame_buffer.tobytes()) + await websocket.send_bytes(output_vision_buffer) async def receive_vision_frames(websocket : WebSocket) -> AsyncIterator[VisionFrame]: websocket_event = await websocket.receive() while websocket_event.get('type') == 'websocket.receive': - frame_buffer = websocket_event.get('bytes') or bytes() - vision_frame = cv2.imdecode(numpy.frombuffer(frame_buffer, numpy.uint8), cv2.IMREAD_COLOR) + vision_buffer = websocket_event.get('bytes') or bytes() + vision_frame = cv2.imdecode(numpy.frombuffer(vision_buffer, numpy.uint8), cv2.IMREAD_COLOR) - if numpy.any(vision_frame): + if is_vision_frame(vision_frame): yield vision_frame websocket_event = await websocket.receive() diff --git a/facefusion/apis/stream_video.py b/facefusion/apis/stream_video.py index b593786b..16ab4a30 100644 --- a/facefusion/apis/stream_video.py +++ b/facefusion/apis/stream_video.py @@ -10,7 +10,7 @@ from facefusion import rtc, state_manager, streamer from facefusion.apis.stream_event import create_receive_event from facefusion.codecs import aom_decoder, aom_encoder, vpx_decoder, vpx_encoder from facefusion.types import AomDecoder, AomEncoder, BitRate, Buffer, BufferPack, Resolution, RtcPeer, RtcPeerVideo, Time, VideoCodec, VisionFrame, VpxDecoder, VpxEncoder -from facefusion.vision import read_static_images +from facefusion.vision import is_vision_frame, read_static_images def run_video_encode_loop(rtc_peer : RtcPeer, video_queue : Queue[Tuple[Time, Future[BufferPack]]]) -> None: @@ -170,7 +170,7 @@ def update_video_encoder_bitrate(video_codec : VideoCodec, video_encoder : VpxEn def handle_video_frame(source_vision_frames : List[VisionFrame], video_codec : VideoCodec, video_decoder : VpxDecoder | AomDecoder, video_queue : Queue[Tuple[Time, Future[BufferPack]]], video_executor : ThreadPoolExecutor, video_buffer : Buffer, video_timestamp : int) -> None: vision_frame = decode_video_frame(video_codec, video_decoder, video_buffer) - if numpy.any(vision_frame) and video_queue.qsize() < video_queue.maxsize: + if is_vision_frame(vision_frame) and video_queue.qsize() < video_queue.maxsize: video_future = video_executor.submit(process_video_frame, source_vision_frames, vision_frame) video_time = rtc.convert_timestamp_to_time(video_codec, video_timestamp) video_queue.put((video_time, video_future)) diff --git a/facefusion/types.py b/facefusion/types.py index d9109b7a..dc973005 100755 --- a/facefusion/types.py +++ b/facefusion/types.py @@ -314,6 +314,9 @@ AssetMetadata : TypeAlias = AudioMetadata | ImageMetadata | VideoMetadata AssetSet : TypeAlias = Dict[AssetId, AudioAsset | ImageAsset | VideoAsset] AssetStore : TypeAlias = Dict[SessionId, AssetSet] +AssetAction = Literal['capture'] +AssetSubject = Literal['frame'] + BenchmarkMode = Literal['warm', 'cold'] BenchmarkResolution = Literal['240p', '360p', '540p', '720p', '1080p', '1440p', '2160p'] BenchmarkSet : TypeAlias = Dict[BenchmarkResolution, str] diff --git a/facefusion/vision.py b/facefusion/vision.py index 24cce8a6..24ebb598 100644 --- a/facefusion/vision.py +++ b/facefusion/vision.py @@ -11,7 +11,7 @@ from facefusion.common_helper import is_windows from facefusion.filesystem import get_file_extension, is_image, is_video from facefusion.media_helper import restrict_trim_frame from facefusion.thread_helper import thread_lock, thread_semaphore -from facefusion.types import ColorMode, Duration, Fps, Mask, Orientation, Resolution, Scale, VisionFrame +from facefusion.types import Buffer, ColorMode, Duration, Fps, Mask, Orientation, Resolution, Scale, VisionFrame def read_static_images(image_paths : List[str], color_mode : ColorMode = 'rgb') -> List[VisionFrame]: @@ -202,6 +202,18 @@ def detect_frame_orientation(vision_frame : VisionFrame) -> Orientation: return 'portrait' +def is_vision_frame(vision_frame : VisionFrame) -> bool: + return numpy.ndim(vision_frame) == 3 + + +def is_vision_frames(vision_frames : List[VisionFrame]) -> bool: + for vision_frame in vision_frames: + if not is_vision_frame(vision_frame): + return False + + return True + + def restrict_frame(vision_frame : VisionFrame, resolution : Resolution) -> VisionFrame: height, width = vision_frame.shape[:2] restrict_width, restrict_height = resolution @@ -264,6 +276,7 @@ def match_frame_color(source_vision_frame : VisionFrame, target_vision_frame : V for color_difference_size in color_difference_sizes: source_vision_frame = equalize_frame_color(source_vision_frame, target_vision_frame, normalize_resolution((color_difference_size, color_difference_size))) + target_vision_frame = equalize_frame_color(source_vision_frame, target_vision_frame, target_vision_frame.shape[:2][::-1]) return target_vision_frame @@ -293,8 +306,16 @@ def create_empty_vision_frame() -> VisionFrame: return numpy.zeros((1, 1, 3)).astype(numpy.uint8) -def is_vision_frame(vision_frame : VisionFrame) -> bool: - return numpy.ndim(vision_frame) == 3 +def to_buffer(vision_frame : VisionFrame) -> Buffer: + is_success, vision_buffer = cv2.imencode('.jpg', vision_frame) + + if is_success: + return vision_buffer.tobytes() + return bytes() + + +def to_strip_buffer(vision_frames : List[VisionFrame]) -> Buffer: + return to_buffer(cv2.hconcat(vision_frames)) def create_tile_frames(vision_frame : VisionFrame, size : Size) -> Tuple[List[VisionFrame], int, int]: diff --git a/tests/test_api_stream_video.py b/tests/test_api_stream_video.py index 50b49110..58169fe9 100644 --- a/tests/test_api_stream_video.py +++ b/tests/test_api_stream_video.py @@ -8,7 +8,6 @@ from typing import Tuple from unittest.mock import patch import cv2 -import numpy import pytest from facefusion import rtc, rtc_store, state_manager @@ -19,7 +18,7 @@ from facefusion.download import conditional_download from facefusion.hash_helper import create_hash from facefusion.libraries import aom as aom_module, datachannel as datachannel_module, vpx as vpx_module from facefusion.types import Buffer, BufferPack, FrameHandler, RtcPeer, RtcPeerVideo, Time, VideoCodec -from facefusion.vision import read_video_frame +from facefusion.vision import is_vision_frame, read_video_frame from .assert_helper import get_test_example_file, get_test_examples_directory @@ -178,7 +177,7 @@ def test_create_and_destroy_video_decoder(video_codec : VideoCodec) -> None: video_decoder = create_video_decoder(video_codec) - assert numpy.any(decode_video_frame(video_codec, video_decoder, encode_buffer)) + assert is_vision_frame(decode_video_frame(video_codec, video_decoder, encode_buffer)) destroy_video_decoder(video_codec, video_decoder)