AAlexP commited on
Commit
f2fa3e5
·
verified ·
1 Parent(s): f2d19d8

Create app.py

Browse files
Files changed (1) hide show
  1. app.py +105 -0
app.py ADDED
@@ -0,0 +1,105 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os, io, torch
2
+ from typing import Optional
3
+ from fastapi import FastAPI, HTTPException, Request, Header
4
+ from fastapi.middleware.cors import CORSMiddleware
5
+ from PIL import Image
6
+
7
+ app = FastAPI(title="Skin Cancer Demo Inference")
8
+
9
+ # Model & secret
10
+ MODEL_ID = "Anwarkh1/Skin_Cancer-Image_Classification"
11
+ SECRET = os.getenv("SECRET", "") # set this in Space Settings → Variables & secrets
12
+
13
+ # Lazy-loaded at startup
14
+ processor = None
15
+ model = None
16
+ id2label = None
17
+ startup_error: Optional[str] = None
18
+
19
+ # CORS (handy if you later hit from a web app)
20
+ app.add_middleware(
21
+ CORSMiddleware,
22
+ allow_origins=["*"], allow_methods=["*"], allow_headers=["*"]
23
+ )
24
+
25
+ @app.on_event("startup")
26
+ async def load_model():
27
+ global processor, model, id2label, startup_error
28
+ try:
29
+ from transformers import AutoImageProcessor, AutoModelForImageClassification
30
+ processor = AutoImageProcessor.from_pretrained(MODEL_ID)
31
+ model = AutoModelForImageClassification.from_pretrained(MODEL_ID)
32
+ id2label = model.config.id2label
33
+ startup_error = None
34
+ print("[startup] model loaded")
35
+ except Exception as e:
36
+ startup_error = f"{type(e).__name__}: {e}"
37
+ print("[startup] failed:", startup_error)
38
+
39
+ @app.get("/")
40
+ def root():
41
+ return {"service": "skin-cancer-demo", "endpoints": ["/health", "/predict"]}
42
+
43
+ @app.get("/health")
44
+ def health():
45
+ return {"status": "ok" if not startup_error else "degraded", "error": startup_error}
46
+
47
+ def _check_ready():
48
+ if startup_error or processor is None or model is None:
49
+ raise HTTPException(status_code=503, detail=f"model not ready: {startup_error}")
50
+
51
+ @app.post("/predict")
52
+ async def predict(
53
+ request: Request,
54
+ token: str = "", # query token for quick tests
55
+ x_api_key: str = Header(default="") # preferred: header auth (X-API-Key)
56
+ ):
57
+ # auth
58
+ auth = x_api_key or token
59
+ if SECRET and auth != SECRET:
60
+ raise HTTPException(status_code=401, detail="unauthorized")
61
+
62
+ _check_ready()
63
+
64
+ # content-type & size guards (good for Salesforce callouts)
65
+ ctype = request.headers.get("content-type", "")
66
+ if "application/octet-stream" not in ctype:
67
+ raise HTTPException(status_code=415, detail="use application/octet-stream")
68
+
69
+ img_bytes = await request.body()
70
+ if len(img_bytes) == 0:
71
+ raise HTTPException(status_code=400, detail="empty body")
72
+ if len(img_bytes) > 5 * 1024 * 1024:
73
+ raise HTTPException(status_code=413, detail="image too large (>5MB)")
74
+
75
+ try:
76
+ img = Image.open(io.BytesIO(img_bytes)).convert("RGB")
77
+ except Exception as e:
78
+ raise HTTPException(status_code=400, detail=f"invalid image: {e}")
79
+
80
+ inputs = processor(images=img, return_tensors="pt")
81
+ with torch.no_grad():
82
+ logits = model(**inputs).logits
83
+ probs_t = torch.softmax(logits, dim=1)[0]
84
+
85
+ top_idx = int(torch.argmax(probs_t).item())
86
+ probs = probs_t.tolist()
87
+
88
+ def idx_to_label(i: int):
89
+ return id2label.get(str(i), id2label.get(i, str(i)))
90
+
91
+ return {
92
+ "prediction": {
93
+ "label": idx_to_label(top_idx),
94
+ "confidence": float(probs[top_idx])
95
+ },
96
+ "all_probs": {idx_to_label(i): float(probs[i]) for i in range(len(probs))},
97
+ "meta": {
98
+ "model": MODEL_ID
99
+ }
100
+ }
101
+
102
+ if __name__ == "__main__":
103
+ import uvicorn
104
+ port = int(os.getenv("PORT", "7860"))
105
+ uvicorn.run("app:app", host="0.0.0.0", port=port)