m-rishabh/iris-api-string-labels
0
1from fastapi import FastAPI, HTTPException, Header2from pydantic import BaseModel3from typing import List, Dict, Any, Optional4import base645 6app = FastAPI()7 8 9from typing import List10import numpy as np11import joblib12import pandas as pd13import os14import traceback15 16# Load the real model17MODEL_PATH = os.path.join(os.path.dirname(__file__), "iris_knn_pipeline.pkl")18model = None19load_error = None20 21try:22 if os.path.exists(MODEL_PATH):23 model = joblib.load(MODEL_PATH)24 else:25 load_error = f"Model file not found at {MODEL_PATH}"26except Exception as e:27 load_error = f"Error loading model: {str(e)}\n{traceback.format_exc()}"28 29def predict_iris(features: List[float]) -> tuple:30 if load_error:31 raise Exception(f"Model not loaded: {load_error}")32 33 # Feature names expected by the scikit-learn pipeline34 FEATURE_NAMES = [35 "sepal length (cm)",36 "sepal width (cm)",37 "petal length (cm)",38 "petal width (cm)"39 ]40 41 # Convert to DataFrame with correct column names42 df = pd.DataFrame([features], columns=FEATURE_NAMES)43 44 pred = int(model.predict(df)[0])45 probs = model.predict_proba(df)[0].tolist()46 47 return pred, probs48 49CLASS_NAMES = ["setosa", "versicolor", "virginica"]50 51 52class ArrayRequest(BaseModel):53 features: List[float]54 55class ObjectRequest(BaseModel):56 sepal_length: float57 sepal_width: float58 petal_length: float59 petal_width: float60 61@app.get("/")62def root():63 return {"message": "Iris API - Format 6 - String Labels"}64 65 66@app.post("/predict")67async def predict(request: ArrayRequest):68 pred, probs = predict_iris(request.features)69 return {70 "predicted_class": CLASS_NAMES[pred],71 "probabilities": probs72 }73 74 75if __name__ == "__main__":76 import uvicorn77 uvicorn.run(app, host="0.0.0.0", port=7860)