mirror of
https://github.com/facefusion/facefusion.git
synced 2026-08-07 17:48:39 +02:00
Asset Tools Part1 (#1217)
* fix import * introduce collect vision frame from asset, add is_vision_frame everywhere, rename to vision_buffer * fix lint * fix lint * fix lint
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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')
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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]
|
||||
|
||||
+24
-3
@@ -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]:
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user