File size: 4,762 Bytes
215a86b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b0c71a8
215a86b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b0c71a8
215a86b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
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)