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)
|