Spaces:
Sleeping
Sleeping
| 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=["*"] | |
| ) | |
| 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) | |
| def root(): | |
| return {"service": "skin-cancer-demo", "endpoints": ["/health", "/predict"]} | |
| 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}") | |
| 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) |