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)