DeveloperSA/PlantDiseaseFastAPI
0
1from fastapi import FastAPI, UploadFile, File2from contextlib import asynccontextmanager3import uvicorn4import tensorflow as tf5import numpy as np6from PIL import Image7import io8import json9import os10import gdown11 12MODEL_PATH = "plant_disease_prediction_model.h5"13MODEL_URL = "https://drive.google.com/uc?id=1rKh-IElSdHTqax7XdfSdZTn-r8T_qWPf"14IMG_SIZE = (224, 224)15 16# Load disease details (no change needed here)17with open("diseases.json", "r") as f:18 DISEASE_DATA = json.load(f)19 DISEASE_DATA = {int(k): v for k, v in DISEASE_DATA.items()}20 21# Model variable is declared but not assigned here22model = None23 24# Download model if it doesn't exist (no change needed here)25if not os.path.exists(MODEL_PATH):26 gdown.download(MODEL_URL, MODEL_PATH, quiet=False)27 28# 1. CRITICAL OPTIMIZATION: Use lifespan to load model outside global scope29@asynccontextmanager30async def lifespan(app: FastAPI):31 global model32 # Model loading happens when the app starts up33 model = tf.keras.models.load_model(MODEL_PATH, compile=False)34 yield35 # Clean up on shutdown (optional but good practice)36 model = None37 38# Pass the lifespan function to FastAPI39app = FastAPI(lifespan=lifespan)40 41def preprocess(img_bytes):42 img = Image.open(io.BytesIO(img_bytes)).convert("RGB")43 img = img.resize(IMG_SIZE)44 img = np.array(img) / 255.045 img = np.expand_dims(img, axis=0)46 return img47 48@app.post("/predict")49async def predict(file: UploadFile = File(...)):50 img_bytes = await file.read()51 img = preprocess(img_bytes)52 53 # Use the globally available model loaded via lifespan54 preds = model.predict(img)[0]55 class_index = int(np.argmax(preds))56 confidence = float(np.max(preds))57 58 info = DISEASE_DATA[class_index]59 60 return {61 "class_index": class_index,62 "disease_name": info["name"],63 "description": info["description"],64 "cause": info["cause"],65 "solution": info["solution"],66 "prevention": info["prevention"],67 "confidence": confidence68 }69 70# 2. MINOR ENHANCEMENT: Improve home endpoint71@app.get("/")72def home():73 disease_list = [{"id": k, "name": v["name"]} for k, v in DISEASE_DATA.items()]74 return {75 "message": "Plant Disease API is running!",76 "model_loaded_status": "Successfully loaded via lifespan" if model is not None else "Loading...",77 "supported_disease_count": len(disease_list)78 }79 80if __name__ == "__main__":81 uvicorn.run(app, host="0.0.0.0", port=8000)