Files
timesfm/stock_backtest_optimized.py
T
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

185 lines
6.9 KiB
Python
Raw 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.
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)