6041e3ff27
Python package build / build (push) Has been cancelled
- 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
218 lines
8.3 KiB
Python
218 lines
8.3 KiB
Python
"""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)
|