diogenemudenge's picture
Upload app.py with huggingface_hub
b0c71a8 verified
Raw
History Blame
4.76 kB
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)