| |
| 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 |
|
|
| |
| ROOT = Path(__file__).resolve().parent.parent |
| sys.path.append(str(ROOT)) |
|
|
| |
| from scripts.inference import ModelInference |
|
|
| |
| |
| model = ModelInference( |
| checkpoint_path=str(ROOT / "checkpoints" / "best_checkpoint.pth"), |
| multi_task=True |
| ) |
|
|
| |
| app = FastAPI(title="VizRef API") |
| app.add_middleware( |
| CORSMiddleware, |
| allow_origins=["*"], allow_credentials=True, |
| allow_methods=["*"], allow_headers=["*"], |
| ) |
|
|
| |
| @app.get("/") |
| def root(): |
| return FileResponse(str(ROOT / "ui" / "index.html")) |
|
|
| |
| @app.get("/api/health") |
| def health_check(): |
| return { |
| "status": "healthy", |
| "model_loaded": True, |
| "model_info": {"model_name": "Efficientnet - B0", "multi_task": True}, |
| } |
|
|
| |
| @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} |
|
|
| |
| @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) |
|
|
| |
| 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 |
|
|
| |
| @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) |
|
|
| |
| 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 |
|
|