# api/app.py import io import json import traceback from pathlib import Path from fastapi import FastAPI, File, UploadFile from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import FileResponse, JSONResponse, Response from PIL import Image import numpy as np import sys # make sure project root is importable (scripts/inference.py lives at repo root) ROOT = Path(__file__).resolve().parent.parent sys.path.append(str(ROOT)) # ---- YOUR ORIGINAL MODEL LOADER (unchanged in spirit) ---- from scripts.inference import ModelInference # create model ONCE (like your original code) # if your ModelInference supports map_location/device, set to CPU for Spaces model = ModelInference( checkpoint_path=str(ROOT / "checkpoints" / "best_checkpoint.pth"), multi_task=True ) # -------- FastAPI app & CORS ---------- app = FastAPI(title="VizRef API") app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) # -------- Serve your UI ---------- @app.get("/") def root(): return FileResponse(str(ROOT / "ui" / "index.html")) # -------- API: health ---------- @app.get("/api/health") def health_check(): return { "status": "healthy", "model_loaded": True, "model_info": {"model_name": "Efficientnet - B0", "multi_task": True}, } # -------- API: training data ---------- @app.get("/api/training_data") def get_training_data(): with open(ROOT / "data" / "splits" / "train.json", "r") as f: data = json.load(f) return {"status": "success", "data": data} # -------- API: proxy image ---------- @app.get("/api/proxy_image") def proxy_image(url: str): import requests try: if url.startswith("http://"): https_url = url.replace("http://", "https://") try: r = requests.get(https_url, timeout=5) if r.status_code == 200: return Response(content=r.content, media_type="image/jpeg") except: pass r = requests.get(url, timeout=5) return Response(content=r.content, media_type="image/jpeg") except: return Response(status_code=404) # --- helper to make numpy types JSON-safe --- def _pythonify(obj): if isinstance(obj, dict): return {k: _pythonify(v) for k, v in obj.items()} if isinstance(obj, list): return [_pythonify(v) for v in obj] if isinstance(obj, (np.integer,)): return int(obj) if isinstance(obj, (np.floating,)): return float(obj) if isinstance(obj, (np.bool_)): return bool(obj) return obj # -------- API: predict (your original logic, adapted to JSON) ---------- @app.post("/api/predict") async def predict(file: UploadFile = File(...)): try: raw = await file.read() try: img = Image.open(io.BytesIO(raw)).convert("RGB") except Exception as e: return JSONResponse({"status": "error", "message": f"Cannot read image: {e}"}, status_code=400) tmp = Path("/tmp/vizref_upload.jpg") img.save(tmp) # your original call result = model.predict_image(str(tmp)) return JSONResponse({ "status": "success", "predictions": _pythonify(result), "model_info": {"model_name": "Efficientnet - B0", "multi_task": True}, }) except Exception as e: print("PREDICT ERROR\n", traceback.format_exc()) return JSONResponse({"status": "error", "message": str(e)}, status_code=500) finally: try: if 'tmp' in locals() and tmp.exists(): tmp.unlink() except: pass