CoolFace
Apppublic

bhavin273/demand-prediction-api

sourceHugging Facemitupdated 8mo agoView on Hugging Face
0likes
train_model.py207 linesDownload Raw Back to root
1import numpy as np2import pandas as pd3from xgboost import XGBRegressor4from sklearn.preprocessing import LabelEncoder5from sklearn.metrics import r2_score6import pickle7from datetime import datetime8from db_client import get_db_connection, COLLECTION9 10# ================= CONFIG =================11MODEL_PATH = "demand_prediction_model.pkl"12RANDOM_STATE = 4213le = LabelEncoder()14# =========================================15 16 17# ------------------------------------------------------------------18# DATA LOADING19# ------------------------------------------------------------------20def export_data_from_mongo():21    print("๐Ÿ”น Exporting data from MongoDB...")22    with get_db_connection() as db:23        collection = db[COLLECTION]24        df = pd.DataFrame(list(collection.find({}, {"_id": 0})))25        print(f"โœ… Loaded {len(df)} rows from MongoDB")26        return df27 28 29# ------------------------------------------------------------------30# TRAIN / VALIDATION SPLIT31# ------------------------------------------------------------------32def per_h3_time_split(data):33    train_raw = []34    val_raw = []35    val_true_list = []36 37    for h3_cell, group in data.groupby("h3_cell"):38        group = group.sort_values("timestamp")39        val_part = group.tail(24).copy()40        train_part = group.iloc[:-24].copy()41 42        val_true_list.append(val_part[['h3_cell', 'timestamp', 'demand']].copy())43        val_part["demand"] = np.nan44 45        train_raw.append(train_part)46        val_raw.append(val_part)47 48    return (49        pd.concat(train_raw).reset_index(drop=True),50        pd.concat(val_raw).reset_index(drop=True),51        pd.concat(val_true_list).reset_index(drop=True),52    )53 54 55# ------------------------------------------------------------------56# TRAINING FEATURE PIPELINE (HISTORICAL DATA ONLY)57# ------------------------------------------------------------------58def prepare_training_features(data):59    print("Preparing TRAINING features...")60 61    data["timestamp"] = pd.to_datetime(data["timestamp"])62    data = data.sort_values(["h3_cell", "timestamp"]).reset_index(drop=True)63 64    data["Weekday"] = data["timestamp"].dt.weekday65    data["Month"] = data["timestamp"].dt.month66    data["Quarter"] = data["timestamp"].dt.quarter67    data["day_number"] = (data["timestamp"] - data["timestamp"].min()).dt.days68    data["trend_sq"] = data["day_number"] ** 269 70    data["h3_cell_enc"] = le.fit_transform(data["h3_cell"])71 72    # ๐Ÿšจ TRAINING ONLY โ€” safe to drop NaNs here73    data = data.dropna().reset_index(drop=True)74 75    feature_columns = [76        "hour_sin",77        "hour_cos",78        "is_weekend",79        "isHoliday",80        "neighbor_availability",81        "h3_cell_enc",82        "Weekday",83        "Month",84        "Quarter",85        "day_number",86        "trend_sq",87    ]88 89    X = data[feature_columns]90    y = data["demand"]91 92    return data, X, y, feature_columns93 94 95# ------------------------------------------------------------------96# MODEL TRAINING97# ------------------------------------------------------------------98def train_model():99    data = export_data_from_mongo()100    data, X, y, feature_columns = prepare_training_features(data)101    train_raw, val_raw, val_true = per_h3_time_split(data)102 103    model = XGBRegressor(104        n_estimators=600,105        learning_rate=0.05,106        max_depth=8,107        subsample=0.9,108        colsample_bytree=0.8,109        objective="reg:squarederror",110        random_state=RANDOM_STATE,111        n_jobs=-1,112    )113 114    X_train = train_raw[feature_columns]115    y_train = train_raw["demand"]116    X_val = val_raw[feature_columns]117    y_val = val_true["demand"]118 119    model.fit(X_train, y_train)120 121    preds = model.predict(X_val)122    print("โœ… Validation Rยฒ:", r2_score(y_val, preds))123 124    with open(MODEL_PATH, "wb") as f:125        pickle.dump(126            {127                "model": model,128                "features": feature_columns,129                "encoder": le,130                "trained_at": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),131            },132            f,133        )134 135    print(f"๐Ÿ’พ Model saved โ†’ {MODEL_PATH}")136    return model137 138 139# ------------------------------------------------------------------140# INFERENCE FEATURE BUILDER141# ------------------------------------------------------------------142def build_inference_features(doc, encoder, feature_columns):143    df = pd.DataFrame([doc])144    df["timestamp"] = pd.to_datetime(df["timestamp"])145 146    df["Weekday"] = df["timestamp"].dt.weekday147    df["Month"] = df["timestamp"].dt.month148    df["Quarter"] = df["timestamp"].dt.quarter149 150    # Neutral trend values for inference151    df["day_number"] = 0152    df["trend_sq"] = 0153 154    # Encode H3155    df["h3_cell_enc"] = encoder.transform(df["h3_cell"])156 157    # Safe defaults158    df["neighbor_availability"] = df.get("neighbor_availability", 1.0)159 160    return df[feature_columns]161 162 163# ------------------------------------------------------------------164# DEMAND + PRICING PREDICTION165# ------------------------------------------------------------------166def predict_demand(h3_cell, timestamp):167    if isinstance(timestamp, str):168        timestamp = pd.to_datetime(timestamp)169 170    with open(MODEL_PATH, "rb") as f:171        model_data = pickle.load(f)172 173    model = model_data["model"]174    encoder = model_data["encoder"]175    feature_columns = model_data["features"]176 177    with get_db_connection() as db:178        doc = db[COLLECTION].find_one({179            "h3_cell": h3_cell,180            "timestamp": timestamp181        })182 183        if not doc:184            return {"error": "Record not found"}185 186        X = build_inference_features(doc, encoder, feature_columns)187 188        predicted_demand = float(model.predict(X)[0])189 190        capacity = max(doc.get("total_capacity", 1), 1)191        availability_ratio = doc.get("availability_ratio", 1)192 193        demand_pressure = predicted_demand / capacity194        scarcity = max(0, 1 - availability_ratio)195 196        demand_factor = min(max(197            1 + 0.6 * demand_pressure + 0.4 * scarcity,198            1.0199        ), 2.0)200 201        return demand_factor202 203 204# ------------------------------------------------------------------205if __name__ == "__main__":206    train_model()207