AAlexP's picture
Update app.py
6338e36 verified
Raw
History Blame Contribute Delete
3.73 kB
import os, io, torch
from typing import Optional
from fastapi import FastAPI, HTTPException, Request, Header
from fastapi.middleware.cors import CORSMiddleware
from PIL import Image
# Put all HF caches in a writable place
os.environ["HF_HOME"] = "/tmp" # preferred going forward
os.environ["HUGGINGFACE_HUB_CACHE"] = "/tmp/hub" # optional, explicit
os.environ["TRANSFORMERS_CACHE"] = "/tmp/transformers" # backward-compat
app = FastAPI(title="Skin Cancer Demo Inference")
# Model & secret
MODEL_ID = "Anwarkh1/Skin_Cancer-Image_Classification"
SECRET = os.getenv("SECRET", "") # set this in Space Settings → Variables & secrets
# Lazy-loaded at startup
processor = None
model = None
id2label = None
startup_error: Optional[str] = None
# CORS (handy if you later hit from a web app)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"], allow_methods=["*"], allow_headers=["*"]
)
@app.on_event("startup")
async def load_model():
global processor, model, id2label, startup_error
try:
from transformers import AutoImageProcessor, AutoModelForImageClassification
processor = AutoImageProcessor.from_pretrained(MODEL_ID)
model = AutoModelForImageClassification.from_pretrained(MODEL_ID)
id2label = model.config.id2label
startup_error = None
print("[startup] model loaded")
except Exception as e:
startup_error = f"{type(e).__name__}: {e}"
print("[startup] failed:", startup_error)
@app.get("/")
def root():
return {"service": "skin-cancer-demo", "endpoints": ["/health", "/predict"]}
@app.get("/health")
def health():
return {"status": "ok" if not startup_error else "degraded", "error": startup_error}
def _check_ready():
if startup_error or processor is None or model is None:
raise HTTPException(status_code=503, detail=f"model not ready: {startup_error}")
@app.post("/predict")
async def predict(
request: Request,
token: str = "", # query token for quick tests
x_api_key: str = Header(default="") # preferred: header auth (X-API-Key)
):
# auth
auth = x_api_key or token
if SECRET and auth != SECRET:
raise HTTPException(status_code=401, detail="unauthorized")
_check_ready()
# content-type & size guards (good for Salesforce callouts)
ctype = request.headers.get("content-type", "")
if "application/octet-stream" not in ctype:
raise HTTPException(status_code=415, detail="use application/octet-stream")
img_bytes = await request.body()
if len(img_bytes) == 0:
raise HTTPException(status_code=400, detail="empty body")
if len(img_bytes) > 5 * 1024 * 1024:
raise HTTPException(status_code=413, detail="image too large (>5MB)")
try:
img = Image.open(io.BytesIO(img_bytes)).convert("RGB")
except Exception as e:
raise HTTPException(status_code=400, detail=f"invalid image: {e}")
inputs = processor(images=img, return_tensors="pt")
with torch.no_grad():
logits = model(**inputs).logits
probs_t = torch.softmax(logits, dim=1)[0]
top_idx = int(torch.argmax(probs_t).item())
probs = probs_t.tolist()
def idx_to_label(i: int):
return id2label.get(str(i), id2label.get(i, str(i)))
return {
"prediction": {
"label": idx_to_label(top_idx),
"confidence": float(probs[top_idx])
},
"all_probs": {idx_to_label(i): float(probs[i]) for i in range(len(probs))},
"meta": {
"model": MODEL_ID
}
}
if __name__ == "__main__":
import uvicorn
port = int(os.getenv("PORT", "7860"))
uvicorn.run("app:app", host="0.0.0.0", port=port)