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)) # pad front to match length return np.concatenate([np.full(period, 50.0), rsi]) def calc_sma(prices, period): sma = np.convolve(prices, np.ones(period) / period, mode="valid") return np.concatenate([np.full(period - 1, sma[0] if len(sma) > 0 else 0.0), sma]) # ---------- 1. 获取数据 ---------- print("下载 AAPL 股票数据...", flush=True) end = datetime.now() start = end - timedelta(days=365) df = yf.download("AAPL", start=start.strftime("%Y-%m-%d"), end=end.strftime("%Y-%m-%d"), progress=False) close = df["Close"].values.flatten().astype(np.float32) volume = df["Volume"].values.flatten().astype(np.float32) dates = df.index # 分割 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_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)} 天(截至 {train_dates[-1].date()})", flush=True) print(f"实际数据: {HORIZON} 天({actual_dates[0].date()} ~ {actual_dates[-1].date()})", flush=True) # ---------- 2. 计算对数收益率 ---------- # log returns = ln(P_t / P_{t-1}),更平稳 train_logret = np.diff(np.log(train_close)).astype(np.float32) # length = N-1 actual_logret = np.diff(np.log(np.concatenate([train_close[-1:], actual_close]))).astype(np.float32) print(f"对数收益率: 均值={train_logret.mean():.6f}, 标准差={train_logret.std():.6f}", flush=True) # ---------- 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. 多窗口集成预测(对数收益率) ---------- CONTEXTS = [128, 256, 512] # 不同 context 长度 all_preds = [] all_quantiles = [] for ctx_len in CONTEXTS: ctx_data = train_logret[-ctx_len:] if len(train_logret) >= ctx_len else train_logret actual_ctx = min(ctx_len, len(ctx_data)) # 对收益率: infer_is_positive=False (可负), normalize_inputs=True model.compile( timesfm.ForecastConfig( max_context=actual_ctx, max_horizon=128, normalize_inputs=True, use_continuous_quantile_head=True, force_flip_invariance=False, infer_is_positive=False, # 收益率可正可负 fix_quantile_crossing=True, ) ) pf, qf = model.forecast(horizon=HORIZON, inputs=[ctx_data]) all_preds.append(pf[0][:HORIZON]) all_quantiles.append(qf[0][:HORIZON]) print(f" context={actual_ctx} 预测完成", flush=True) # 集成: 取平均 ensemble_pred_logret = np.mean(all_preds, axis=0) ensemble_q_logret = np.mean(all_quantiles, axis=0) # ---------- 5. 转换回价格 ---------- # P_t = P_{t-1} * exp(r_t) last_price = train_close[-1] pred_prices = [] for i in range(HORIZON): last_price = last_price * np.exp(ensemble_pred_logret[i]) pred_prices.append(last_price) pred_prices = np.array(pred_prices) # 分位数价格 q_prices = np.zeros((HORIZON, 10)) for qi in range(10): p = train_close[-1] for i in range(HORIZON): p = p * np.exp(ensemble_q_logret[i, qi]) q_prices[i, qi] = p # ---------- 6. 计算误差 ---------- actual = actual_close pred = pred_prices 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_prices[:, 1]) & (actual <= q_prices[:, 9])) * 100 in_40 = np.mean((actual >= q_prices[:, 3]) & (actual <= q_prices[:, 7])) * 100 print() print("=== 优化版回测结果:6/16 ~ 6/30 ===", 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("=== 误差指标(优化版 vs 原始版)===", flush=True) print(f" MAE: ${mae:.2f} (原始: $8.46)", flush=True) print(f" RMSE: ${rmse:.2f} (原始: $11.29)", flush=True) print(f" MAPE: {mape:.2f}% (原始: 2.98%)", flush=True) print(f" 方向准确率: {dir_acc:.1f}% (原始: 55.6%)", flush=True) print(f" 80%CI覆盖率: {in_80:.1f}% (原始: 90.0%)", flush=True) print(f" 40%CI覆盖率: {in_40:.1f}% (原始: 60.0%)", flush=True) # ---------- 7. 画对比图 ---------- fig, axes = plt.subplots(2, 1, figsize=(14, 10), sharex=False) # 上图: 价格对比 ax = axes[0] show_n = min(60, len(train_close)) 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="Optimized Forecast", color="tomato", linewidth=2, linestyle="--") ax.fill_between(x_actual, q_prices[:, 1], q_prices[:, 9], alpha=0.15, color="tomato", label="80% CI") ax.fill_between(x_actual, q_prices[:, 3], q_prices[:, 7], alpha=0.3, color="tomato", label="40% CI") ax.axvline(x=show_n - 1, color="gray", linestyle="--", alpha=0.5) ax.set_title(f"Optimized: Log-Return + Ensemble (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=["tomato" if e > 0 else "steelblue" for e in errors], alpha=0.7) ax2.axhline(y=0, color="black", linewidth=0.8) ax2.set_title("Forecast Error per Day (Optimized)", 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_optimized.png", dpi=150) print() print("Chart saved: aapl_backtest_optimized.png", flush=True) print("Done!", flush=True)