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
204 lines
7.4 KiB
Python
204 lines
7.4 KiB
Python
"""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)
|