diogenemudenge commited on
Commit
215a86b
·
verified ·
1 Parent(s): 844d778

Upload app.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. app.py +154 -0
app.py ADDED
@@ -0,0 +1,154 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import time
3
+ import json
4
+ import joblib
5
+ import logging
6
+ import traceback
7
+ from typing import List, Any
8
+
9
+ import numpy as np
10
+ import pandas as pd
11
+ from flask import Flask, request, jsonify
12
+
13
+ # ----------------------------
14
+ # Config
15
+ # ----------------------------
16
+ MODEL_PATH = os.environ.get(
17
+ "MODEL_PATH",
18
+ "/content/superkart_best_tuned_RandomForest.joblib"
19
+ )
20
+ ALLOWED_EXTENSIONS = {"csv"}
21
+ EXPECTED_COLUMNS: List[str] = None
22
+ ID_COLUMNS = ["Product_Id", "Store_Id"]
23
+
24
+ logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
25
+ log = logging.getLogger("superkart_api")
26
+
27
+ # ----------------------------
28
+ # Load model on startup
29
+ # ----------------------------
30
+ if not os.path.exists(MODEL_PATH):
31
+ raise FileNotFoundError(
32
+ f"MODEL_PATH not found: {MODEL_PATH}\n"
33
+ "Make sure your Drive is mounted and the path is correct."
34
+ )
35
+
36
+ t0 = time.time()
37
+ pipe = joblib.load(MODEL_PATH)
38
+ load_secs = time.time() - t0
39
+ log.info(f"Loaded pipeline from {MODEL_PATH} in {load_secs:.2f}s")
40
+
41
+ assert hasattr(pipe, "predict"), "Loaded object does not have .predict(...)"
42
+ assert "pre" in pipe.named_steps and "model" in pipe.named_steps, \
43
+ "Pipeline must contain 'pre' and 'model' steps."
44
+
45
+ try:
46
+ raw_feature_names = getattr(pipe.named_steps["pre"], "feature_names_in_", None)
47
+ except Exception:
48
+ raw_feature_names = None
49
+
50
+ app = Flask(__name__)
51
+
52
+ # ----------------------------
53
+ # Helpers
54
+ # ----------------------------
55
+ def _allowed_file(filename: str) -> bool:
56
+ return "." in filename and filename.rsplit(".", 1)[1].lower() in ALLOWED_EXTENSIONS
57
+
58
+ def _align_columns(df: pd.DataFrame) -> pd.DataFrame:
59
+ global raw_feature_names
60
+ if raw_feature_names is not None:
61
+ for col in raw_feature_names:
62
+ if col not in df.columns:
63
+ df[col] = np.nan
64
+ df = df[list(raw_feature_names)]
65
+ if EXPECTED_COLUMNS:
66
+ for col in EXPECTED_COLUMNS:
67
+ if col not in df.columns:
68
+ df[col] = np.nan
69
+ df = df[EXPECTED_COLUMNS]
70
+ return df
71
+
72
+ def _to_dataframe_from_json(payload: Any) -> pd.DataFrame:
73
+ if isinstance(payload, dict):
74
+ df = pd.DataFrame([payload])
75
+ elif isinstance(payload, list):
76
+ if not payload or not all(isinstance(x, dict) for x in payload):
77
+ raise ValueError("Payload must be a dict or list of dicts.")
78
+ df = pd.DataFrame(payload)
79
+ else:
80
+ raise ValueError("JSON payload must be object or list of objects.")
81
+ return _align_columns(df)
82
+
83
+ def _predict_df(df: pd.DataFrame) -> np.ndarray:
84
+ preds = pipe.predict(df)
85
+ return np.asarray(preds).reshape(-1).astype(float)
86
+
87
+ # ----------------------------
88
+ # Routes
89
+ # ----------------------------
90
+ @app.get("/")
91
+ def index():
92
+ return jsonify({
93
+ "service": "SuperKart Forecast API",
94
+ "status": "ok",
95
+ "endpoints": ["/health", "/model-info", "/predict", "/predict-csv"]
96
+ }), 200
97
+
98
+ @app.get("/health")
99
+ def health():
100
+ return jsonify({
101
+ "status": "ok",
102
+ "model_path": MODEL_PATH,
103
+ "loaded_in_seconds": round(load_secs, 3),
104
+ "training_features": list(raw_feature_names) if raw_feature_names is not None else None
105
+ })
106
+
107
+ @app.get("/model-info")
108
+ def model_info():
109
+ mdl = pipe.named_steps["model"]
110
+ try:
111
+ params = mdl.get_params()
112
+ except Exception:
113
+ params = str(mdl)
114
+ return jsonify({
115
+ "type": mdl.__class__.__name__,
116
+ "params": params
117
+ })
118
+
119
+ @app.post("/predict")
120
+ def predict():
121
+ try:
122
+ payload = request.get_json(silent=True)
123
+ if payload is None:
124
+ return jsonify({"error": "Invalid or empty JSON body."}), 400
125
+ df = _to_dataframe_from_json(payload)
126
+ preds = _predict_df(df)
127
+ echo_ids = {idc: df[idc].tolist() for idc in ID_COLUMNS if idc in df.columns}
128
+ return jsonify({"n": int(len(preds)), **({"ids": echo_ids} if echo_ids else {}), "predictions": preds.tolist()})
129
+ except Exception as e:
130
+ log.error("Predict error: %s\n%s", e, traceback.format_exc())
131
+ return jsonify({"error": str(e)}), 500
132
+
133
+ @app.post("/predict-csv")
134
+ def predict_csv():
135
+ try:
136
+ if "file" not in request.files:
137
+ return jsonify({"error": "No file part named 'file'."}), 400
138
+ file = request.files["file"]
139
+ if file.filename == "":
140
+ return jsonify({"error": "Empty filename."}), 400
141
+ if not _allowed_file(file.filename):
142
+ return jsonify({"error": "Only .csv files are allowed."}), 400
143
+ df = pd.read_csv(file)
144
+ df = _align_columns(df)
145
+ preds = _predict_df(df)
146
+ ids = {idc: df[idc].astype(str).tolist() for idc in ID_COLUMNS if idc in df.columns}
147
+ return jsonify({"n": int(len(preds)), **({"ids": ids} if ids else {}), "predictions": preds.tolist()})
148
+ except Exception as e:
149
+ log.error("Predict-CSV error: %s\n%s", e, traceback.format_exc())
150
+ return jsonify({"error": str(e)}), 500
151
+
152
+ if __name__ == "__main__":
153
+ port = int(os.environ.get("PORT", "8000"))
154
+ app.run(host="0.0.0.0", port=port, debug=False)