Qionk/a-share-quant
0
1"""21-3月中长期收益率预测模块3============================4与短期预测(1-5天)完全隔离:5- 数据粒度:日线 → 周线(W-FRI重采样)6- 预测目标:单日收益率 → 未来4/8/12周累计收益率7- 特征体系:高频量价 → 趋势性/周期性特征8- 主力模型:树模型(LightGBM/XGBoost/CatBoost),禁用DL和统计模型9"""10 11import numpy as np12import pandas as pd13from sklearn.linear_model import LinearRegression14from sklearn.model_selection import TimeSeriesSplit15from sklearn.metrics import mean_squared_error, r2_score16 17# ─── 周线数据生成 ──────────────────────────────────────────────18 19def resample_to_weekly(df_daily: pd.DataFrame) -> dict:20 """21 将日线数据重采样为周线。22 23 参数:24 df_daily: 日线 DataFrame,需包含列 open, high, low, close, volume, amount25 26 返回:27 dict with:28 - df_weekly: 周线 DataFrame(中文列名)29 - data_warning: 数据不足警告 msg or None30 - total_weeks: 总周数31 """32 rename_map = {33 "open": "开盘", "high": "最高", "low": "最低",34 "close": "收盘", "volume": "成交量", "amount": "成交额",35 }36 cols_needed = list(rename_map.keys())37 available = [c for c in cols_needed if c in df_daily.columns]38 39 df_weekly = df_daily[available].resample("W-FRI").agg({40 "open": "first", "high": "max", "low": "min",41 "close": "last", "volume": "sum", "amount": "sum",42 })43 44 # 处理可能缺失的amount列45 df_weekly = df_weekly.rename(columns=rename_map)46 47 # 周收益率48 df_weekly["周收益率"] = df_weekly["收盘"].pct_change() * 10049 df_weekly = df_weekly.dropna()50 51 total_weeks = len(df_weekly)52 warning = None53 if total_weeks < 156:54 warning = f"⚠️ 中长期预测建议使用至少3年的历史数据(当前{total_weeks}周),数据量不足可能导致模型效果不佳"55 56 return {57 "df_weekly": df_weekly,58 "data_warning": warning,59 "total_weeks": total_weeks,60 }61 62 63# ─── 预测目标 ──────────────────────────────────────────────────64 65def create_weekly_targets(df_weekly: pd.DataFrame) -> pd.DataFrame:66 """67 生成未来 4/8/12 周累计收益率目标(用小数表示)。68 69 返回添加了列的 DataFrame:70 目标_1月, 目标_2月, 目标_3月71 """72 df = df_weekly.copy()73 close = df["收盘"]74 75 df["目标_1月"] = close.shift(-4) / close - 176 df["目标_2月"] = close.shift(-8) / close - 177 df["目标_3月"] = close.shift(-12) / close - 178 79 df = df.iloc[:-12] # 去掉没有未来数据的最后12行80 return df81 82 83# ─── 特征工程 ──────────────────────────────────────────────────84 85def compute_weekly_features(df_weekly: pd.DataFrame) -> pd.DataFrame:86 """87 计算周线级别的中长期特征(~18个)。88 89 返回: feature DataFrame (index 与 df_weekly 对齐)90 """91 wk = df_weekly.copy()92 close = wk["收盘"]93 volume = wk["成交量"]94 returns = wk["周收益率"]95 96 features = {}97 98 # 1. 均线系统(相对价格偏差)99 for w in [5, 10, 20, 60]:100 features[f"MA{w}"] = close.rolling(w).mean() / close - 1101 102 # 2. 趋势强度103 for w in [10, 20]:104 def _slope(series, window=w):105 out = np.full(len(series), np.nan)106 for i in range(window - 1, len(series)):107 y = series.iloc[i - window + 1:i + 1].values108 x = np.arange(window)109 out[i] = np.polyfit(x, y, 1)[0]110 return out111 features[f"趋势斜率_{w}周"] = _slope(close) / close.values112 113 # 3. 波动率特征114 features["波动率_4周"] = returns.rolling(4).std()115 features["波动率_12周"] = returns.rolling(12).std()116 117 # 4. 成交量特征118 features["成交量_4周均值"] = volume.rolling(4).mean() / volume - 1119 features["成交量_12周均值"] = volume.rolling(12).mean() / volume - 1120 121 # 量价配合度122 ret_ma4 = returns.rolling(4).mean()123 vol_dir = np.sign(features["成交量_4周均值"])124 features["量价配合度_4周"] = ret_ma4 * vol_dir125 126 # 5. RSI(14)127 delta = returns128 gain = delta.where(delta > 0, 0.0)129 loss = (-delta.where(delta < 0, 0.0))130 avg_gain = gain.rolling(14).mean()131 avg_loss = loss.rolling(14).mean()132 rs = avg_gain / avg_loss.replace(0, np.nan)133 features["RSI_14"] = 100 - (100 / (1 + rs))134 135 # 6. MACD 周线136 ema12 = close.ewm(span=12, adjust=False).mean()137 ema26 = close.ewm(span=26, adjust=False).mean()138 macd_line = ema12 - ema26139 macd_signal = macd_line.ewm(span=9, adjust=False).mean()140 features["MACD"] = macd_line141 features["MACD信号"] = macd_signal142 features["MACD柱"] = macd_line - macd_signal143 144 df_feat = pd.DataFrame(features, index=wk.index)145 df_feat = df_feat.fillna(0)146 return df_feat147 148 149# ─── 模型构建 ──────────────────────────────────────────────────150 151def build_longterm_lgb():152 """LightGBM 中长期模型(主力),严格防过拟合参数"""153 import lightgbm as lgb154 return lgb.LGBMRegressor(155 objective="regression",156 metric="rmse",157 num_leaves=15,158 max_depth=4,159 learning_rate=0.01,160 n_estimators=300,161 subsample=0.7,162 colsample_bytree=0.7,163 reg_alpha=0.1,164 reg_lambda=0.1,165 random_state=42,166 verbose=-1,167 )168 169 170def build_longterm_xgb():171 """XGBoost 中长期模型(辅助)"""172 import xgboost as xgb173 return xgb.XGBRegressor(174 objective="reg:squarederror",175 max_depth=4,176 learning_rate=0.01,177 n_estimators=300,178 subsample=0.7,179 colsample_bytree=0.7,180 reg_alpha=0.1,181 reg_lambda=1.0,182 random_state=42,183 verbosity=0,184 )185 186 187def build_longterm_catboost():188 """CatBoost 中长期模型(辅助)"""189 from catboost import CatBoostRegressor190 return CatBoostRegressor(191 depth=4,192 learning_rate=0.01,193 iterations=300,194 subsample=0.7,195 l2_leaf_reg=3,196 random_seed=42,197 verbose=0,198 allow_writing_files=False,199 )200 201 202def build_longterm_lr():203 """线性回归(基准模型)"""204 return LinearRegression()205 206 207MODEL_BUILDERS = {208 "LightGBM": build_longterm_lgb,209 "XGBoost": build_longterm_xgb,210 "CatBoost": build_longterm_catboost,211 "LinearRegression": build_longterm_lr,212}213 214# ─── Bootstrap 置信区间(特征扰动法) ─────────────────────────215 216def _bootstrap_ci(model, X: np.ndarray, n_iter: int = 100) -> tuple:217 """对树模型做特征扰动法 bootstrap,返回 (mean, lower, upper)"""218 rng = np.random.RandomState(42)219 preds = []220 for _ in range(n_iter):221 noise = rng.normal(0, 0.01, X.shape)222 p = model.predict(X + noise)223 preds.append(p)224 preds = np.array(preds)225 mean = preds.mean(axis=0)226 std = preds.std(axis=0)227 return mean, mean - 1.96 * std, mean + 1.96 * std228 229 230def _bootstrap_ci_linear(model, X: np.ndarray, n_iter: int = 100) -> tuple:231 """对线性回归做残差 bootstrap"""232 y_pred = model.predict(X)233 residuals = y_pred - y_pred # 占位,实际用零均值正态234 rng = np.random.RandomState(42)235 preds = []236 for _ in range(n_iter):237 noise = rng.normal(0, np.std(y_pred) * 0.1, len(y_pred))238 preds.append(y_pred + noise)239 preds = np.array(preds)240 std = preds.std(axis=0)241 return y_pred, y_pred - 1.96 * std, y_pred + 1.96 * std242 243 244# ─── 训练管线 ──────────────────────────────────────────────────245 246def train_long_term_models(df_weekly: pd.DataFrame, selected_models: list,247 horizon_weeks: int = 4,248 progress_callback=None) -> dict:249 """250 训练所有选中的中长期模型。251 252 参数:253 df_weekly: 已处理好的周线数据(含目标列)254 selected_models: ["LightGBM", "XGBoost", "CatBoost", "LinearRegression"] 的子集255 horizon_weeks: 4, 8, 或 12256 progress_callback: fn(pct, msg) 进度回调257 258 返回:259 {model_name: {"model": obj, "cv_rmse": float, "cv_r2": float,260 "prediction": float (小数), "confidence_interval": (low, high),261 "direction_accuracy": float, "feature_importance": dict}}262 """263 horizon_map = {4: "目标_1月", 8: "目标_2月", 12: "目标_3月"}264 target_col = horizon_map[horizon_weeks]265 266 df_with_targets = create_weekly_targets(df_weekly)267 features_df = compute_weekly_features(df_with_targets)268 269 # 对齐 index270 common_idx = features_df.index.intersection(df_with_targets.index)271 X_all = features_df.loc[common_idx].values272 y_all = df_with_targets.loc[common_idx, target_col].values273 feature_names = list(features_df.columns)274 latest_close = float(df_with_targets.loc[common_idx[-1], "收盘"])275 276 results = {}277 total = len(selected_models)278 for i, model_name in enumerate(selected_models):279 if progress_callback:280 progress_callback(i / total, f"训练 {model_name}...")281 282 builder = MODEL_BUILDERS.get(model_name)283 if builder is None:284 continue285 286 try:287 # TimeSeriesSplit 5折交叉验证288 tscv = TimeSeriesSplit(n_splits=min(5, len(X_all) // 24))289 cv_rmse_scores = []290 cv_r2_scores = []291 direction_accs = []292 293 for fold_i, (train_idx, val_idx) in enumerate(tscv.split(X_all)):294 X_tr, X_vl = X_all[train_idx], X_all[val_idx]295 y_tr, y_vl = y_all[train_idx], y_all[val_idx]296 297 model = builder()298 # LightGBM/XGBoost/CatBoost 有 early_stopping_rounds299 if model_name in ("LightGBM", "XGBoost"):300 eval_set = [(X_vl, y_vl)]301 model.fit(X_tr, y_tr, eval_set=eval_set, verbose=False)302 elif model_name == "CatBoost":303 model.fit(X_tr, y_tr, eval_set=(X_vl, y_vl),304 early_stopping_rounds=20, verbose=False)305 else: # LinearRegression306 model.fit(X_tr, y_tr)307 308 y_pred = model.predict(X_vl)309 cv_rmse_scores.append(np.sqrt(mean_squared_error(y_vl, y_pred)))310 cv_r2_scores.append(r2_score(y_vl, y_pred))311 312 # 方向准确率(预测方向 vs 实际方向)313 pred_dir = np.sign(y_pred)314 actual_dir = np.sign(y_vl)315 dir_acc = np.mean(pred_dir == actual_dir)316 direction_accs.append(dir_acc)317 318 if progress_callback and fold_i == 0:319 progress_callback((i + 0.3) / total, f"{model_name}: CV中...")320 321 avg_rmse = float(np.mean(cv_rmse_scores))322 avg_r2 = float(np.mean(cv_r2_scores))323 avg_dir_acc = float(np.mean(direction_accs))324 325 # 全量重训326 final_model = builder()327 if model_name in ("LightGBM", "XGBoost"):328 final_model.fit(X_all, y_all, verbose=False)329 elif model_name == "CatBoost":330 final_model.fit(X_all, y_all, verbose=False)331 else:332 final_model.fit(X_all, y_all)333 334 # 对未来 horizon_weeks 的预测335 latest_features = X_all[-1:].copy()336 future_pred = float(final_model.predict(latest_features)[0])337 338 # Bootstrap 置信区间339 if model_name == "LinearRegression":340 _, ci_low, ci_high = _bootstrap_ci_linear(final_model, X_all)341 else:342 _, ci_low, ci_high = _bootstrap_ci(final_model, X_all)343 ci_low = float(ci_low[-1]) if len(ci_low) > 0 else future_pred * 0.8344 ci_high = float(ci_high[-1]) if len(ci_high) > 0 else future_pred * 1.2345 346 # 特征重要性347 feat_imp = {}348 if model_name == "LightGBM":349 feat_imp = dict(zip(feature_names,350 final_model.feature_importances_.tolist()))351 elif model_name == "XGBoost":352 feat_imp = dict(zip(feature_names,353 final_model.feature_importances_.tolist()))354 elif model_name == "CatBoost":355 feat_imp = dict(zip(feature_names,356 final_model.get_feature_importance().tolist()))357 358 results[model_name] = {359 "model": final_model,360 "cv_rmse": avg_rmse,361 "cv_r2": avg_r2,362 "prediction": future_pred, # 小数形式累计收益率363 "confidence_interval": (ci_low, ci_high),364 "direction_accuracy": avg_dir_acc,365 "feature_importance": feat_imp,366 }367 368 except Exception as e:369 results[model_name] = {370 "model": None,371 "cv_rmse": float("nan"),372 "cv_r2": float("nan"),373 "prediction": 0.0,374 "confidence_interval": (0.0, 0.0),375 "direction_accuracy": 0.0,376 "feature_importance": {},377 "error": str(e),378 }379 380 if progress_callback:381 progress_callback((i + 1) / total, f"{model_name}: 完成")382 383 # 集成:等权平均384 valid_preds = [r["prediction"] for r in results.values()385 if not np.isnan(r.get("cv_rmse", float("nan")))]386 if valid_preds:387 ensemble_pred = np.mean(valid_preds)388 else:389 ensemble_pred = 0.0390 391 results["_ensemble"] = {392 "prediction": ensemble_pred,393 "latest_close": latest_close,394 "horizon_weeks": horizon_weeks,395 }396 397 return results398 399 400# ─── 风险评估 ──────────────────────────────────────────────────401 402def assess_risk(df_weekly: pd.DataFrame) -> dict:403 """404 计算中长线风险指标。405 406 返回:407 annual_vol_pct: 年化波动率(%)408 max_drawdown_pct: 最大回撤(%)409 risk_level: 低/中/高410 """411 returns = df_weekly["周收益率"].dropna()412 if len(returns) < 20:413 return {"annual_vol_pct": 0, "max_drawdown_pct": 0, "risk_level": "数据不足"}414 415 annual_vol = float(np.std(returns) * np.sqrt(52))416 417 # 最大回撤418 close = df_weekly["收盘"].values419 peak = np.maximum.accumulate(close)420 drawdown = (close - peak) / peak421 max_dd = float(np.min(drawdown) * 100)422 423 if annual_vol < 15:424 level = "低"425 elif annual_vol < 25:426 level = "中"427 else:428 level = "高"429 430 return {431 "annual_vol_pct": round(annual_vol, 1),432 "max_drawdown_pct": round(abs(max_dd), 1),433 "risk_level": level,434 }435 436 437def get_rating(pred_return_pct: float, direction_acc: float) -> str:438 """439 根据预测收益率和方向准确率给出综合评级。440 441 返回: 强烈看涨 / 看涨 / 中性 / 看跌 / 强烈看跌442 """443 if pred_return_pct > 10 and direction_acc > 0.6:444 return "强烈看涨"445 elif pred_return_pct > 3:446 return "看涨"447 elif pred_return_pct > -3:448 return "中性"449 elif pred_return_pct > -10:450 return "看跌"451 else:452 return "强烈看跌"