import spaces # must be first — ZeroGPU requires this before any CUDA package import os import gc import time import traceback import json import gradio as gr import insightface import torch from insightface.app import FaceAnalysis from image_utils import _b64_to_pil, _collect_source_images from validators import ( _validate_target, _validate_selection, _validate_sources, _validate_bboxes, ) from logging_utils import spawn_log from pipeline import ( detect as _pipeline_detect, detect_and_crop_source as _pipeline_detect_and_crop_source, ) from pipeline_flux import ( prepare_swap_crops, run_swap_flux, load_flux_pipeline_cpu, move_flux_to_device, ) MODEL_DIR = "models" os.makedirs(MODEL_DIR, exist_ok=True) import torchvision print( f"[startup] torch={torch.__version__}, CUDA={torch.cuda.is_available()}", flush=True ) print(f"[startup] torchvision={torchvision.__version__}", flush=True) print(f"[startup] insightface={insightface.__version__}", flush=True) # ── CPU face analyzer (for detection step, no GPU required) ─────────────────── print("[startup] loading CPU FaceAnalysis...", flush=True) _cpu_face_analyzer = FaceAnalysis(name="buffalo_l", providers=["CPUExecutionProvider"]) _cpu_face_analyzer.prepare(ctx_id=0, det_size=(640, 640)) print("[startup] CPU FaceAnalysis ready.", flush=True) # ── Flux2 pipeline: weights downloaded on CPU at startup (no GPU cost), # then moved to CUDA inside @spaces.GPU on the first swap request. ────────── print( "[startup] loading Flux2 pipeline on CPU (download may take a while)…", flush=True ) _flux_pipe_cpu = load_flux_pipeline_cpu() print("[startup] Flux2 CPU pipeline ready.", flush=True) # ── Public Gradio functions ──────────────────────────────────────────────────── def detect_faces(target_b64): """Detect faces in target image (CPU). Returns JSON.""" return _pipeline_detect(target_b64, _cpu_face_analyzer) def process_source(source_json): """Detect and crop source face (CPU). Returns JSON.""" return _pipeline_detect_and_crop_source(source_json, _cpu_face_analyzer) @spaces.GPU(duration=lambda _img, crops: 50 + len(crops) * 20) def _swap_faces_gpu(target_image, prepared_crops): """GPU inference only. Accepts pre-cropped inputs; yields progress then (result_image, "").""" t_gpu = time.perf_counter() print("[flux] ===== START =====", flush=True) yield None, "Moving model to GPU…" flux_pipe, flux_device = move_flux_to_device(_flux_pipe_cpu) print("[flux] pipeline on GPU ready", flush=True) yield None, "Swapping faces…" try: result_image = run_swap_flux( target_image, prepared_crops, flux_pipe, flux_device ) print(f"[flux] done in {time.perf_counter() - t_gpu:.2f}s", flush=True) yield result_image, "" except gr.Error: raise except Exception as e: print(f"[flux] ERROR: {e}\n{traceback.format_exc()}", flush=True) raise finally: gc.collect() torch.cuda.empty_cache() print("[flux] ===== END =====", flush=True) def swap_faces( target_b64, selected_json, sources_json, bboxes_json, original_sources_json ): """Validate inputs, run GPU swap, then log after the GPU scope is released.""" _validate_target(target_b64) selected_indices = json.loads(selected_json) if selected_json else [] _validate_selection(selected_indices) sources = json.loads(sources_json) if sources_json else {} _validate_sources(selected_indices, sources) bboxes = json.loads(bboxes_json) if bboxes_json else {} _validate_bboxes(selected_indices, bboxes) target_image = _b64_to_pil(target_b64) source_images = _collect_source_images(selected_indices, sources) try: orig_sources = ( json.loads(original_sources_json) if original_sources_json else {} ) log_source_images = ( _collect_source_images(selected_indices, orig_sources) if orig_sources else source_images ) except Exception: log_source_images = source_images prepared_crops = prepare_swap_crops(target_image, selected_indices, sources, bboxes) last_yield = None success = False error_msg = "" duration = 0.0 t0 = time.perf_counter() try: for image, status in _swap_faces_gpu(target_image, prepared_crops): last_yield = (image, status) yield (gr.update() if image is None else image), status duration = time.perf_counter() - t0 success = True except gr.Error: raise except Exception as e: duration = time.perf_counter() - t0 error_msg = str(e) raise finally: if success or error_msg: result_image = last_yield[0] if last_yield is not None else None spawn_log( target_image, log_source_images, result_image, selected_indices, duration, success, error_msg, ) # ── Static asset loading ─────────────────────────────────────────────────────── def _read(path): with open(path) as f: return f.read() css = _read("static/app.css") wire_outputs_js = _read("static/wire_outputs.js") render_faces_js = _read("static/render_faces.js") handle_source_result_js = _read("static/handle_source_result.js") app_html = _read("templates/app.html") # ── Gradio app ───────────────────────────────────────────────────────────────── with gr.Blocks() as demo: # Hidden data channels target_b64_inp = gr.Textbox( value="", elem_id="target-b64", elem_classes="hidden-input", container=False ) detected_json_out = gr.Textbox( value="", elem_id="detected-json", elem_classes="hidden-input", container=False ) selected_json_inp = gr.Textbox( value="", elem_id="selected-json", elem_classes="hidden-input", container=False ) sources_json_inp = gr.Textbox( value="", elem_id="sources-json", elem_classes="hidden-input", container=False ) bboxes_json_inp = gr.Textbox( value="", elem_id="bboxes-json", elem_classes="hidden-input", container=False ) swap_status_out = gr.Textbox( value="", elem_id="swap-status", elem_classes="hidden-input", container=False ) result_img = gr.Image( elem_id="result-img", elem_classes="hidden-input", container=False, format="png" ) source_inp = gr.Textbox( value="", elem_id="source-inp", elem_classes="hidden-input", container=False ) source_result_out = gr.Textbox( value="", elem_id="source-result", elem_classes="hidden-input", container=False ) original_sources_json_inp = gr.Textbox( value="", elem_id="original-sources-json", elem_classes="hidden-input", container=False, ) # Hidden trigger buttons detect_btn = gr.Button("Detect", elem_id="detect-btn", elem_classes="hidden-input") run_btn = gr.Button("Run", elem_id="run-btn", elem_classes="hidden-input") process_source_btn = gr.Button( "Process Source", elem_id="process-source-btn", elem_classes="hidden-input" ) gr.HTML(app_html) detect_btn.click( fn=detect_faces, inputs=[target_b64_inp], outputs=[detected_json_out], ) detected_json_out.change( fn=None, inputs=[detected_json_out], js=render_faces_js, ) run_btn.click( fn=swap_faces, inputs=[ target_b64_inp, selected_json_inp, sources_json_inp, bboxes_json_inp, original_sources_json_inp, ], outputs=[result_img, swap_status_out], ) process_source_btn.click( fn=process_source, inputs=[source_inp], outputs=[source_result_out], ) source_result_out.change( fn=None, inputs=[source_result_out], js=handle_source_result_js, ) demo.load(fn=None, js=wire_outputs_js) if __name__ == "__main__": demo.queue(max_size=10).launch( css=css, show_error=True, js="() => { document.documentElement.classList.add('dark'); }", )