Qionk/a-share-quant
0
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