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:
Henry Ruhs
2026-08-07 15:23:09 +02:00
committed by GitHub
parent a66609b388
commit 1ca152b934
9 changed files with 86 additions and 24 deletions
+33 -2
View File
@@ -8,8 +8,8 @@ from starlette.datastructures import UploadFile
import facefusion.choices import facefusion.choices
from facefusion import ffmpeg, process_manager, state_manager 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.filesystem import create_directory, get_file_extension, get_file_format, is_audio, is_image, is_video
from facefusion.types import ImageMetadata, MediaType from facefusion.types import ImageAsset, ImageMetadata, MediaType, VideoAsset, VisionFrame
from facefusion.vision import detect_image_resolution 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: 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() process_manager.end()
return asset_paths 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
+13 -2
View File
@@ -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 import session_context, session_manager
from facefusion.apis import asset_store from facefusion.apis import asset_store
from facefusion.apis.asset_helper import save_asset_files, validate_asset_files from facefusion.apis.asset_helper import capture_asset_frames, save_asset_files, validate_asset_files
from facefusion.apis.endpoints.session import extract_access_token from facefusion.apis.session_helper import extract_access_token
from facefusion.filesystem import remove_file from facefusion.filesystem import remove_file
from facefusion.vision import is_vision_frames, to_strip_buffer
async def upload_asset(request : Request) -> Response: 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) asset = asset_store.get_asset(session_id, asset_id)
if asset: 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': if request.query_params.get('action') == 'download':
asset_path = asset.get('path') asset_path = asset.get('path')
+1 -1
View File
@@ -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 import args_helper, capability_store, session_manager, state_manager, translator
from facefusion.apis import asset_store 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: async def get_state(request : Request) -> JSONResponse:
+1 -3
View File
@@ -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: 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_handler(ctypes.string_at(data, size), ctypes.cast(info, ctypes.POINTER(ctypes.c_uint32)).contents.value)
frame_timestamp = ctypes.cast(info, ctypes.POINTER(ctypes.c_uint32)).contents.value
frame_handler(frame_buffer, frame_timestamp)
def dispatch_event(event : threading.Event, track : int, pointer : ctypes.c_void_p) -> None: def dispatch_event(event : threading.Event, track : int, pointer : ctypes.c_void_p) -> None:
+7 -8
View File
@@ -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.apis.stream_video import receive_video_frames, run_video_encode_loop
from facefusion.libraries import datachannel as datachannel_module 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.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: async def process_image(websocket : WebSocket) -> None:
capture_vision_frame = await anext(receive_vision_frames(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')) 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) 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_vision_buffer)
await websocket.send_bytes(output_frame_buffer.tobytes())
async def receive_vision_frames(websocket : WebSocket) -> AsyncIterator[VisionFrame]: async def receive_vision_frames(websocket : WebSocket) -> AsyncIterator[VisionFrame]:
websocket_event = await websocket.receive() websocket_event = await websocket.receive()
while websocket_event.get('type') == 'websocket.receive': while websocket_event.get('type') == 'websocket.receive':
frame_buffer = websocket_event.get('bytes') or bytes() vision_buffer = websocket_event.get('bytes') or bytes()
vision_frame = cv2.imdecode(numpy.frombuffer(frame_buffer, numpy.uint8), cv2.IMREAD_COLOR) 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 yield vision_frame
websocket_event = await websocket.receive() websocket_event = await websocket.receive()
+2 -2
View File
@@ -10,7 +10,7 @@ from facefusion import rtc, state_manager, streamer
from facefusion.apis.stream_event import create_receive_event from facefusion.apis.stream_event import create_receive_event
from facefusion.codecs import aom_decoder, aom_encoder, vpx_decoder, vpx_encoder 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.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: 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: 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) 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_future = video_executor.submit(process_video_frame, source_vision_frames, vision_frame)
video_time = rtc.convert_timestamp_to_time(video_codec, video_timestamp) video_time = rtc.convert_timestamp_to_time(video_codec, video_timestamp)
video_queue.put((video_time, video_future)) video_queue.put((video_time, video_future))
+3
View File
@@ -314,6 +314,9 @@ AssetMetadata : TypeAlias = AudioMetadata | ImageMetadata | VideoMetadata
AssetSet : TypeAlias = Dict[AssetId, AudioAsset | ImageAsset | VideoAsset] AssetSet : TypeAlias = Dict[AssetId, AudioAsset | ImageAsset | VideoAsset]
AssetStore : TypeAlias = Dict[SessionId, AssetSet] AssetStore : TypeAlias = Dict[SessionId, AssetSet]
AssetAction = Literal['capture']
AssetSubject = Literal['frame']
BenchmarkMode = Literal['warm', 'cold'] BenchmarkMode = Literal['warm', 'cold']
BenchmarkResolution = Literal['240p', '360p', '540p', '720p', '1080p', '1440p', '2160p'] BenchmarkResolution = Literal['240p', '360p', '540p', '720p', '1080p', '1440p', '2160p']
BenchmarkSet : TypeAlias = Dict[BenchmarkResolution, str] BenchmarkSet : TypeAlias = Dict[BenchmarkResolution, str]
+24 -3
View File
@@ -11,7 +11,7 @@ from facefusion.common_helper import is_windows
from facefusion.filesystem import get_file_extension, is_image, is_video from facefusion.filesystem import get_file_extension, is_image, is_video
from facefusion.media_helper import restrict_trim_frame from facefusion.media_helper import restrict_trim_frame
from facefusion.thread_helper import thread_lock, thread_semaphore 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]: 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' 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: def restrict_frame(vision_frame : VisionFrame, resolution : Resolution) -> VisionFrame:
height, width = vision_frame.shape[:2] height, width = vision_frame.shape[:2]
restrict_width, restrict_height = resolution 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: 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))) 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]) target_vision_frame = equalize_frame_color(source_vision_frame, target_vision_frame, target_vision_frame.shape[:2][::-1])
return target_vision_frame return target_vision_frame
@@ -293,8 +306,16 @@ def create_empty_vision_frame() -> VisionFrame:
return numpy.zeros((1, 1, 3)).astype(numpy.uint8) return numpy.zeros((1, 1, 3)).astype(numpy.uint8)
def is_vision_frame(vision_frame : VisionFrame) -> bool: def to_buffer(vision_frame : VisionFrame) -> Buffer:
return numpy.ndim(vision_frame) == 3 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]: def create_tile_frames(vision_frame : VisionFrame, size : Size) -> Tuple[List[VisionFrame], int, int]:
+2 -3
View File
@@ -8,7 +8,6 @@ from typing import Tuple
from unittest.mock import patch from unittest.mock import patch
import cv2 import cv2
import numpy
import pytest import pytest
from facefusion import rtc, rtc_store, state_manager 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.hash_helper import create_hash
from facefusion.libraries import aom as aom_module, datachannel as datachannel_module, vpx as vpx_module 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.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 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) 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) destroy_video_decoder(video_codec, video_decoder)