RefDiffNet / app.py
vinayedula's picture
Update app.py
40d93b0 verified
Raw
History Blame Contribute Delete
5.06 kB
"""Gradio demo: PCB defect + golden reference -> enriched image (A11_CA prebackbone)."""
from __future__ import annotations
import os
import socket
from pathlib import Path
# HF Spaces: SSR needs Node and often serves unstyled HTML when it fails.
os.environ.setdefault("GRADIO_SSR_MODE", "false")
_ROOT = Path(__file__).resolve().parent
def _is_hf_space() -> bool:
return bool(os.environ.get("SPACE_ID") or os.environ.get("SYSTEM") == "spaces")
def _configure_runtime() -> None:
os.environ.setdefault(
"PREBACKBONE_ONLY_WEIGHTS",
str(_ROOT / "weights" / "prebackbone_a11_ca.pt"),
)
if _is_hf_space():
os.environ.setdefault("PRELOAD_MODEL", "1")
_configure_runtime()
import gradio_patch
gradio_patch.apply()
import gradio as gr
import numpy as np
from PIL import Image
from prebackbone_infer import enrich_pair, get_enricher
TITLE = "RefDiffNet — PCB Reference–Defect Enrichment"
DESCRIPTION = """
Upload a **defect PCB image** and its **golden reference** (same H×W, pre-aligned — e.g. `*_input.jpg` / `*_reference.jpg` from training).
The RefDiffNet prebackbone outputs an **enriched** image: `enriched = defect + α · gate · delta`
"""
def run_demo(defect_path: str | None, reference_path: str | None):
if not defect_path or not reference_path:
raise gr.Error("Please upload both defect and reference images.")
defect_rgb, ref_rgb, enriched_rgb = enrich_pair(defect_path, reference_path)
defect_rgb = np.array(defect_rgb, dtype=np.uint8, copy=True)
ref_rgb = np.array(ref_rgb, dtype=np.uint8, copy=True)
enriched_rgb = np.array(enriched_rgb, dtype=np.uint8, copy=True)
h = max(defect_rgb.shape[0], ref_rgb.shape[0], enriched_rgb.shape[0])
def _pad(im: np.ndarray) -> np.ndarray:
if im.shape[0] == h:
return im
pad = h - im.shape[0]
return np.pad(im, ((0, pad), (0, 0), (0, 0)), mode="constant", constant_values=114)
row = np.concatenate([_pad(defect_rgb), _pad(ref_rgb), _pad(enriched_rgb)], axis=1)
return Image.fromarray(enriched_rgb), Image.fromarray(row)
def _example_pairs() -> list[list[str]]:
"""All *_input / *_reference pairs bundled in examples/ (from yolo26 prebackbone_samples)."""
pairs: list[list[str]] = []
ex_dir = _ROOT / "examples"
if not ex_dir.is_dir():
return pairs
for inp in sorted(ex_dir.glob("*_input.jpg")):
ref = inp.with_name(inp.name.replace("_input.", "_reference."))
if ref.is_file():
pairs.append([str(inp), str(ref)])
return pairs
def _preload_model() -> None:
try:
get_enricher()
print("[RefDiffNet] Prebackbone loaded.")
except FileNotFoundError as e:
print(f"[RefDiffNet] WARN: {e}")
def build_demo() -> gr.Interface:
return gr.Interface(
fn=run_demo,
inputs=[
gr.Image(type="filepath", label="Defect image (input)", height=320),
gr.Image(type="filepath", label="Golden reference", height=320),
],
outputs=[
gr.Image(type="pil", label="Enriched output", height=360),
gr.Image(type="pil", label="Defect | Reference | Enriched", height=360),
],
title=TITLE,
description=DESCRIPTION.strip(),
examples=_example_pairs() or None,
examples_per_page=6,
cache_examples=False,
flagging_mode="never",
theme=gr.themes.Soft(),
)
demo = build_demo()
def _local_launch_kwargs() -> dict:
server_name = os.environ.get("GRADIO_SERVER_NAME", "127.0.0.1")
preferred_port = int(os.environ.get("PORT", os.environ.get("GRADIO_SERVER_PORT", "7860")))
bind_host = "127.0.0.1" if server_name in ("127.0.0.1", "localhost") else server_name
server_port = preferred_port
if server_name in ("127.0.0.1", "localhost"):
for port in range(preferred_port, preferred_port + 20):
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
try:
sock.bind((bind_host, port))
server_port = port
break
except OSError:
continue
share = os.environ.get("GRADIO_SHARE", "").lower() in ("1", "true", "yes")
return {
"server_name": server_name,
"server_port": server_port,
"share": share,
}
def _launch() -> None:
"""Launch with SSR off (required for themed UI on HF Spaces)."""
kwargs: dict = {"show_error": True, "ssr_mode": False}
if _is_hf_space():
demo.launch(**kwargs)
else:
demo.launch(
**_local_launch_kwargs(),
inbrowser=False,
prevent_thread_lock=False,
**kwargs,
)
if __name__ == "__main__":
print(f"[RefDiffNet] gradio {getattr(gr, '__version__', 'unknown')} ssr_mode=False")
if os.environ.get("PRELOAD_MODEL", "0").lower() in ("1", "true", "yes"):
_preload_model()
_launch()