thinkingEverytime/QuantOracle
1
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 