CoolFace
Apppublic

thinkingEverytime/QuantOracle

sourceHugging Faceupdated 6mo agoView on Hugging Face
1likes
train_gbdt.py88 linesDownload Raw Back to scripts
1#!/usr/bin/env python32"""Train a GBDT regressor (sklearn) on the feature table (training-only)."""3 4from __future__ import annotations5 6import argparse7 8import numpy as np9import pandas as pd10 11from quant.registry import model_version_dir, save_meta, version_id, write_latest12 13 14FEATURES = [15    "ret_1d",16    "ret_5d",17    "ret_20d",18    "vol_20d",19    "price_sma20",20    "price_sma50",21    "rsi_14",22]23 24 25def main():26    ap = argparse.ArgumentParser()27    ap.add_argument("--data", default="data/features.parquet")28    ap.add_argument("--horizon", type=int, default=5)29    ap.add_argument("--max-leaf-nodes", type=int, default=31)30    ap.add_argument("--learning-rate", type=float, default=0.05)31    args = ap.parse_args()32 33    try:34        from sklearn.ensemble import HistGradientBoostingRegressor35        import joblib36    except Exception as e:  # pragma: no cover37        raise SystemExit(f"Install training deps: pip install -r requirements-ml.txt ({e})")38 39    df = pd.read_parquet(args.data)40    df["Date"] = pd.to_datetime(df["Date"])41    df = df.sort_values("Date").dropna()42 43    dates = df["Date"].drop_duplicates().sort_values()44    cutoff = dates.iloc[int(len(dates) * 0.8)]45    train = df[df["Date"] <= cutoff]46    test = df[df["Date"] > cutoff]47 48    Xtr = train[FEATURES].to_numpy(dtype=float)49    ytr = train["target"].to_numpy(dtype=float)50    Xte = test[FEATURES].to_numpy(dtype=float)51    yte = test["target"].to_numpy(dtype=float)52 53    model = HistGradientBoostingRegressor(54        max_leaf_nodes=args.max_leaf_nodes,55        learning_rate=args.learning_rate,56        random_state=7,57    )58    model.fit(Xtr, ytr)59 60    yhat = model.predict(Xte)61    ic = float(np.corrcoef(yhat, yte)[0, 1]) if len(yte) > 10 else 0.062    hit = float((np.sign(yhat) == np.sign(yte)).mean()) if len(yte) else 0.063 64    model_id = f"gbdt_h{args.horizon}"65    v = version_id()66    out_dir = model_version_dir(model_id, v)67    out_dir.mkdir(parents=True, exist_ok=True)68 69    joblib.dump(model, out_dir / "model.joblib")70    meta = {71        "model": "gbdt",72        "horizon": args.horizon,73        "features": FEATURES,74        "cutoff": cutoff.strftime("%Y-%m-%d"),75        "rows_train": int(len(train)),76        "rows_test": int(len(test)),77        "ic": ic,78        "hit_rate": hit,79        "params": {"max_leaf_nodes": args.max_leaf_nodes, "learning_rate": args.learning_rate},80    }81    save_meta(out_dir, meta)82    write_latest(model_id, v)83    print(f"Wrote {model_id}@{v} IC={ic:.4f} hit={hit:.3f}")84 85 86if __name__ == "__main__":87    main()88