| from __future__ import annotations |
|
|
| from pathlib import Path |
|
|
| import gradio as gr |
| import joblib |
| import pandas as pd |
|
|
|
|
| ARTIFACT_PATH = Path(__file__).with_name("model.joblib") |
|
|
|
|
| def _load_bundle() -> dict: |
| if not ARTIFACT_PATH.exists(): |
| from train import train_and_save |
|
|
| train_and_save(ARTIFACT_PATH) |
|
|
| return joblib.load(ARTIFACT_PATH) |
|
|
|
|
| BUNDLE = _load_bundle() |
| MODEL = BUNDLE["model"] |
| TARGET_NAMES = BUNDLE["target_names"] |
| FEATURE_NAMES = BUNDLE["feature_names"] |
|
|
|
|
| def predict(sepal_length: float, sepal_width: float, petal_length: float, petal_width: float): |
| x = pd.DataFrame( |
| [[sepal_length, sepal_width, petal_length, petal_width]], |
| columns=FEATURE_NAMES, |
| ) |
|
|
| pred_idx = int(MODEL.predict(x)[0]) |
| pred_label = TARGET_NAMES[pred_idx] |
|
|
| if hasattr(MODEL, "predict_proba"): |
| proba = MODEL.predict_proba(x)[0] |
| proba_dict = {str(TARGET_NAMES[i]): float(proba[i]) for i in range(len(TARGET_NAMES))} |
| else: |
| proba_dict = {TARGET_NAMES[pred_idx]: 1.0} |
|
|
| return pred_label, proba_dict |
|
|
|
|
| demo = gr.Interface( |
| fn=predict, |
| inputs=[ |
| gr.Number(label=FEATURE_NAMES[0], value=5.8), |
| gr.Number(label=FEATURE_NAMES[1], value=3.0), |
| gr.Number(label=FEATURE_NAMES[2], value=4.0), |
| gr.Number(label=FEATURE_NAMES[3], value=1.2), |
| ], |
| outputs=[ |
| gr.Textbox(label="Predicted class"), |
| gr.Label(label="Class probabilities"), |
| ], |
| title="Iris Classification (KNN)", |
| description=( |
| "KNN classifier trained on the classic Iris dataset. " |
| "Enter measurements and get a predicted species + probabilities." |
| ), |
| ) |
|
|
|
|
| if __name__ == "__main__": |
| demo.launch() |
|
|