someone-in-the-world's picture
Claude Sonnet 4.6
Log full source image instead of cropped face
d0cfa98
Raw History Blame Contribute Delete
8.72 kB
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'); }",
)