| import os |
| import time |
| import json |
| import joblib |
| import logging |
| import traceback |
| from typing import List, Any |
|
|
| import numpy as np |
| import pandas as pd |
| from flask import Flask, request, jsonify |
|
|
| MODEL_PATH = os.environ.get( |
| "MODEL_PATH", |
| "/content/superkart_best_tuned_RandomForest.joblib" |
| ) |
| ALLOWED_EXTENSIONS = {"csv"} |
| EXPECTED_COLUMNS: List[str] = None |
| ID_COLUMNS = ["Product_Id", "Store_Id"] |
|
|
| logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s") |
| log = logging.getLogger("superkart_api") |
|
|
|
|
| if not os.path.exists(MODEL_PATH): |
| raise FileNotFoundError( |
| f"MODEL_PATH not found: {MODEL_PATH}\n" |
| "Make sure your Drive is mounted and the path is correct." |
| ) |
|
|
| t0 = time.time() |
| pipe = joblib.load(MODEL_PATH) |
| load_secs = time.time() - t0 |
| log.info(f"Loaded pipeline from {MODEL_PATH} in {load_secs:.2f}s") |
|
|
| assert hasattr(pipe, "predict"), "Loaded object does not have .predict(...)" |
| assert "pre" in pipe.named_steps and "model" in pipe.named_steps, \ |
| "Pipeline must contain 'pre' and 'model' steps." |
|
|
| try: |
| raw_feature_names = getattr(pipe.named_steps["pre"], "feature_names_in_", None) |
| except Exception: |
| raw_feature_names = None |
|
|
| app = Flask(__name__) |
|
|
|
|
| def _allowed_file(filename: str) -> bool: |
| return "." in filename and filename.rsplit(".", 1)[1].lower() in ALLOWED_EXTENSIONS |
|
|
| def _align_columns(df: pd.DataFrame) -> pd.DataFrame: |
| global raw_feature_names |
| if raw_feature_names is not None: |
| for col in raw_feature_names: |
| if col not in df.columns: |
| df[col] = np.nan |
| df = df[list(raw_feature_names)] |
| if EXPECTED_COLUMNS: |
| for col in EXPECTED_COLUMNS: |
| if col not in df.columns: |
| df[col] = np.nan |
| df = df[EXPECTED_COLUMNS] |
| return df |
|
|
| def _to_dataframe_from_json(payload: Any) -> pd.DataFrame: |
| if isinstance(payload, dict): |
| df = pd.DataFrame([payload]) |
| elif isinstance(payload, list): |
| if not payload or not all(isinstance(x, dict) for x in payload): |
| raise ValueError("Payload must be a dict or list of dicts.") |
| df = pd.DataFrame(payload) |
| else: |
| raise ValueError("JSON payload must be object or list of objects.") |
| return _align_columns(df) |
|
|
| def _predict_df(df: pd.DataFrame) -> np.ndarray: |
| preds = pipe.predict(df) |
| return np.asarray(preds).reshape(-1).astype(float) |
|
|
| @app.get("/") |
| def index(): |
| return jsonify({ |
| "service": "SuperKart Forecast API", |
| "status": "ok", |
| "endpoints": ["/health", "/model-info", "/predict", "/predict-csv"] |
| }), 200 |
|
|
| @app.get("/health") |
| def health(): |
| return jsonify({ |
| "status": "ok", |
| "model_path": MODEL_PATH, |
| "loaded_in_seconds": round(load_secs, 3), |
| "training_features": list(raw_feature_names) if raw_feature_names is not None else None |
| }) |
|
|
| @app.get("/model-info") |
| def model_info(): |
| mdl = pipe.named_steps["model"] |
| try: |
| params = mdl.get_params() |
| except Exception: |
| params = str(mdl) |
| return jsonify({ |
| "type": mdl.__class__.__name__, |
| "params": params |
| }) |
|
|
| @app.post("/predict") |
| def predict(): |
| try: |
| payload = request.get_json(silent=True) |
| if payload is None: |
| return jsonify({"error": "Invalid or empty JSON body."}), 400 |
| df = _to_dataframe_from_json(payload) |
| preds = _predict_df(df) |
| echo_ids = {idc: df[idc].tolist() for idc in ID_COLUMNS if idc in df.columns} |
| return jsonify({"n": int(len(preds)), **({"ids": echo_ids} if echo_ids else {}), "predictions": preds.tolist()}) |
| except Exception as e: |
| log.error("Predict error: %s\n%s", e, traceback.format_exc()) |
| return jsonify({"error": str(e)}), 500 |
|
|
| @app.post("/predict-csv") |
| def predict_csv(): |
| try: |
| if "file" not in request.files: |
| return jsonify({"error": "No file part named 'file'."}), 400 |
| file = request.files["file"] |
| if file.filename == "": |
| return jsonify({"error": "Empty filename."}), 400 |
| if not _allowed_file(file.filename): |
| return jsonify({"error": "Only .csv files are allowed."}), 400 |
| df = pd.read_csv(file) |
| df = _align_columns(df) |
| preds = _predict_df(df) |
| ids = {idc: df[idc].astype(str).tolist() for idc in ID_COLUMNS if idc in df.columns} |
| return jsonify({"n": int(len(preds)), **({"ids": ids} if ids else {}), "predictions": preds.tolist()}) |
| except Exception as e: |
| log.error("Predict-CSV error: %s\n%s", e, traceback.format_exc()) |
| return jsonify({"error": str(e)}), 500 |
|
|
| if __name__ == "__main__": |
| port = int(os.environ.get("PORT", "8000")) |
| app.run(host="0.0.0.0", port=port, debug=False) |
|
|