Team Ai
Apppublic

DeveloperSA/PlantDiseaseFastAPI

sourceHugging Faceupdated 10mo agoView on Hugging Face
0likes
app.py81 linesDownload Raw Back to root
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)