CoolFace
Apppublic

Qionk/a-share-quant

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
predict_app.py3547 linesDownload Raw Back to root
1"""2A股价格预测工具3==============4功能: 11种模型(LSTM/GRU/1D-CNN/CNN-GRU/PatchTST/TFT/XGBoost/LightGBM/ARIMA/SARIMA/GARCH)预测A股个股未来收盘价5启动: streamlit run predict_app.py6依赖: pip install -r requirements.txt7"""8 9import os10import sys11import io12import json13import time as _time14import numpy as np15import pandas as pd16import streamlit as st17import plotly.graph_objects as go18from plotly.subplots import make_subplots19from datetime import datetime, timedelta, date20 21ROOT = os.path.dirname(os.path.abspath(__file__))22sys.path.insert(0, ROOT)23 24from src.predict.data_input import (25    load_from_akshare, load_from_excel, generate_template, get_stock_name,26)27from src.predict.features import compute_technical_indicators, prepare_features, create_sequences28from src.predict.preprocessing import preprocess_data29from src.predict.models import ModelConfig30from src.predict.training import (31    train_all_models, compute_ensemble_weights, ensemble_predict,32    calc_metrics, backtest_predictions, TrainingCallbacks,33    validate_training_data,34)35from src.predict.model_store import save_model, list_models, delete_model, load_model36from src.predict.continuous import (37    rolling_train, track_performance, should_retrain, get_model_status, cleanup_old_models,38)39from src.predict.mysql_store import (40    is_configured as cloud_configured,41    save_training_results as cloud_save,42    load_by_session_id as cloud_load_session,43    list_available_stocks as cloud_list_stocks,44    restore_to_session_state as cloud_restore,45)46from src.predict.stock_data_store import (47    list_db_stocks, load_stock_from_db, fetch_and_store, has_stock_data,48    list_stocks_with_status, list_stock_sessions, delete_stock_data,49)50from src.predict.fibonacci_wave import (51    detect_wave_levels, calculate_wave_fibonacci, generate_wave_fib_signals,52)53from src.predict.long_term_prediction import (54    resample_to_weekly, train_long_term_models,55    assess_risk, get_rating,56)57from src.predict.ensemble_classifier import (58    run_classifier_pipeline, get_recommended_params,59    check_params_deviation, CLF_FEATURE_COLS, create_clf_features,60    calculate_classification_metrics,61)62from src.data import load_config63 64# ═══════ 页面配置 ═══════65 66st.set_page_config(page_title="A股价格预测", page_icon="📈", layout="wide")67st.title("A股价格预测工具")68st.caption("支持 LSTM / GRU / 1D-CNN / CNN-GRU / PatchTST / TFT / XGBoost / LightGBM / ARIMA / SARIMA / GARCH 多模型集成预测")69 70config = load_config()71predict_cfg = config.get("predict", {})72 73 74def validate_target_variable(df: pd.DataFrame = None, silent: bool = False):75    """校验日收益率是否为小数形式(非百分比)。76 77    在启动时和数据加载后调用,防止回退到价格预测或百分比形式。78    返回 (is_valid: bool, message: str)79    """80    # 检查默认模型参数中是否有价格相关配置残留81    default_price_keys = ["predicted_close", "last_price", "returns_to_price"]82    for key in default_price_keys:83        if key in predict_cfg:84            return False, f"配置中包含已废弃的价格字段 '{key}',请清理"85 86    # 检查数据中的日收益率范围87    if df is not None and "日收益率" in df.columns:88        returns = df["日收益率"].dropna()89        if len(returns) > 0:90            abs_mean = abs(returns).mean()91            abs_max = abs(returns).max()92 93            # 小数形式:均值约0.005-0.03,最大值<0.11(A股涨跌停限制)94            if abs_mean > 1.0:95                return False, (96                    f"日收益率均值为 {abs_mean:.4f},疑似百分比形式(应为小数)。"97                    f"请检查 preprocessing.py 中是否误乘了 100。"98                )99            if abs_max > 0.15:100                return False, (101                    f"日收益率最大绝对值为 {abs_max:.4f},超出A股涨跌停范围。"102                    f"请检查数据是否存在异常值。"103                )104 105    if not silent:106        print("[校验] 日收益率格式检查通过(小数形式)")107    return True, "ok"108 109 110def _serialize_results() -> bytes:111    """将训练结果序列化为 JSON(不含模型对象,便于下载保存)"""112    data = {113        "stock_code": st.session_state.stock_code,114        "stock_name": st.session_state.stock_name,115        "ensemble_weights": st.session_state.ensemble_weights,116        "save_time": datetime.now().isoformat(),117    }118    if st.session_state.predictions:119        preds = st.session_state.predictions120        data["predictions"] = {121            k: v.tolist() if isinstance(v, np.ndarray) else v122            for k, v in preds.items()123            if k != "model_predictions"124        }125        if preds.get("model_predictions"):126            data["predictions"]["model_predictions"] = {127                k: v.tolist() if isinstance(v, np.ndarray) else v128                for k, v in preds["model_predictions"].items()129            }130 131    if st.session_state.train_results:132        data["train_results"] = {}133        for name, r in st.session_state.train_results.items():134            data["train_results"][name] = {135                "model_name": r.model_name,136                "cv_metrics": r.cv_metrics,137                "training_time": r.training_time,138                "test_predictions": r.test_predictions.tolist() if hasattr(r.test_predictions, 'tolist') else [],139                "test_actuals": r.test_actuals.tolist() if hasattr(r.test_actuals, 'tolist') else [],140                "test_returns": r.test_returns.tolist() if hasattr(r.test_returns, 'tolist') else [],141                "test_returns_actual": r.test_returns_actual.tolist() if hasattr(r.test_returns_actual, 'tolist') else [],142                "_last_close": r._last_close,143                "confidence_lower": r.confidence_lower.tolist() if hasattr(r.confidence_lower, 'tolist') else [],144                "confidence_upper": r.confidence_upper.tolist() if hasattr(r.confidence_upper, 'tolist') else [],145                "future_predictions": r.future_predictions.tolist() if hasattr(r.future_predictions, 'tolist') else [],146                "future_conf_lower": r.future_conf_lower.tolist() if hasattr(r.future_conf_lower, 'tolist') else [],147                "future_conf_upper": r.future_conf_upper.tolist() if hasattr(r.future_conf_upper, 'tolist') else [],148                "train_history": r.train_history,149                "feature_cols": r.feature_cols,150                "n_features": r.n_features,151            }152 153    if st.session_state.stock_data is not None:154        df = st.session_state.stock_data.copy()155        df.index = df.index.strftime("%Y-%m-%d")156        data["stock_data"] = df.to_dict(orient="split")157 158    return json.dumps(data, ensure_ascii=False, indent=2, default=str).encode("utf-8")159 160 161def _deserialize_results(content: bytes):162    """从 JSON 恢复训练结果到 session_state"""163    from src.predict.training import TrainResult164    data = json.loads(content.decode("utf-8"))165 166    st.session_state.stock_code = data.get("stock_code")167    st.session_state.stock_name = data.get("stock_name")168    st.session_state.ensemble_weights = data.get("ensemble_weights")169 170    if "stock_data" in data:171        sd = data["stock_data"]172        df = pd.DataFrame(sd["data"], columns=sd["columns"], index=sd["index"])173        df.index = pd.to_datetime(df.index)174        df.index.name = "date"175        st.session_state.stock_data = df176 177    if "train_results" in data:178        results = {}179        for name, rd in data["train_results"].items():180            tr = TrainResult(model_name=rd["model_name"])181            tr.cv_metrics = rd.get("cv_metrics", {})182            tr.training_time = rd.get("training_time", 0)183            tr.test_predictions = np.array(rd.get("test_predictions", []))184            tr.test_actuals = np.array(rd.get("test_actuals", []))185            tr.test_returns = np.array(rd.get("test_returns", []))186            tr.test_returns_actual = np.array(rd.get("test_returns_actual", []))187            tr._last_close = rd.get("_last_close", 0.0)188            tr.confidence_lower = np.array(rd.get("confidence_lower", []))189            tr.confidence_upper = np.array(rd.get("confidence_upper", []))190            tr.future_predictions = np.array(rd.get("future_predictions", []))191            tr.future_conf_lower = np.array(rd.get("future_conf_lower", []))192            tr.future_conf_upper = np.array(rd.get("future_conf_upper", []))193            tr.train_history = rd.get("train_history", {})194            tr.feature_cols = rd.get("feature_cols", [])195            tr.n_features = rd.get("n_features", 0)196            results[name] = tr197        st.session_state.train_results = results198 199    if "predictions" in data:200        preds = data["predictions"]201        restored = {}202        for k, v in preds.items():203            if k == "model_predictions":204                restored[k] = {mk: np.array(mv) for mk, mv in v.items()}205            elif isinstance(v, list):206                restored[k] = np.array(v)207            else:208                restored[k] = v209        st.session_state.predictions = restored210 211# ═══════ Session State 初始化 ═══════212 213TRAINING_LOCK = os.path.join(ROOT, "models", ".training_lock.json")214 215for key in ["stock_data", "stock_code", "stock_name", "train_results",216            "ensemble_weights", "predictions", "training_active"]:217    if key not in st.session_state:218        st.session_state[key] = None219if "training_active" not in st.session_state:220    st.session_state.training_active = False221if "cloud_stocks_cache" not in st.session_state:222    st.session_state.cloud_stocks_cache = None223    st.session_state.cloud_stocks_ts = 0224if "db_stocks" not in st.session_state:225    st.session_state.db_stocks = []226 227# ── 模型参数默认值(供 st.dialog 弹窗使用) ──228DEFAULT_MODEL_PARAMS = {229    "XGBoost": {"n_estimators": 100, "max_depth": 6, "learning_rate": 0.1, "subsample": 0.8},230    "LightGBM": {"n_estimators": 100, "max_depth": 6, "learning_rate": 0.1, "num_leaves": 31, "subsample": 0.8},231    "1D-CNN": {"look_back": 30, "filters": 32, "kernel_size": 3, "dropout": 0.2, "learning_rate": 0.001},232    "CNN-GRU": {"cnn_filters": 32, "kernel_size": 3, "gru_units": 24, "dropout": 0.2, "look_back": 30, "learning_rate": 0.0006},233    "GRU": {"units": 32, "look_back": 30, "dropout": 0.2, "learning_rate": 0.001},234    "LSTM": {"units": 32, "look_back": 30, "dropout": 0.2, "learning_rate": 0.001},235    "PatchTST": {"d_model": 128, "n_heads": 4, "n_layers": 2, "patch_size": 16, "dropout": 0.1, "look_back": 30, "learning_rate": 0.001},236    "TFT": {"hidden_size": 64, "n_heads": 4, "dropout": 0.2, "lstm_layers": 1, "look_back": 30, "learning_rate": 0.001},237    "ARIMA": {"auto": True, "p": 1, "d": 1, "q": 1},238    "SARIMA": {"p": 1, "d": 1, "q": 1, "P": 1, "D": 1, "Q": 1, "s": 5},239    "GARCH": {"p": 1, "q": 1, "dist": "t"},240}241DL_LEARNING_RATE = 0.001242DL_EPOCHS = 100243DL_BATCH_SIZE = 32244 245if "model_params" not in st.session_state:246    st.session_state.model_params = {k: dict(v) for k, v in DEFAULT_MODEL_PARAMS.items()}247if "modified_models" not in st.session_state:248    st.session_state.modified_models = set()249if "dl_epochs" not in st.session_state:250    st.session_state.dl_epochs = DL_EPOCHS251if "dl_batch_size" not in st.session_state:252    st.session_state.dl_batch_size = DL_BATCH_SIZE253if "dl_learning_rate" not in st.session_state:254    st.session_state.dl_learning_rate = DL_LEARNING_RATE255if "longterm_results" not in st.session_state:256    st.session_state.longterm_results = None257 258# ── 涨跌预测模块状态 ──259if "clf_results" not in st.session_state:260    st.session_state.clf_results = None261if "clf_params" not in st.session_state:262    st.session_state.clf_params = {263        "XGBoost": {264            "n_estimators": 100, "max_depth": 6, "learning_rate": 0.1,265            "subsample": 0.8, "colsample_bytree": 0.8,266            "min_child_weight": 1, "reg_alpha": 0.0, "reg_lambda": 1.0,267        },268        "ElasticNet": {"C": 1.0, "l1_ratio": 0.15, "max_iter": 5000, "tol": 1e-3},269    }270if "clf_recommended_params" not in st.session_state:271    st.session_state.clf_recommended_params = None272if "clf_selected_models" not in st.session_state:273    st.session_state.clf_selected_models = ["XGBoost", "ElasticNet"]274if "clf_look_back" not in st.session_state:275    st.session_state.clf_look_back = 20276if "clf_n_splits" not in st.session_state:277    st.session_state.clf_n_splits = 5278if "clf_modified_models" not in st.session_state:279    st.session_state.clf_modified_models = set()280if "clf_training_active" not in st.session_state:281    st.session_state.clf_training_active = False282if "clf_forecast_days" not in st.session_state:283    st.session_state.clf_forecast_days = 1284if "clf_threshold" not in st.session_state:285    st.session_state.clf_threshold = 0.50286if "clf_ensemble_result" not in st.session_state:287    st.session_state.clf_ensemble_result = None288if "clf_results_stock_code" not in st.session_state:289    st.session_state.clf_results_stock_code = None290if "clf_autotune_active" not in st.session_state:291    st.session_state.clf_autotune_active = False292if "clf_autotune_results" not in st.session_state:293    st.session_state.clf_autotune_results = None294if "_clf_tune_train_active" not in st.session_state:295    st.session_state._clf_tune_train_active = False296if "clf_feature_screen" not in st.session_state:297    st.session_state.clf_feature_screen = False298if "clf_top_n_features" not in st.session_state:299    st.session_state.clf_top_n_features = 50300if "clf_screening_info" not in st.session_state:301    st.session_state.clf_screening_info = None302 303 304def _param_changed(model_name, key, value, default_val):305    """检查参数是否被修改,更新 modified_models"""306    if value != default_val:307        st.session_state.modified_models.add(model_name)308    else:309        # 检查所有参数是否都恢复默认310        all_default = all(311            st.session_state.model_params[model_name].get(k) == DEFAULT_MODEL_PARAMS[model_name].get(k)312            for k in DEFAULT_MODEL_PARAMS[model_name]313        )314        if all_default:315            st.session_state.modified_models.discard(model_name)316 317 318def _training_lock_read():319    if os.path.exists(TRAINING_LOCK):320        import json as _json321        with open(TRAINING_LOCK) as f:322            return _json.load(f)323    return None324 325 326def _training_lock_write(data):327    os.makedirs(os.path.dirname(TRAINING_LOCK), exist_ok=True)328    import json as _json329    with open(TRAINING_LOCK, "w") as f:330        _json.dump(data, f)331 332 333def _training_lock_clear():334    if os.path.exists(TRAINING_LOCK):335        os.remove(TRAINING_LOCK)336 337 338def _check_orphaned_training():339    lock = _training_lock_read()340    if not lock:341        return None342    pid = lock.get("pid")343    if pid:344        try:345            os.kill(pid, 0)  # 不发送信号,只检查存在性346            return "running"347        except OSError:348            pass349    _training_lock_clear()350    return "orphaned"351 352 353orphan_status = _check_orphaned_training()354if orphan_status == "running":355    st.session_state.training_active = True356elif orphan_status == "orphaned":357    st.session_state.training_active = False358    st.warning("检测到上次训练中断(可能刷新了页面),已自动重置训练状态")359 360 361# ═══════ 模型参数弹窗 (st.dialog) ═══════362 363@st.dialog("XGBoost 参数设置")364def xgboost_dialog():365    params = st.session_state.model_params["XGBoost"]366    defaults = DEFAULT_MODEL_PARAMS["XGBoost"]367    new_lr = st.slider("学习率", 0.01, 0.30, params["learning_rate"], 0.01, format="%.2f", key="dg_xgb_lr")368    new_n = st.slider("树数量", 50, 300, params["n_estimators"], 10, key="dg_xgb_n")369    new_md = st.slider("最大深度", 3, 10, params["max_depth"], 1, key="dg_xgb_md")370    new_ss = st.slider("子样本", 0.5, 1.0, params["subsample"], 0.05, key="dg_xgb_ss")371    c1, c2 = st.columns(2)372    if c1.button("恢复默认", use_container_width=True):373        st.session_state.model_params["XGBoost"] = dict(defaults)374        st.session_state.modified_models.discard("XGBoost")375        st.rerun()376    if c2.button("确认保存", use_container_width=True, type="primary"):377        st.session_state.model_params["XGBoost"] = {378            "n_estimators": new_n, "max_depth": new_md,379            "learning_rate": new_lr, "subsample": new_ss}380        _param_changed("XGBoost", "n_estimators", new_n, defaults["n_estimators"])381        st.rerun()382 383 384@st.dialog("LightGBM 参数设置")385def lightgbm_dialog():386    params = st.session_state.model_params["LightGBM"]387    defaults = DEFAULT_MODEL_PARAMS["LightGBM"]388    new_lr = st.slider("学习率", 0.01, 0.30, params["learning_rate"], 0.01, format="%.2f", key="dg_lgb_lr")389    new_n = st.slider("树数量", 50, 300, params["n_estimators"], 10, key="dg_lgb_n")390    new_md = st.slider("最大深度", 3, 10, params["max_depth"], 1, key="dg_lgb_md")391    new_nl = st.slider("叶子数", 15, 127, params["num_leaves"], 2, key="dg_lgb_nl")392    new_ss = st.slider("子样本", 0.5, 1.0, params["subsample"], 0.05, key="dg_lgb_ss")393    c1, c2 = st.columns(2)394    if c1.button("恢复默认", use_container_width=True):395        st.session_state.model_params["LightGBM"] = dict(defaults)396        st.session_state.modified_models.discard("LightGBM")397        st.rerun()398    if c2.button("确认保存", use_container_width=True, type="primary"):399        st.session_state.model_params["LightGBM"] = {400            "n_estimators": new_n, "max_depth": new_md,401            "learning_rate": new_lr, "num_leaves": new_nl, "subsample": new_ss}402        _param_changed("LightGBM", "n_estimators", new_n, defaults["n_estimators"])403        st.rerun()404 405 406@st.dialog("1D-CNN 参数设置")407def cnn_dialog():408    params = st.session_state.model_params["1D-CNN"]409    defaults = DEFAULT_MODEL_PARAMS["1D-CNN"]410    new_lb = st.slider("时间步长", 1, 60, params["look_back"], key="dg_cnn_lb")411    new_fl = st.slider("卷积核", 16, 64, params["filters"], 8, key="dg_cnn_fl")412    new_ks = st.slider("核大小", 2, 5, params["kernel_size"], 1, key="dg_cnn_ks")413    new_do = st.slider("Dropout", 0.1, 0.4, params["dropout"], 0.05, key="dg_cnn_do")414    new_lr = st.slider("学习率", 0.0001, 0.005, params["learning_rate"], 0.0001, format="%.4f", key="dg_cnn_lr")415    c1, c2 = st.columns(2)416    if c1.button("恢复默认", use_container_width=True):417        st.session_state.model_params["1D-CNN"] = dict(defaults)418        st.session_state.modified_models.discard("1D-CNN")419        st.rerun()420    if c2.button("确认保存", use_container_width=True, type="primary"):421        st.session_state.model_params["1D-CNN"] = {422            "look_back": new_lb, "filters": new_fl, "kernel_size": new_ks,423            "dropout": new_do, "learning_rate": new_lr}424        _param_changed("1D-CNN", "filters", new_fl, defaults["filters"])425        st.rerun()426 427 428@st.dialog("CNN-GRU 参数设置")429def cnn_gru_dialog():430    params = st.session_state.model_params["CNN-GRU"]431    defaults = DEFAULT_MODEL_PARAMS["CNN-GRU"]432    new_cf = st.slider("卷积核", 16, 64, params["cnn_filters"], 8, key="dg_cg_cf")433    new_ks = st.slider("核大小", 2, 5, params["kernel_size"], 1, key="dg_cg_ks")434    new_gu = st.slider("GRU单元", 16, 64, params["gru_units"], 8, key="dg_cg_gu")435    new_lb = st.slider("时间步长", 1, 60, params["look_back"], key="dg_cg_lb")436    new_do = st.slider("Dropout", 0.1, 0.4, params["dropout"], 0.05, key="dg_cg_do")437    new_lr = st.slider("学习率", 0.0001, 0.005, params["learning_rate"], 0.0001, format="%.4f", key="dg_cg_lr")438    c1, c2 = st.columns(2)439    if c1.button("恢复默认", use_container_width=True):440        st.session_state.model_params["CNN-GRU"] = dict(defaults)441        st.session_state.modified_models.discard("CNN-GRU")442        st.rerun()443    if c2.button("确认保存", use_container_width=True, type="primary"):444        st.session_state.model_params["CNN-GRU"] = {445            "cnn_filters": new_cf, "kernel_size": new_ks,446            "gru_units": new_gu, "look_back": new_lb,447            "dropout": new_do, "learning_rate": new_lr}448        _param_changed("CNN-GRU", "cnn_filters", new_cf, defaults["cnn_filters"])449        st.rerun()450 451 452@st.dialog("GRU 参数设置")453def gru_dialog():454    params = st.session_state.model_params["GRU"]455    defaults = DEFAULT_MODEL_PARAMS["GRU"]456    new_un = st.slider("神经元", 16, 64, params["units"], 8, key="dg_gru_un")457    new_lb = st.slider("时间步长", 1, 60, params["look_back"], key="dg_gru_lb")458    new_do = st.slider("Dropout", 0.1, 0.4, params["dropout"], 0.05, key="dg_gru_do")459    new_lr = st.slider("学习率", 0.0001, 0.005, params["learning_rate"], 0.0001, format="%.4f", key="dg_gru_lr")460    c1, c2 = st.columns(2)461    if c1.button("恢复默认", use_container_width=True):462        st.session_state.model_params["GRU"] = dict(defaults)463        st.session_state.modified_models.discard("GRU")464        st.rerun()465    if c2.button("确认保存", use_container_width=True, type="primary"):466        st.session_state.model_params["GRU"] = {467            "units": new_un, "look_back": new_lb,468            "dropout": new_do, "learning_rate": new_lr}469        _param_changed("GRU", "units", new_un, defaults["units"])470        st.rerun()471 472 473@st.dialog("LSTM 参数设置")474def lstm_dialog():475    params = st.session_state.model_params["LSTM"]476    defaults = DEFAULT_MODEL_PARAMS["LSTM"]477    new_un = st.slider("神经元", 16, 64, params["units"], 8, key="dg_lstm_un")478    new_lb = st.slider("时间步长", 1, 60, params["look_back"], key="dg_lstm_lb")479    new_do = st.slider("Dropout", 0.1, 0.4, params["dropout"], 0.05, key="dg_lstm_do")480    new_lr = st.slider("学习率", 0.0001, 0.005, params["learning_rate"], 0.0001, format="%.4f", key="dg_lstm_lr")481    c1, c2 = st.columns(2)482    if c1.button("恢复默认", use_container_width=True):483        st.session_state.model_params["LSTM"] = dict(defaults)484        st.session_state.modified_models.discard("LSTM")485        st.rerun()486    if c2.button("确认保存", use_container_width=True, type="primary"):487        st.session_state.model_params["LSTM"] = {488            "units": new_un, "look_back": new_lb,489            "dropout": new_do, "learning_rate": new_lr}490        _param_changed("LSTM", "units", new_un, defaults["units"])491        st.rerun()492 493 494@st.dialog("PatchTST 参数设置")495def patchtst_dialog():496    params = st.session_state.model_params["PatchTST"]497    defaults = DEFAULT_MODEL_PARAMS["PatchTST"]498    new_dm = st.select_slider("d_model", [32, 64, 128, 256],499                              value=params["d_model"], key="dg_pt_dm")500    new_nh = st.select_slider("注意力头数", [2, 4, 8],501                              value=params["n_heads"], key="dg_pt_nh")502    new_nl = st.slider("编码器层数", 1, 4, params["n_layers"], 1, key="dg_pt_nl")503    new_ps = st.select_slider("Patch大小", [8, 16, 32],504                              value=params["patch_size"], key="dg_pt_ps")505    new_lb = st.slider("时间步长", 1, 60, params["look_back"], key="dg_pt_lb")506    new_do = st.slider("Dropout", 0.05, 0.3, params["dropout"], 0.05, key="dg_pt_do")507    new_lr = st.slider("学习率", 0.0001, 0.005, params["learning_rate"], 0.0001, format="%.4f", key="dg_pt_lr")508    c1, c2 = st.columns(2)509    if c1.button("恢复默认", use_container_width=True):510        st.session_state.model_params["PatchTST"] = dict(defaults)511        st.session_state.modified_models.discard("PatchTST")512        st.rerun()513    if c2.button("确认保存", use_container_width=True, type="primary"):514        st.session_state.model_params["PatchTST"] = {515            "d_model": new_dm, "n_heads": new_nh, "n_layers": new_nl,516            "patch_size": new_ps, "look_back": new_lb,517            "dropout": new_do, "learning_rate": new_lr}518        _param_changed("PatchTST", "d_model", new_dm, defaults["d_model"])519        st.rerun()520 521 522@st.dialog("TFT 参数设置")523def tft_dialog():524    params = st.session_state.model_params["TFT"]525    defaults = DEFAULT_MODEL_PARAMS["TFT"]526    new_hs = st.select_slider("隐藏层大小", [32, 64, 128],527                              value=params["hidden_size"], key="dg_tft_hs")528    new_nh = st.select_slider("注意力头数", [2, 4, 8],529                              value=params["n_heads"], key="dg_tft_nh")530    new_do = st.slider("Dropout", 0.1, 0.4, params["dropout"], 0.05, key="dg_tft_do")531    new_nl = st.slider("LSTM层数", 1, 3, params["lstm_layers"], 1, key="dg_tft_nl")532    new_lb = st.slider("时间步长", 1, 60, params["look_back"], key="dg_tft_lb")533    new_lr = st.slider("学习率", 0.0001, 0.005, params["learning_rate"], 0.0001, format="%.4f", key="dg_tft_lr")534    c1, c2 = st.columns(2)535    if c1.button("恢复默认", use_container_width=True):536        st.session_state.model_params["TFT"] = dict(defaults)537        st.session_state.modified_models.discard("TFT")538        st.rerun()539    if c2.button("确认保存", use_container_width=True, type="primary"):540        st.session_state.model_params["TFT"] = {541            "hidden_size": new_hs, "n_heads": new_nh,542            "dropout": new_do, "lstm_layers": new_nl,543            "look_back": new_lb, "learning_rate": new_lr}544        _param_changed("TFT", "hidden_size", new_hs, defaults["hidden_size"])545        st.rerun()546 547 548@st.dialog("ARIMA 参数设置")549def arima_dialog():550    params = st.session_state.model_params["ARIMA"]551    defaults = DEFAULT_MODEL_PARAMS["ARIMA"]552    new_auto = st.toggle("自动选参", value=params["auto"], key="dg_ar_auto")553    if not new_auto:554        c1, c2, c3 = st.columns(3)555        with c1:556            new_p = st.slider("p", 0, 5, params["p"], 1, key="dg_ar_p")557        with c2:558            new_d = st.slider("d", 0, 2, params["d"], 1, key="dg_ar_d")559        with c3:560            new_q = st.slider("q", 0, 5, params["q"], 1, key="dg_ar_q")561    else:562        new_p, new_d, new_q = params["p"], params["d"], params["q"]563    c1, c2 = st.columns(2)564    if c1.button("恢复默认", use_container_width=True):565        st.session_state.model_params["ARIMA"] = dict(defaults)566        st.session_state.modified_models.discard("ARIMA")567        st.rerun()568    if c2.button("确认保存", use_container_width=True, type="primary"):569        st.session_state.model_params["ARIMA"] = {570            "auto": new_auto, "p": new_p, "d": new_d, "q": new_q}571        _param_changed("ARIMA", "auto", new_auto, defaults["auto"])572        st.rerun()573 574 575@st.dialog("SARIMA 参数设置")576def sarima_dialog():577    params = st.session_state.model_params["SARIMA"]578    defaults = DEFAULT_MODEL_PARAMS["SARIMA"]579    c1, c2, c3 = st.columns(3)580    with c1:581        new_p = st.slider("p", 0, 3, params["p"], 1, key="dg_sa_p")582        new_d = st.slider("d", 0, 2, params["d"], 1, key="dg_sa_d")583        new_q = st.slider("q", 0, 3, params["q"], 1, key="dg_sa_q")584    with c2:585        new_P = st.slider("季节P", 0, 3, params["P"], 1, key="dg_sa_P")586        new_D = st.slider("季节D", 0, 2, params["D"], 1, key="dg_sa_D")587        new_Q = st.slider("季节Q", 0, 3, params["Q"], 1, key="dg_sa_Q")588    with c3:589        new_s = st.slider("季节周期s", 3, 66, params["s"], 1, key="dg_sa_s")590    c1, c2 = st.columns(2)591    if c1.button("恢复默认", use_container_width=True):592        st.session_state.model_params["SARIMA"] = dict(defaults)593        st.session_state.modified_models.discard("SARIMA")594        st.rerun()595    if c2.button("确认保存", use_container_width=True, type="primary"):596        st.session_state.model_params["SARIMA"] = {597            "p": new_p, "d": new_d, "q": new_q,598            "P": new_P, "D": new_D, "Q": new_Q, "s": new_s}599        _param_changed("SARIMA", "p", new_p, defaults["p"])600        st.rerun()601 602 603@st.dialog("GARCH 参数说明")604def garch_dialog():605    st.info("GARCH(1,1) 模型参数固定,不可修改")606    st.markdown("""607    - **p=1**: ARCH阶数608    - **q=1**: GARCH阶数609    - **dist='t'**: 学生t分布(捕获厚尾特性)610    - **mean='constant'**: 允许非零均值收益611    """)612    st.caption("GARCH模型用于波动率预测和风险指标计算")613 614 615# ── 涨跌预测分类器参数对话框 ──616 617@st.dialog("XGBoost 分类器参数设置")618def clf_xgboost_dialog():619    params = st.session_state.clf_params["XGBoost"]620    defaults = {621        "n_estimators": 100, "max_depth": 6, "learning_rate": 0.1,622        "subsample": 0.8, "colsample_bytree": 0.8,623        "min_child_weight": 1, "reg_alpha": 0.0, "reg_lambda": 1.0,624    }625 626    new_lr = st.number_input("学习率 learning_rate", min_value=0.001, max_value=0.50, value=params.get("learning_rate", 0.1), step=0.001, key="dg_clf_xgb_lr", format="%.3f")627    new_n = st.number_input("树数量 n_estimators", min_value=10, max_value=1000, value=params.get("n_estimators", 100), step=10, key="dg_clf_xgb_n")628    new_md = st.number_input("最大深度 max_depth", min_value=2, max_value=15, value=params.get("max_depth", 6), step=1, key="dg_clf_xgb_md")629    new_ss = st.number_input("子样本比例 subsample", min_value=0.3, max_value=1.0, value=params.get("subsample", 0.8), step=0.05, key="dg_clf_xgb_ss", format="%.2f")630    new_cbt = st.number_input("列采样比例 colsample_bytree", min_value=0.3, max_value=1.0, value=params.get("colsample_bytree", 0.8), step=0.05, key="dg_clf_xgb_cbt", format="%.2f")631    new_mcw = st.number_input("最小子节点权重 min_child_weight", min_value=0, max_value=20, value=params.get("min_child_weight", 1), step=1, key="dg_clf_xgb_mcw")632    new_ra = st.number_input("L1正则 reg_alpha", min_value=0.0, max_value=10.0, value=params.get("reg_alpha", 0.0), step=0.01, key="dg_clf_xgb_ra", format="%.2f")633    new_rl = st.number_input("L2正则 reg_lambda", min_value=0.0, max_value=10.0, value=params.get("reg_lambda", 1.0), step=0.01, key="dg_clf_xgb_rl", format="%.2f")634 635    c1, c2 = st.columns(2)636    if c1.button("恢复默认", use_container_width=True, key="clf_xgb_reset"):637        st.session_state.clf_params["XGBoost"] = dict(defaults)638        st.session_state.clf_modified_models.discard("XGBoost")639        st.rerun()640    if c2.button("确认保存", use_container_width=True, type="primary", key="clf_xgb_save"):641        st.session_state.clf_params["XGBoost"] = {642            "n_estimators": new_n, "max_depth": new_md,643            "learning_rate": new_lr, "subsample": new_ss,644            "colsample_bytree": new_cbt, "min_child_weight": new_mcw,645            "reg_alpha": new_ra, "reg_lambda": new_rl,646        }647        all_default = all(648            st.session_state.clf_params["XGBoost"].get(k) == defaults.get(k)649            for k in defaults650        )651        if all_default:652            st.session_state.clf_modified_models.discard("XGBoost")653        else:654            st.session_state.clf_modified_models.add("XGBoost")655        st.rerun()656 657 658@st.dialog("ElasticNet 分类器参数设置")659def clf_elasticnet_dialog():660    params = st.session_state.clf_params["ElasticNet"]661    defaults = {"C": 1.0, "l1_ratio": 0.15, "max_iter": 5000, "tol": 1e-3}662 663    new_C = st.number_input("正则化强度 C (越小越强)", min_value=0.0, max_value=2.0, value=params.get("C", 1.0), step=0.01,664                      key="dg_clf_en_c", format="%.2f")665    new_l1 = st.number_input("L1比例 l1_ratio", min_value=0.0, max_value=1.0, value=params.get("l1_ratio", 0.15), step=0.01,666                       key="dg_clf_en_l1", format="%.2f")667    new_mi = st.number_input("最大迭代次数 max_iter", min_value=500, max_value=20000, value=params.get("max_iter", 5000), step=500,668                       key="dg_clf_en_mi")669    new_tol = st.number_input("收敛容差 tol", min_value=1e-6, max_value=1e-2, value=params.get("tol", 1e-3), step=1e-4,670                              key="dg_clf_en_tol", format="%.6f")671 672    c1, c2 = st.columns(2)673    if c1.button("恢复默认", use_container_width=True, key="clf_en_reset"):674        st.session_state.clf_params["ElasticNet"] = dict(defaults)675        st.session_state.clf_modified_models.discard("ElasticNet")676        st.rerun()677    if c2.button("确认保存", use_container_width=True, type="primary", key="clf_en_save"):678        st.session_state.clf_params["ElasticNet"] = {679            "C": new_C, "l1_ratio": new_l1, "max_iter": new_mi, "tol": new_tol}680        all_default = all(681            st.session_state.clf_params["ElasticNet"].get(k) == defaults.get(k)682            for k in defaults683        )684        if all_default:685            st.session_state.clf_modified_models.discard("ElasticNet")686        else:687            st.session_state.clf_modified_models.add("ElasticNet")688        st.rerun()689 690 691# ═══════ 侧边栏 ═══════692 693btn_train = False694btn_clf_train = False695btn_clf_autotune = False696btn_clf_optuna = False697btn_clf_parallel = False698 699with st.sidebar:700    st.header("配置参数")701 702    # 数据来源703    st.subheader("数据输入")704    data_source = st.radio("数据来源", ["数据库加载", "Excel上传"], horizontal=True)705 706    if data_source == "数据库加载":707        # 日期区间选择(max_value 设为明天,确保今天可选)708        _tomorrow = date.today() + __import__('datetime').timedelta(days=1)709        col_d1, col_d2 = st.columns(2)710        with col_d1:711            custom_start = st.date_input("起始日期", value=date(2020, 1, 1),712                                         min_value=date(1990, 1, 1), max_value=_tomorrow,713                                         key="custom_start_date")714        with col_d2:715            custom_end = st.date_input("结束日期", value=date.today(),716                                       min_value=date(1990, 1, 1), max_value=_tomorrow,717                                       key="custom_end_date")718        start_d = custom_start.strftime("%Y%m%d")719        end_d = custom_end.strftime("%Y%m%d")720 721        if st.button("刷新列表", key="refresh_db_stocks"):722            try:723                st.session_state.db_stocks = list_stocks_with_status()724            except Exception as e:725                st.error(f"加载失败: {e}")726 727        if not st.session_state.db_stocks:728            try:729                st.session_state.db_stocks = list_stocks_with_status()730            except Exception:731                pass732 733        if st.session_state.db_stocks:734            stock_list = st.session_state.db_stocks735            stock_labels = {}736            for s in stock_list:737                # 数据新鲜度738                today_str = date.today().strftime("%Y-%m-%d")739                end_date = s["end_date"]740                if end_date == today_str:741                    freshness = "最新"742                elif end_date >= today_str:743                    freshness = "最新"744                else:745                    days_behind = (date.today() - date.fromisoformat(end_date)).days746                    freshness = f"{days_behind}天前"747                data_part = f"{s['name']} ({s['code']})  {s['rows']}天  最新: {s['end_date']} ({freshness})"748                if s["trained"]:749                    models = ", ".join(s.get("trained_models", []))750                    train_part = f"  |  已训练: {models}"751                else:752                    train_part = "  |  未训练"753                stock_labels[data_part + train_part] = s754 755            selected_label = st.selectbox(756                "选择股票", options=list(stock_labels.keys()), key="db_stock_select",757                label_visibility="collapsed",758            )759            if selected_label:760                info = stock_labels[selected_label]761 762                # 训练版本选择763                selected_session_id = None764                if info["trained"]:765                    sessions = list_stock_sessions(info["code"])766                    if sessions:767                        if len(sessions) == 1:768                            selected_session_id = sessions[0]["session_id"]769                            s = sessions[0]770                            models = ", ".join(s.get("trained_models", []))771                            st.caption(f"训练记录: {s['trained_at']}  {models}")772                        else:773                            session_labels = {}774                            for s in sessions:775                                models = ", ".join(s.get("trained_models", []))776                                session_labels[f"{s['trained_at']}  {models}"] = s777                            selected_session = st.selectbox(778                                "训练版本", options=list(session_labels.keys()),779                                index=0, key=f"session_{info['code']}",780                            )781                            selected_session_id = session_labels[selected_session]["session_id"]782 783                if st.button("加载", type="primary", use_container_width=True, key="load_db_btn"):784                    with st.spinner(f"加载 {info['name']}...并检查数据更新..."):785                        try:786                            _is_fresh = info["end_date"] >= date.today().strftime("%Y-%m-%d")787                            if not _is_fresh:788                                st.toast(f"正在更新 {info['name']} 的数据...")789                            # 若要求的起始日早于DB最早日,则删除重建790                            if start_d < info["start_date"]:791                                delete_stock_data(info["code"])792                            df, _, _ = fetch_and_store(info["code"], start_date=start_d,793                                                        end_date=end_d, max_days=10000)794                            st.session_state.stock_data = df795                            st.session_state.stock_code = info["code"]796                            st.session_state.stock_name = info["name"]797 798                            if selected_session_id:799                                try:800                                    result = cloud_load_session(selected_session_id)801                                    if result:802                                        cloud_restore(result[0], result[1])803                                    else:804                                        st.session_state.train_results = None805                                        st.session_state.predictions = None806                                        st.session_state.clf_results = None807                                        st.session_state.clf_ensemble_result = None808                                except Exception:809                                    st.session_state.train_results = None810                                    st.session_state.predictions = None811                                    st.session_state.clf_results = None812                                    st.session_state.clf_ensemble_result = None813                            else:814                                st.session_state.train_results = None815                                st.session_state.predictions = None816                                st.session_state.clf_results = None817                                st.session_state.clf_ensemble_result = None818 819                            st.success(f"加载成功: {len(df)} 条")820                            st.rerun()821                        except Exception as e:822                            st.error(f"加载失败: {e}")823        else:824            st.info("暂无数据,请先获取股票")825 826        st.divider()827        st.caption("添加新股")828        new_code = st.text_input("股票代码", key="new_stock_code", placeholder="6位代码")829 830        if new_code and len(new_code) == 6:831            exists = False832            try:833                exists = has_stock_data(new_code) or any(s["code"] == new_code for s in st.session_state.db_stocks)834            except Exception:835                pass836            if exists:837                st.info("该股票数据已存在,刷新列表即可看到。如需重新获取请点击下方按钮")838                btn_label = "重新获取数据"839            else:840                btn_label = "获取数据"841            if st.button(btn_label, type="primary", use_container_width=True, key="fetch_new_btn"):842                with st.spinner(f"正在获取 {new_code} ({start_d} ~ {end_d}) 数据..."):843                    try:844                        if exists:845                            delete_stock_data(new_code)846                        df, name, _ = fetch_and_store(new_code, start_date=start_d,847                                                       end_date=end_d, max_days=10000)848                        st.session_state.stock_data = df849                        st.session_state.stock_code = new_code850                        st.session_state.stock_name = name851                        st.session_state.train_results = None852                        st.session_state.predictions = None853                        st.session_state.clf_results = None854                        st.session_state.clf_ensemble_result = None855                        st.success(f"获取成功: {name} ({len(df)} 条)")856                        st.session_state.db_stocks = list_stocks_with_status()857                        st.rerun()858                    except Exception as e:859                        st.error(f"获取失败: {e}")860 861    else:862        st.download_button(863            "下载Excel模板",864            data=generate_template(),865            file_name="stock_data_template.xlsx",866            mime="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",867        )868        uploaded = st.file_uploader("上传Excel文件", type=["xlsx", "xls"])869        if uploaded and st.button("解析数据", type="primary", use_container_width=True):870            with st.spinner("解析中..."):871                try:872                    df = load_from_excel(uploaded)873                    st.session_state.stock_data = df874                    st.session_state.stock_code = "CUSTOM"875                    st.session_state.stock_name = uploaded.name.split(".")[0]876                    st.session_state.train_results = None877                    st.session_state.predictions = None878                    st.session_state.clf_results = None879                    st.session_state.clf_ensemble_result = None880                    st.success(f"解析成功: {len(df)} 条数据")881                except Exception as e:882                    st.error(str(e))883 884    st.divider()885 886    # 模型选择887    st.subheader("模型选择")888    all_models = [889        "LSTM", "GRU", "1D-CNN", "CNN-GRU", "PatchTST", "TFT",890        "XGBoost", "LightGBM",891        "ARIMA", "SARIMA", "GARCH",892    ]893    selected_models = st.multiselect("选择模型", all_models, default=all_models,894                                      help="DL: LSTM/GRU/1D-CNN/CNN-GRU/PatchTST/TFT | 树模型: XGBoost/LightGBM | 统计: ARIMA/SARIMA/GARCH")895 896    # 模型参数设置897    dialog_map = {898        "XGBoost": xgboost_dialog, "LightGBM": lightgbm_dialog,899        "1D-CNN": cnn_dialog, "CNN-GRU": cnn_gru_dialog,900        "GRU": gru_dialog, "LSTM": lstm_dialog,901        "PatchTST": patchtst_dialog, "TFT": tft_dialog,902        "ARIMA": arima_dialog, "SARIMA": sarima_dialog, "GARCH": garch_dialog,903    }904 905    st.caption("参数设置(🔵 = 已修改)")906    cols = st.columns(3)907    for i, m in enumerate(all_models):908        with cols[i % 3]:909            mod = " 🔵" if m in st.session_state.modified_models else ""910            if st.button(f"⚙ {m}{mod}", key=f"set_{m}", use_container_width=True):911                dialog_map[m]()912 913    use_ensemble = st.toggle("集成预测", value=True)914 915    st.divider()916 917    # 通用训练参数918    st.subheader("训练参数")919    forecast_days = st.slider("预测天数", 1, 10, 5)920    look_back = st.slider("时间步长(天)", 1, 60, predict_cfg.get("default_look_back", 30))921 922    quick_mode = st.toggle("快速模式", value=False, help="减少训练轮次和模型参数,适合快速测试")923 924    # DL通用参数925    col_e, col_b = st.columns(2)926    with col_e:927        st.session_state.dl_epochs = st.slider(928            "训练轮次", 5, 200,929            st.session_state.dl_epochs if not quick_mode else min(st.session_state.dl_epochs, 30),930            5, key="sidebar_epochs",931            help="深度学习模型的训练轮次")932    with col_b:933        st.session_state.dl_batch_size = st.select_slider(934            "批量大小", [8, 16, 32, 64],935            value=st.session_state.dl_batch_size, key="sidebar_batch")936 937    st.divider()938 939    # 模型状态940    if st.session_state.stock_code:941        status = get_model_status(st.session_state.stock_code)942        st.info(f"模型状态: {status}")943 944    # 操作按钮945    st.subheader("操作")946    if st.session_state.training_active:947        st.warning("训练进行中,请勿刷新页面")948        if st.button("强制停止训练", use_container_width=True, type="secondary"):949            lock = _training_lock_read()950            if lock and lock.get("pid"):951                try:952                    import signal953                    os.kill(lock["pid"], signal.SIGKILL)954                except Exception:955                    pass956            _training_lock_clear()957            st.session_state.training_active = False958            st.rerun()959    else:960        btn_train = st.button("训练所有模型", type="primary", use_container_width=True,961                              disabled=st.session_state.stock_data is None)962    btn_export = st.button("导出所有结果", use_container_width=True,963                           disabled=st.session_state.train_results is None)964 965    # ── 涨跌预测设置 ──966    st.divider()967    st.subheader("涨跌预测设置")968 969    clf_model_options = ["XGBoost", "ElasticNet"]970    st.session_state.clf_selected_models = st.multiselect(971        "分类模型", clf_model_options,972        default=st.session_state.clf_selected_models,973        help="XGBoost 二分类器 + ElasticNet LogisticRegression")974 975    st.session_state.clf_look_back = st.number_input(976        "特征回溯天数", min_value=1, max_value=60, value=st.session_state.clf_look_back,977        step=1, key="clf_look_back_slider",978        help="每个样本使用过去 N 天的特征")979 980    st.session_state.clf_n_splits = st.number_input(981        "扩展窗口折数", min_value=3, max_value=10, value=st.session_state.clf_n_splits,982        step=1, key="clf_n_splits_slider",983        help="时间序列扩展窗口验证的折数")984 985    st.session_state.clf_forecast_days = st.number_input(986        "预测持有天数", min_value=1, max_value=20, value=st.session_state.clf_forecast_days,987        step=1, key="clf_forecast_days_slider",988        help="T日收盘买入,持有N天后T+N日收盘卖出。1=次日卖出")989 990    st.session_state.clf_threshold = st.number_input(991        "概率阈值", min_value=0.30, max_value=0.70, value=st.session_state.clf_threshold,992        step=0.01, format="%.2f", key="clf_threshold_slider",993        help="融合概率 >= 阈值时做多,否则空仓(默认0.5)")994 995    st.caption("分类器参数(🔵 = 已修改)")996    c1, c2 = st.columns(2)997    with c1:998        mod_xgb = " 🔵" if "XGBoost" in st.session_state.clf_modified_models else ""999        if st.button(f"⚙ XGBoost{mod_xgb}", key="set_clf_xgb", use_container_width=True):1000            clf_xgboost_dialog()1001    with c2:1002        mod_en = " 🔵" if "ElasticNet" in st.session_state.clf_modified_models else ""1003        if st.button(f"⚙ ElasticNet{mod_en}", key="set_clf_en", use_container_width=True):1004            clf_elasticnet_dialog()1005 1006    # 智能推荐1007    st.divider()1008    btn_clf_recommend = st.button("智能推荐", key="clf_smart_recommend", use_container_width=True)1009    if btn_clf_recommend:1010        if st.session_state.stock_data is not None:1011            n_samples = len(st.session_state.stock_data)1012            n_features = len([c for c in CLF_FEATURE_COLS if c in st.session_state.stock_data.columns])1013            st.session_state.clf_recommended_params = get_recommended_params(n_samples, n_features)1014            st.toast(f"推荐模式: {st.session_state.clf_recommended_params['mode']}")1015 1016            st.session_state.clf_params["XGBoost"] = dict(1017                st.session_state.clf_recommended_params["xgb"])1018            st.session_state.clf_params["ElasticNet"] = dict(1019                st.session_state.clf_recommended_params["elasticnet"])1020            st.session_state.clf_look_back = st.session_state.clf_recommended_params["look_back"]1021            st.session_state.clf_n_splits = st.session_state.clf_recommended_params["n_splits"]1022            st.rerun()1023        else:1024            st.warning("请先加载数据")1025 1026    if st.session_state.clf_recommended_params is not None:1027        rec = st.session_state.clf_recommended_params1028        st.info(f"当前推荐: **{rec['mode']}** 模式 (样本数反馈)")1029        cur = st.session_state.clf_params1030        dev_warnings = check_params_deviation(cur, rec)1031        if dev_warnings:1032            for w in dev_warnings:1033                st.warning(w)1034 1035    # 训练触发1036    if st.session_state.clf_training_active:1037        st.warning("涨跌预测训练中...")1038    else:1039        btn_clf_train = st.button("开始涨跌训练", type="primary", use_container_width=True,1040            disabled=st.session_state.stock_data is None or len(st.session_state.clf_selected_models) == 0)1041 1042    # 自动调参1043    st.divider()1044    st.caption("自动调参")1045    clf_target_auc = st.number_input("目标AUC", min_value=0.50, max_value=0.70,1046        value=0.53, step=0.01, format="%.2f", key="clf_target_auc")1047    _at_c1, _at_c2 = st.columns(2)1048    with _at_c1:1049        clf_max_trials = st.number_input("随机/并行次数", min_value=5, max_value=100,1050            value=30, step=5, key="clf_max_trials")1051    with _at_c2:1052        clf_max_trials_optuna = st.number_input("贝叶斯次数", min_value=5, max_value=100,1053            value=15, step=5, key="clf_max_trials_optuna")1054    _at_b1, _at_b2 = st.columns(2)1055    with _at_b1:1056        btn_clf_autotune = st.button("随机调参", key="btn_clf_autotune", use_container_width=True,1057            disabled=st.session_state.stock_data is None or st.session_state.clf_autotune_active)1058    with _at_b2:1059        btn_clf_optuna = st.button("智能调参(贝叶斯)", key="btn_clf_optuna", use_container_width=True,1060        disabled=st.session_state.stock_data is None or st.session_state.clf_autotune_active)1061    btn_clf_parallel = st.button("⚡ 并行调参 (多核加速)", key="btn_clf_parallel", use_container_width=True,1062        disabled=st.session_state.stock_data is None or st.session_state.clf_autotune_active)1063 1064 1065# ═══════ 构建 ModelConfig ═══════1066 1067def _build_config():1068    dl_cfg = predict_cfg.get("dl", {})1069    qm = predict_cfg.get("quick_mode", {})1070    pt_cfg = predict_cfg.get("patchtst", {})1071    tf_cfg = predict_cfg.get("tft", {})1072    mp = st.session_state.model_params1073 1074    epochs = st.session_state.dl_epochs1075    batch_size = st.session_state.dl_batch_size1076 1077    if quick_mode:1078        epochs = min(epochs, qm.get("epochs", 10))1079        units_lstm = [qm.get("lstm_units", [32, 16])[0], qm.get("lstm_units", [32, 16])[1]]1080        units_gru = [qm.get("gru_units", [32, 16])[0], qm.get("gru_units", [32, 16])[1]]1081        cnn_filters = [qm.get("cnn_filters", [32, 16])[0], qm.get("cnn_filters", [32, 16])[1]]1082        cnn_gru_cf = [qm.get("cnn_gru_filters", [32, 16])[0], qm.get("cnn_gru_filters", [32, 16])[1]]1083        cnn_gru_gu = [qm.get("cnn_gru_gru_units", [24])[0]]1084        patchtst_d_model = qm.get("patchtst_d_model", 16)1085        patchtst_n_layers = qm.get("patchtst_n_encoder_layers", 1)1086        tft_hidden = qm.get("tft_hidden_size", 8)1087        tft_n_heads_val = qm.get("tft_n_heads", 2)1088        xgb_n_estimators = min(mp["XGBoost"]["n_estimators"], 100)1089        lgb_n_estimators = min(mp["LightGBM"]["n_estimators"], 100)1090        xgb_max_depth = mp["XGBoost"]["max_depth"]1091        xgb_lr = mp["XGBoost"]["learning_rate"]1092        lgb_max_depth = mp["LightGBM"]["max_depth"]1093        lgb_num_leaves = mp["LightGBM"]["num_leaves"]1094    else:1095        # DL: 使用模型专属参数1096        lstm_p = mp["LSTM"]1097        gru_p = mp["GRU"]1098        cnn_p = mp["1D-CNN"]1099        cg_p = mp["CNN-GRU"]1100        pt_p = mp["PatchTST"]1101        tft_p = mp["TFT"]1102        xgb_p = mp["XGBoost"]1103        lgb_p = mp["LightGBM"]1104 1105        units_lstm = [lstm_p["units"], lstm_p["units"] // 2]1106        units_gru = [gru_p["units"], gru_p["units"] // 2]1107        cnn_filters = [cnn_p["filters"], cnn_p["filters"] // 2]1108        cnn_gru_cf = [cg_p["cnn_filters"], cg_p["cnn_filters"] // 2]1109        cnn_gru_gu = [cg_p["gru_units"]]1110        patchtst_d_model = pt_p["d_model"]1111        patchtst_n_layers = pt_p["n_layers"]1112        tft_hidden = tft_p["hidden_size"]1113        tft_n_heads_val = tft_p["n_heads"]1114        xgb_n_estimators = xgb_p["n_estimators"]1115        xgb_max_depth = xgb_p["max_depth"]1116        xgb_lr = xgb_p["learning_rate"]1117        lgb_n_estimators = lgb_p["n_estimators"]1118        lgb_max_depth = lgb_p["max_depth"]1119        lgb_num_leaves = lgb_p["num_leaves"]1120 1121    # 小样本自动检测1122    data = st.session_state.stock_data1123    sm_cfg = predict_cfg.get("small_sample", {})1124    if data is not None and len(data) < sm_cfg.get("threshold", 200):1125        units_lstm = sm_cfg.get("lstm_units", [32, 16])1126        units_gru = sm_cfg.get("gru_units", [32, 16])1127 1128    sarima_p = mp["SARIMA"]1129    arima_p = mp["ARIMA"]1130    garch_p = mp["GARCH"]1131 1132    # dropout:从第一个选中的DL模型获取,或使用默认值1133    dropout = dl_cfg.get("dropout", 0.2)1134    dl_selected = [m for m in selected_models if m in ("LSTM", "GRU", "1D-CNN", "CNN-GRU", "PatchTST", "TFT")]1135    if dl_selected:1136        first_dl = dl_selected[0]1137        dl_params = mp.get(first_dl, {})1138        dropout = dl_params.get("dropout", dropout)1139 1140    def _get_lr(model_name):1141        return mp.get(model_name, {}).get("learning_rate", DL_LEARNING_RATE)1142 1143    return ModelConfig(1144        look_back=look_back,1145        epochs=epochs,1146        batch_size=batch_size,1147        learning_rate=DL_LEARNING_RATE,1148        dropout=dropout,1149        lstm_units=units_lstm,1150        gru_units=units_gru,1151        cnn_filters=cnn_filters,1152        cnn_kernel_size=mp.get("1D-CNN", {}).get("kernel_size", dl_cfg.get("cnn_kernel_size", 3)),1153        cnn_gru_filters=cnn_gru_cf,1154        cnn_gru_gru_units=cnn_gru_gu,1155        cnn_gru_kernel_size=mp.get("CNN-GRU", {}).get("kernel_size", dl_cfg.get("cnn_kernel_size", 3)),1156        early_stop_patience=3 if quick_mode else dl_cfg.get("early_stop_patience", 10),1157        patchtst_patch_size=mp.get("PatchTST", {}).get("patch_size", pt_cfg.get("patch_size", 16)),1158        patchtst_d_model=patchtst_d_model,1159        patchtst_n_heads=mp.get("PatchTST", {}).get("n_heads", pt_cfg.get("n_heads", 4)),1160        patchtst_n_encoder_layers=patchtst_n_layers,1161        patchtst_ff_dim=pt_cfg.get("ff_dim", 256),1162        patchtst_dropout=mp.get("PatchTST", {}).get("dropout", pt_cfg.get("dropout", 0.1)),1163        tft_hidden_size=tft_hidden,1164        tft_n_heads=tft_n_heads_val,1165        tft_dropout=mp.get("TFT", {}).get("dropout", tf_cfg.get("dropout", 0.2)),1166        tft_lstm_layers=mp.get("TFT", {}).get("lstm_layers", tf_cfg.get("lstm_layers", 1)),1167        # Per-model DL learning rates1168        lstm_lr=_get_lr("LSTM"),1169        gru_lr=_get_lr("GRU"),1170        cnn_lr=_get_lr("1D-CNN"),1171        cnn_gru_lr=_get_lr("CNN-GRU"),1172        patchtst_lr=_get_lr("PatchTST"),1173        tft_lr=_get_lr("TFT"),1174        xgboost_n_estimators=xgb_n_estimators,1175        xgboost_max_depth=xgb_max_depth,1176        xgboost_learning_rate=xgb_lr,1177        xgboost_subsample=mp["XGBoost"]["subsample"],1178        lightgbm_n_estimators=lgb_n_estimators,1179        lightgbm_max_depth=lgb_max_depth,1180        lightgbm_learning_rate=mp["LightGBM"]["learning_rate"],1181        lightgbm_num_leaves=lgb_num_leaves,1182        lightgbm_subsample=mp["LightGBM"]["subsample"],1183        sarima_order=(sarima_p["p"], sarima_p["d"], sarima_p["q"]),1184        sarima_seasonal_order=(sarima_p["P"], sarima_p["D"], sarima_p["Q"], sarima_p["s"]),1185        garch_p=garch_p["p"],1186        garch_q=garch_p["q"],1187        garch_dist=garch_p["dist"],1188    )1189 1190 1191# ═══════ 实时训练回调 ═══════1192 1193class StreamlitTrainingCallbacks(TrainingCallbacks):1194    """将训练回调连接到Streamlit UI容器"""1195 1196    def __init__(self, model_containers, overall_progress, overall_status, log_container):1197        self.containers = model_containers1198        self.progress = overall_progress1199        self.status = overall_status1200        self.log = log_container

Showing the first 1,200 of 3547 lines. Download the file for the rest.