Spaces:
Sleeping
Sleeping
File size: 3,725 Bytes
f2fa3e5 6338e36 1017190 78eb655 f2fa3e5 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 | 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) |