VizRef / api /app.py
chenx906's picture
Fix the deployment space problem
2ef860a
Raw
History Blame Contribute Delete
3.72 kB
# 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