Qionk/a-share-quant
0
1"""2SQLPub / MySQL 云端存储 — 训练结果共享3对 predict_app.py 保持与 supabase_store.py 相同的函数签名4"""5 6import os7import json8import uuid as _uuid9import numpy as np10from datetime import datetime11 12 13def _get_conn():14 try:15 import streamlit as st16 host = st.secrets.get("MYSQL_HOST", os.environ.get("MYSQL_HOST", ""))17 port = int(st.secrets.get("MYSQL_PORT", os.environ.get("MYSQL_PORT", "3306")))18 user = st.secrets.get("MYSQL_USER", os.environ.get("MYSQL_USER", ""))19 password = st.secrets.get("MYSQL_PASSWORD", os.environ.get("MYSQL_PASSWORD", ""))20 database = st.secrets.get("MYSQL_DATABASE", os.environ.get("MYSQL_DATABASE", ""))21 except Exception:22 host = os.environ.get("MYSQL_HOST", "")23 port = int(os.environ.get("MYSQL_PORT", "3306"))24 user = os.environ.get("MYSQL_USER", "")25 password = os.environ.get("MYSQL_PASSWORD", "")26 database = os.environ.get("MYSQL_DATABASE", "")27 28 if not host or not user or not database:29 return None30 31 import pymysql32 return pymysql.connect(33 host=host, port=port, user=user, password=password,34 database=database, charset="utf8mb4",35 connect_timeout=10, read_timeout=15, write_timeout=15,36 autocommit=True,37 )38 39 40def is_configured() -> bool:41 try:42 import streamlit as st43 host = st.secrets.get("MYSQL_HOST", os.environ.get("MYSQL_HOST", ""))44 user = st.secrets.get("MYSQL_USER", os.environ.get("MYSQL_USER", ""))45 database = st.secrets.get("MYSQL_DATABASE", os.environ.get("MYSQL_DATABASE", ""))46 except Exception:47 host = os.environ.get("MYSQL_HOST", "")48 user = os.environ.get("MYSQL_USER", "")49 database = os.environ.get("MYSQL_DATABASE", "")50 return bool(host and user and database)51 52 53def _sanitize(v):54 if isinstance(v, np.ndarray):55 v = v.tolist()56 if isinstance(v, dict):57 return {k: _sanitize(val) for k, val in v.items()}58 if isinstance(v, list):59 return [_sanitize(x) for x in v]60 if isinstance(v, float) and (np.isnan(v) or np.isinf(v)):61 return None62 if isinstance(v, (np.integer,)):63 return int(v)64 if isinstance(v, (np.floating,)):65 f = float(v)66 return None if (np.isnan(f) or np.isinf(f)) else f67 return v68 69 70def _json_dumps(obj):71 return json.dumps(obj, ensure_ascii=False, default=str)72 73 74def _json_loads(s):75 if s is None:76 return None77 if isinstance(s, (dict, list)):78 return s79 return json.loads(s)80 81 82def _serialize_stock_data(df):83 import pandas as pd84 if df is None:85 return None86 trimmed = df.tail(200)87 return {88 "index": [d.strftime("%Y-%m-%d") for d in trimmed.index],89 "columns": list(trimmed.columns),90 "data": _sanitize(trimmed.values.tolist()),91 }92 93 94def _deserialize_stock_data(obj):95 import pandas as pd96 if not obj:97 return None98 obj = _json_loads(obj) if isinstance(obj, str) else obj99 idx = pd.to_datetime(obj["index"])100 df = pd.DataFrame(obj["data"], index=idx, columns=obj["columns"])101 for col in df.columns:102 df[col] = pd.to_numeric(df[col], errors="coerce")103 return df104 105 106# ── 对外接口(签名与 supabase_store.py 一致)─────────────────────107 108def save_training_results(stock_code, stock_name, results, ensemble_weights,109 predictions, config, forecast_days, selected_models,110 stock_data=None):111 conn = _get_conn()112 if not conn:113 return None114 115 preds_json = None116 if predictions:117 preds_json = {}118 for k, v in predictions.items():119 if k == "model_predictions":120 preds_json[k] = {mk: _sanitize(mv) for mk, mv in v.items()}121 elif k == "weights":122 preds_json[k] = _sanitize(v) if isinstance(v, dict) else v123 else:124 preds_json[k] = _sanitize(v)125 126 config_summary = {127 "look_back": config.look_back,128 "n_features": config.n_features,129 "epochs": config.epochs,130 "batch_size": config.batch_size,131 "early_stop_patience": config.early_stop_patience,132 "learning_rate": config.learning_rate,133 "dropout": config.dropout,134 "lstm_units": config.lstm_units,135 "gru_units": config.gru_units,136 "cnn_filters": config.cnn_filters,137 "cnn_kernel_size": config.cnn_kernel_size,138 "patchtst_patch_size": config.patchtst_patch_size,139 "patchtst_d_model": config.patchtst_d_model,140 "patchtst_n_heads": config.patchtst_n_heads,141 "patchtst_n_encoder_layers": config.patchtst_n_encoder_layers,142 "patchtst_ff_dim": config.patchtst_ff_dim,143 "patchtst_dropout": config.patchtst_dropout,144 "tft_hidden_size": config.tft_hidden_size,145 "tft_n_heads": config.tft_n_heads,146 "tft_dropout": config.tft_dropout,147 "tft_lstm_layers": config.tft_lstm_layers,148 "cnn_gru_filters": config.cnn_gru_filters,149 "cnn_gru_gru_units": config.cnn_gru_gru_units,150 "cnn_gru_kernel_size": config.cnn_gru_kernel_size,151 "xgboost_n_estimators": config.xgboost_n_estimators,152 "xgboost_max_depth": config.xgboost_max_depth,153 "xgboost_learning_rate": config.xgboost_learning_rate,154 "xgboost_subsample": config.xgboost_subsample,155 "lightgbm_n_estimators": config.lightgbm_n_estimators,156 "lightgbm_max_depth": config.lightgbm_max_depth,157 "lightgbm_learning_rate": config.lightgbm_learning_rate,158 "lightgbm_num_leaves": config.lightgbm_num_leaves,159 "sarima_order": list(config.sarima_order),160 "sarima_seasonal_order": list(config.sarima_seasonal_order),161 }162 163 session_id = str(_uuid.uuid4())164 session_sql = """165 INSERT INTO training_sessions166 (id, stock_code, stock_name, forecast_days, selected_models,167 ensemble_weights, predictions, config_summary, last_close_price, stock_data)168 VALUES (%s,%s,%s,%s,%s,%s,%s,%s,%s,%s)169 """170 cur = conn.cursor()171 cur.execute(session_sql, (172 session_id, stock_code, stock_name, forecast_days,173 _json_dumps(selected_models),174 _json_dumps(ensemble_weights),175 _json_dumps(preds_json),176 _json_dumps(config_summary),177 float(predictions["predicted_return"][0]) if predictions else None,178 _json_dumps(_serialize_stock_data(stock_data)),179 ))180 181 model_sql = """182 INSERT INTO model_results183 (id, session_id, model_name, cv_metrics, training_time,184 future_predictions, future_conf_lower, future_conf_upper,185 test_predictions, test_actuals, test_returns, test_returns_actual,186 confidence_lower, confidence_upper)187 VALUES (%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s,%s)188 """189 for name, r in results.items():190 cur.execute(model_sql, (191 str(_uuid.uuid4()), session_id, name,192 _json_dumps(_sanitize(r.cv_metrics) if r.cv_metrics else {}),193 r.training_time,194 _json_dumps(_sanitize(r.future_predictions)),195 _json_dumps(_sanitize(r.future_conf_lower)),196 _json_dumps(_sanitize(r.future_conf_upper)),197 _json_dumps(_sanitize(r.test_predictions)),198 _json_dumps(_sanitize(r.test_actuals)),199 _json_dumps(_sanitize(r.test_returns)),200 _json_dumps(_sanitize(r.test_returns_actual)),201 _json_dumps(_sanitize(r.confidence_lower)),202 _json_dumps(_sanitize(r.confidence_upper)),203 ))204 205 cur.close()206 conn.close()207 return session_id208 209 210def load_latest_results(stock_code):211 conn = _get_conn()212 if not conn:213 return None214 215 cur = conn.cursor()216 cur.execute(217 "SELECT * FROM training_sessions WHERE stock_code=%s "218 "ORDER BY trained_at DESC LIMIT 1", (stock_code,))219 cols = [d[0] for d in cur.description]220 row = cur.fetchone()221 if not row:222 cur.close(); conn.close()223 return None224 session = dict(zip(cols, row))225 226 cur.execute("SELECT * FROM model_results WHERE session_id=%s", (session["id"],))227 mcols = [d[0] for d in cur.description]228 model_rows = [dict(zip(mcols, r)) for r in cur.fetchall()]229 230 cur.close(); conn.close()231 return session, model_rows232 233 234def list_available_stocks():235 conn = _get_conn()236 if not conn:237 return []238 cur = conn.cursor()239 cur.execute(240 "SELECT id, stock_code, stock_name, trained_at, selected_models, forecast_days "241 "FROM training_sessions ORDER BY trained_at DESC LIMIT 50"242 )243 cols = [d[0] for d in cur.description]244 rows = [dict(zip(cols, r)) for r in cur.fetchall()]245 for r in rows:246 if hasattr(r["trained_at"], "isoformat"):247 r["trained_at"] = r["trained_at"].isoformat()248 r["selected_models"] = _json_loads(r["selected_models"])249 cur.close(); conn.close()250 return rows251 252 253def load_by_session_id(session_id):254 conn = _get_conn()255 if not conn:256 return None257 258 cur = conn.cursor()259 cur.execute("SELECT * FROM training_sessions WHERE id=%s", (session_id,))260 cols = [d[0] for d in cur.description]261 row = cur.fetchone()262 if not row:263 cur.close(); conn.close()264 return None265 session = dict(zip(cols, row))266 267 cur.execute("SELECT * FROM model_results WHERE session_id=%s", (session_id,))268 mcols = [d[0] for d in cur.description]269 model_rows = [dict(zip(mcols, r)) for r in cur.fetchall()]270 271 cur.close(); conn.close()272 return session, model_rows273 274 275def restore_to_session_state(session_row, model_rows):276 import streamlit as st277 from src.predict.training import TrainResult278 279 st.session_state.stock_code = session_row["stock_code"]280 st.session_state.stock_name = session_row["stock_name"]281 st.session_state.ensemble_weights = _json_loads(session_row.get("ensemble_weights"))282 283 # 不覆盖已加载的最新 stock_data(可能比训练时保存的更新)284 if st.session_state.get("stock_data") is None:285 st.session_state.stock_data = _deserialize_stock_data(session_row.get("stock_data"))286 287 predictions_raw = _json_loads(session_row.get("predictions"))288 if predictions_raw:289 restored = {}290 for k, v in predictions_raw.items():291 if k == "model_predictions":292 restored[k] = {mk: np.array(mv) for mk, mv in v.items()}293 elif isinstance(v, list):294 restored[k] = np.array(v)295 else:296 restored[k] = v297 st.session_state.predictions = restored298 299 results = {}300 for mr in model_rows:301 tr = TrainResult(model_name=mr["model_name"])302 tr.cv_metrics = _json_loads(mr.get("cv_metrics")) or {}303 tr.training_time = mr.get("training_time", 0)304 tr.future_predictions = np.array(_json_loads(mr.get("future_predictions")) or [])305 tr.future_conf_lower = np.array(_json_loads(mr.get("future_conf_lower")) or [])306 tr.future_conf_upper = np.array(_json_loads(mr.get("future_conf_upper")) or [])307 tr.test_predictions = np.array(_json_loads(mr.get("test_predictions")) or [])308 tr.test_actuals = np.array(_json_loads(mr.get("test_actuals")) or [])309 tr.test_returns = np.array(_json_loads(mr.get("test_returns")) or [])310 tr.test_returns_actual = np.array(_json_loads(mr.get("test_returns_actual")) or [])311 tr.confidence_lower = np.array(_json_loads(mr.get("confidence_lower")) or [])312 tr.confidence_upper = np.array(_json_loads(mr.get("confidence_upper")) or [])313 tr.train_history = {}314 tr.feature_cols = []315 tr.n_features = 0316 results[mr["model_name"]] = tr317 318 st.session_state.train_results = results