Files
timesfm/stock_backtest_trick.py
freedakgmail 6041e3ff27
Python package build / build (push) Has been cancelled
Add stock prediction experiments: AAPL forecast + 4-version backtest comparison
- stock_forecast.py: Basic AAPL price forecast
- stock_backtest.py: V1 original price prediction backtest
- stock_backtest_optimized.py: V2 log-return + ensemble optimization
- stock_backtest_xreg.py: V3 XReg covariates (volume/RSI/SPY)
- stock_backtest_trick.py: V4 SPY trend guidance + volatility adjustment
- Visualization charts for all versions
2026-07-01 21:01:59 +08:00

218 lines
8.3 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""TimesFM 纯技巧优化:
1. 对数收益率预测(更平稳)
2. 多窗口集成(128/256/512
3. SPY 大盘走势引导:先预测 SPY,用 SPY 预测的趋势辅助判断 AAPL 方向
4. 波动率调整:用近期波动率缩放置信区间
5. infer_is_positive=False(收益率可负)
"""
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
# 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.float32)
spy_close = df_spy["Close"].values.flatten().astype(np.float32)
dates = df_aapl.index
min_len = min(len(close), len(spy_close))
close = close[-min_len:]
spy_close = spy_close[-min_len:]
dates = dates[-min_len:]
# 分割
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_spy = spy_close[: split_idx + 1]
train_dates = dates[: split_idx + 1]
actual_close = close[split_idx + 1 :]
actual_spy = spy_close[split_idx + 1 :]
actual_dates = dates[split_idx + 1 :]
HORIZON = len(actual_close)
print(f"训练数据: {len(train_close)} 天 | 预测: {HORIZON}", flush=True)
# 2. 对数收益率
train_logret = np.diff(np.log(train_close)).astype(np.float32)
train_spy_logret = np.diff(np.log(train_spy)).astype(np.float32)
actual_logret = np.diff(np.log(np.concatenate([train_close[-1:], actual_close]))).astype(np.float32)
# 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. 多窗口集成预测 AAPL 收益率
CONTEXTS = [128, 256, 512]
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))
# 向上取整到 32 的倍数
actual_ctx = ((actual_ctx + 31) // 32) * 32
ctx_data = train_logret[-actual_ctx:]
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" AAPL context={actual_ctx} done", flush=True)
# 5. 预测 SPY 收益率(大盘趋势引导)
spy_ctx = train_spy_logret[-256:]
spy_ctx_len = ((len(spy_ctx) + 31) // 32) * 32
spy_ctx = train_spy_logret[-spy_ctx_len:]
model.compile(
timesfm.ForecastConfig(
max_context=spy_ctx_len,
max_horizon=128,
normalize_inputs=True,
use_continuous_quantile_head=True,
force_flip_invariance=False,
infer_is_positive=False,
fix_quantile_crossing=True,
)
)
spy_pf, _ = model.forecast(horizon=HORIZON, inputs=[spy_ctx])
print(f" SPY context={spy_ctx_len} done", flush=True)
# 6. 集成 + SPY 趋势调整
ensemble_pred_logret = np.mean(all_preds, axis=0)
ensemble_q_logret = np.mean(all_quantiles, axis=0)
spy_pred_logret = spy_pf[0][:HORIZON]
# SPY 趋势调整:如果 SPY 预测下跌,对 AAPL 预测施加向下的调整
# 计算 AAPL 对 SPY 的 beta(敏感度)
beta = np.corrcoef(train_logret[-60:], train_spy_logret[-60:])[0, 1]
print(f" AAPL-SPY 60日相关系数: {beta:.3f}", flush=True)
# 调整:将 SPY 预测的偏离均值部分 * beta 加到 AAPL 预测上
spy_mean = np.mean(train_spy_logret[-60:])
spy_deviation = spy_pred_logret - spy_mean # SPY 偏离其均值的部分
adjustment = beta * spy_deviation * 0.3 # 0.3 是调整强度,避免过度修正
adjusted_pred_logret = ensemble_pred_logret + adjustment
# 7. 波动率调整置信区间
recent_vol = np.std(train_logret[-20:])
long_vol = np.std(train_logret[-60:])
print(f" 近20日波动率: {recent_vol:.5f} | 近60日波动率: {long_vol:.5f}", flush=True)
# 如果近期波动率高于长期,扩大置信区间
vol_ratio = recent_vol / max(long_vol, 1e-8)
vol_scale = max(vol_ratio, 1.0) # 只扩大不缩小
adjusted_q_logret = ensemble_q_logret.copy()
median_idx = 5
for qi in range(10):
if qi != median_idx:
adjusted_q_logret[:, qi] = ensemble_q_logret[:, median_idx] + (
ensemble_q_logret[:, qi] - ensemble_q_logret[:, median_idx]
) * vol_scale
# 8. 转换回价格
last_price = train_close[-1]
pred_prices = []
for i in range(HORIZON):
last_price = last_price * np.exp(adjusted_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(adjusted_q_logret[i, qi])
q_prices[i, qi] = p
# 9. 误差指标
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("=== 纯技巧优化版回测结果 ===", 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, XReg: $9.21)", flush=True)
print(f" RMSE: ${rmse:.2f} (原始: $11.29, 对数收益率: $11.09, XReg: $12.47)", flush=True)
print(f" MAPE: {mape:.2f}% (原始: 2.98%, 对数收益率: 2.82%, XReg: 3.25%)", flush=True)
print(f" 方向准确率: {dir_acc:.1f}% (原始: 55.6%, 对数收益率: 44.4%, XReg: 44.4%)", flush=True)
print(f" 80%CI覆盖率: {in_80:.1f}% (原始: 90.0%, 对数收益率: 100.0%, XReg: 70.0%)", flush=True)
print(f" 40%CI覆盖率: {in_40:.1f}% (原始: 60.0%, 对数收益率: 80.0%, XReg: 40.0%)", flush=True)
# 10. 画图
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="Trick Forecast", color="purple", linewidth=2, linestyle="--")
ax.fill_between(x_actual, q_prices[:, 1], q_prices[:, 9], alpha=0.15, color="purple", label="80% CI")
ax.fill_between(x_actual, q_prices[:, 3], q_prices[:, 7], alpha=0.3, color="purple", label="40% CI")
ax.axvline(x=show_n - 1, color="gray", linestyle="--", alpha=0.5)
ax.set_title(f"Trick: LogRet + Ensemble + SPY + VolAdj (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=["purple" if e > 0 else "steelblue" for e in errors], alpha=0.7)
ax2.axhline(y=0, color="black", linewidth=0.8)
ax2.set_title("Trick 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_trick.png", dpi=150)
print()
print("Chart saved: aapl_backtest_trick.png", flush=True)
print("Done!", flush=True)