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
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
+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.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')
+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.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:
+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:
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:
+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.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()
+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.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))
+3
View File
@@ -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
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.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]:
+2 -3
View File
@@ -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)