crtal's picture
docs: fix README model name, add health endpoint, fix local dev port
5aba070
Raw History Blame Contribute Delete
2.78 kB
from fastapi import FastAPI, File, UploadFile, HTTPException
import torch
import torch.nn as nn
from torchvision import models, transforms
from PIL import Image, UnidentifiedImageError
import io
import uvicorn
app = FastAPI(title="Food Spoilage Detection API (Swin)")
# Load model
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# Swin architecture
model = models.swin_t()
num_ftrs = model.head.in_features
model.head = nn.Linear(num_ftrs, 2)
# Load Swin weights
model.load_state_dict(torch.load("models/spoilage_model_swin_finetuned.pth", map_location=device))
model = model.to(device)
model.eval()
# Preprocessing transforms
preprocess = transforms.Compose([
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])
CLASS_NAMES = ["fresh", "spoiled"]
MAX_FILE_SIZE = 20 * 1024 * 1024 # 20 MB
@app.post("/predict")
async def predict(file: UploadFile = File(...)):
contents = await file.read()
if len(contents) > MAX_FILE_SIZE:
raise HTTPException(status_code=400, detail="File too large. Maximum size is 10 MB.")
try:
image = Image.open(io.BytesIO(contents)).convert("RGB")
except UnidentifiedImageError:
raise HTTPException(status_code=400, detail="Could not read image. Please upload a valid image file.")
try:
input_tensor = preprocess(image)
input_batch = input_tensor.unsqueeze(0).to(device)
with torch.inference_mode():
outputs = model(input_batch)
probabilities = torch.nn.functional.softmax(outputs[0], dim=0)
confidence, index = torch.max(probabilities, 0)
conf_score = float(confidence.item())
if conf_score < 0.70:
return {
"status": "ambiguous",
"message": "Model confidence too low. The image might not be a recognized food item or is too low quality.",
"confidence": conf_score,
"suggestion": "Please provide a clearer image of a fruit or vegetable."
}
return {
"prediction": CLASS_NAMES[index.item()],
"confidence": conf_score,
"status": "success",
"model": "swin_transformer_tiny_finetuned"
}
except Exception as e:
raise HTTPException(status_code=500, detail=f"Error processing image: {str(e)}")
@app.get("/")
def read_root():
return {
"message": "Food Spoilage Detection API is running.",
"usage": "POST an image to /predict",
"supported_classes": CLASS_NAMES
}
@app.get("/health")
def health():
return {"status": "ok"}
if __name__ == "__main__":
uvicorn.run(app, host="0.0.0.0", port=7860)