"""TimesFM + XReg 协变量预测:加入成交量、RSI、SPY 大盘指数""" import yfinance as yf import numpy as np import torch import timesfm import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt from datetime import datetime, timedelta def calc_rsi(prices, period=14): deltas = np.diff(prices) gains = np.where(deltas > 0, deltas, 0.0) losses = np.where(deltas < 0, -deltas, 0.0) avg_gain = np.convolve(gains, np.ones(period) / period, mode="valid") avg_loss = np.convolve(losses, np.ones(period) / period, mode="valid") avg_loss = np.where(avg_loss == 0, 1e-10, avg_loss) rs = avg_gain / avg_loss rsi = 100.0 - (100.0 / (1.0 + rs)) return np.concatenate([np.full(period, 50.0), rsi]) # 1. 获取数据 print("下载 AAPL + SPY 数据...", flush=True) end = datetime.now() start = end - timedelta(days=365) df_aapl = yf.download("AAPL", start=start.strftime("%Y-%m-%d"), end=end.strftime("%Y-%m-%d"), progress=False) df_spy = yf.download("SPY", start=start.strftime("%Y-%m-%d"), end=end.strftime("%Y-%m-%d"), progress=False) close = df_aapl["Close"].values.flatten().astype(np.float64) volume = df_aapl["Volume"].values.flatten().astype(np.float64) spy_close = df_spy["Close"].values.flatten().astype(np.float64) dates = df_aapl.index # 对齐长度 min_len = min(len(close), len(spy_close)) close = close[-min_len:] volume = volume[-min_len:] spy_close = spy_close[-min_len:] dates = dates[-min_len:] # 计算 RSI rsi = calc_rsi(close) # 分割 split_date = "2026-06-15" split_idx = None for i, d in enumerate(dates): if str(d.date()) <= split_date: split_idx = i train_close = close[: split_idx + 1] train_vol = volume[: split_idx + 1] train_spy = spy_close[: split_idx + 1] train_rsi = rsi[: split_idx + 1] train_dates = dates[: split_idx + 1] actual_close = close[split_idx + 1 :] actual_dates = dates[split_idx + 1 :] HORIZON = len(actual_close) print(f"训练数据: {len(train_close)} 天", flush=True) print(f"预测目标: {HORIZON} 天", flush=True) # 2. 准备协变量 # 动态数值协变量需要覆盖 context + horizon 的完整长度 # 对每个序列: train 部分用实际值,test 部分需要"未来值" # 对于 RSI 和 Volume,我们没有未来值,用最后一个值填充 # 对于 SPY,用训练集最后一个值填充(因为我们无法预知未来 SPY) # 协变量需要: 每个协变量是一个 list,每个元素对应一个输入序列的完整长度(context+horizon) full_len = len(train_close) + HORIZON # 成交量: train 用实际值,future 用最近 5 日均值 vol_future = np.mean(train_vol[-5:]) vol_full = np.concatenate([train_vol, np.full(HORIZON, vol_future)]) # RSI: train 用实际值,future 用 50(中性) rsi_full = np.concatenate([train_rsi, np.full(HORIZON, 50.0)]) # SPY: train 用实际值,future 用最后一个值 spy_future = train_spy[-1] spy_full = np.concatenate([train_spy, np.full(HORIZON, spy_future)]) # 3. 加载模型 print() print("加载 TimesFM 模型...", flush=True) torch.set_float32_matmul_precision("high") model = timesfm.TimesFM_2p5_200M_torch.from_pretrained( "google/timesfm-2.5-200m-pytorch", torch_compile=False ) print("模型加载完成", flush=True) # 4. 编译(return_backcast=True 是 XReg 必需的) model.compile( timesfm.ForecastConfig( max_context=512, max_horizon=128, normalize_inputs=True, use_continuous_quantile_head=True, force_flip_invariance=False, infer_is_positive=True, fix_quantile_crossing=True, return_backcast=True, # XReg 需要 ) ) print("编译完成", flush=True) # 5. 用 XReg 预测 print("运行 XReg 协变量预测...", flush=True) # 动态数值协变量: dict[str, list[list[float]]] # 每个协变量是一个 list,其中每个元素是一个序列(对应一个输入) # 这里只有一个输入序列 dynamic_num_covs = { "volume": [vol_full], "rsi": [rsi_full], "spy_close": [spy_full], } # 静态数值协变量 static_num_covs = { "avg_volume": [np.mean(train_vol)], } point_outputs, quantile_outputs = model.forecast_with_covariates( inputs=[train_close], dynamic_numerical_covariates=dynamic_num_covs, static_numerical_covariates=static_num_covs, xreg_mode="xreg + timesfm", # 先回归再预测残差 normalize_xreg_target_per_input=True, ridge=1.0, ) print("XReg 预测完成!", flush=True) # 6. 提取结果 pred = np.array(point_outputs[0][:HORIZON]) q = np.array(quantile_outputs[0]) # (horizon, 10) or (full, 10) # quantile_outputs 可能包含 backcast,取最后 HORIZON 个 if q.shape[0] > HORIZON: q = q[-HORIZON:] # 7. 计算误差 actual = actual_close mae = np.mean(np.abs(actual - pred)) rmse = np.sqrt(np.mean((actual - pred) ** 2)) mape = np.mean(np.abs((actual - pred) / actual)) * 100 actual_dir = np.diff(actual) pred_dir = np.diff(pred) dir_acc = np.mean(actual_dir * pred_dir > 0) * 100 in_80 = np.mean((actual >= q[:, 1]) & (actual <= q[:, 9])) * 100 in_40 = np.mean((actual >= q[:, 3]) & (actual <= q[:, 7])) * 100 print() print("=== XReg 协变量版回测结果 ===", flush=True) print() print(" 日期 实际价 预测价 误差 误差%", flush=True) for i in range(HORIZON): err = pred[i] - actual[i] err_pct = err / actual[i] * 100 print(f" {actual_dates[i].date()} {actual[i]:7.2f} {pred[i]:7.2f} {err:+7.2f} {err_pct:+6.2f}%", flush=True) print() print("=== 误差指标(三版对比)===", flush=True) print(f" MAE: ${mae:.2f} (原始: $8.46, 优化: $8.00)", flush=True) print(f" RMSE: ${rmse:.2f} (原始: $11.29, 优化: $11.09)", flush=True) print(f" MAPE: {mape:.2f}% (原始: 2.98%, 优化: 2.82%)", flush=True) print(f" 方向准确率: {dir_acc:.1f}% (原始: 55.6%, 优化: 44.4%)", flush=True) print(f" 80%CI覆盖率: {in_80:.1f}% (原始: 90.0%, 优化: 100.0%)", flush=True) print(f" 40%CI覆盖率: {in_40:.1f}% (原始: 60.0%, 优化: 80.0%)", flush=True) # 8. 画图 fig, axes = plt.subplots(2, 1, figsize=(14, 10)) show_n = min(60, len(train_close)) ax = axes[0] ax.plot(range(show_n), train_close[-show_n:], label="Historical", color="steelblue", linewidth=1.5) x_actual = range(show_n, show_n + HORIZON) ax.plot(x_actual, actual, label="Actual", color="forestgreen", linewidth=2, marker="o", markersize=4) ax.plot(x_actual, pred, label="XReg Forecast", color="darkorange", linewidth=2, linestyle="--") ax.fill_between(x_actual, q[:, 1], q[:, 9], alpha=0.15, color="darkorange", label="80% CI") ax.fill_between(x_actual, q[:, 3], q[:, 7], alpha=0.3, color="darkorange", label="40% CI") ax.axvline(x=show_n - 1, color="gray", linestyle="--", alpha=0.5) ax.set_title(f"XReg: AAPL + Volume/RSI/SPY (MAPE={mape:.2f}%, Dir={dir_acc:.0f}%)", fontsize=13) ax.set_ylabel("Price (USD)") ax.legend(loc="upper left") # 误差柱状图 ax2 = axes[1] errors = pred - actual ax2.bar(range(HORIZON), errors, color=["darkorange" if e > 0 else "steelblue" for e in errors], alpha=0.7) ax2.axhline(y=0, color="black", linewidth=0.8) ax2.set_title("XReg Forecast Error per Day", fontsize=13) ax2.set_xlabel("Trading Days After Split") ax2.set_ylabel("Error (USD)") ax2.set_xticks(range(HORIZON)) ax2.set_xticklabels([str(d.date()) for d in actual_dates], rotation=45, fontsize=8) plt.tight_layout() plt.savefig("aapl_backtest_xreg.png", dpi=150) print() print("Chart saved: aapl_backtest_xreg.png", flush=True) print("Done!", flush=True)