import os import pickle from sklearn.datasets import load_iris from sklearn.tree import DecisionTreeClassifier MODEL_PATH = "ml_models/iris.model" def train_model() -> DecisionTreeClassifier: iris = load_iris() X = iris.data y = iris.target model = DecisionTreeClassifier() model.fit(X, y) return model def save_model(model: DecisionTreeClassifier) -> bool: os.makedirs(os.path.dirname(MODEL_PATH), exist_ok=True) try: with open(MODEL_PATH, "wb") as f: pickle.dump(model, f) return True except Exception as e: print(f"An error occurred while saving the model: {e}") return False def load_model() -> DecisionTreeClassifier: try: with open(MODEL_PATH, 'rb') as f: model = pickle.load(f) return model except Exception as e: print(f"An error occurred while loading the model: {e}") raise e def predict( model: DecisionTreeClassifier, sepal_length: float, sepal_width: float, petal_length: float, petal_width: float ) -> dict: prediction = model.predict([[sepal_length, sepal_width, petal_length, petal_width]]) prediction_str = "" match prediction: case 0: prediction_str = "setosa" case 1: prediction_str = "versicolor" case 2: prediction_str = "virginica" prediction_prob = model.predict_proba([[sepal_length, sepal_width, petal_length, petal_width]]) return { "prediction": prediction_str, "prediction_probabilities": prediction_prob[0].tolist() }