CoolFace
Apppublic

Qionk/a-share-quant

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
mysql_store.py318 linesDownload Raw Back to predict
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