Files

390 lines
14 KiB
Python

"""
Walk-Forward 验证 — 滚动 IS/OOS + 过拟合检测
核心流程:
1. 将数据切分为滚动窗口 (train + test)
2. 在每个 train 窗口上做参数优化 (IS)
3. 用最优参数在 test 窗口上回测 (OOS)
4. 汇总: IS/OOS 性能对比、衰减比、参数稳定性
判定过拟合: OOS 夏普 / IS 夏普 < 0.5 → 过拟合
用法:
from walk_forward import WalkForwardValidator
from strategies.sma_cross import SmaCrossStrategy
wf = WalkForwardValidator(train_size=300, test_size=100)
result = wf.validate(
strategy_class=SmaCrossStrategy,
df=df,
param_grid={"fast": [5,10], "slow": [20,30]},
)
print(result.summary())
print(f"过拟合: {result.is_overfit}")
"""
from __future__ import annotations
import os
from typing import Type
import numpy as np
import pandas as pd
import raptorbt
from .optimizer import StrategyOptimizer
from strategies.base import Strategy
class WalkForwardWindow:
"""单个 walk-forward 窗口的结果"""
def __init__(self, idx, train_start, train_end, test_start, test_end,
best_params, is_metrics, oos_metrics):
self.idx = idx
self.train_start = train_start
self.train_end = train_end
self.test_start = test_start
self.test_end = test_end
self.best_params = best_params
self.is_metrics = is_metrics # dict
self.oos_metrics = oos_metrics # dict
class WalkForwardResult:
"""Walk-forward 验证汇总结果"""
def __init__(self, windows: list, metric: str):
self.windows = windows
self.metric = metric
self.n_windows = len(windows)
@property
def is_sharpe_avg(self) -> float:
vals = [w.is_metrics.get("sharpe_ratio", np.nan) for w in self.windows]
return float(np.nanmean(vals)) if vals else np.nan
@property
def oos_sharpe_avg(self) -> float:
vals = [w.oos_metrics.get("sharpe_ratio", np.nan) for w in self.windows]
return float(np.nanmean(vals)) if vals else np.nan
@property
def oos_return_avg(self) -> float:
vals = [w.oos_metrics.get("total_return_pct", np.nan) for w in self.windows]
return float(np.nanmean(vals)) if vals else np.nan
@property
def oos_max_drawdown_avg(self) -> float:
vals = [w.oos_metrics.get("max_drawdown_pct", np.nan) for w in self.windows]
return float(np.nanmean(vals)) if vals else np.nan
@property
def oos_trades_total(self) -> int:
return sum(w.oos_metrics.get("total_trades", 0) for w in self.windows)
@property
def oos_win_rate_avg(self) -> float:
vals = [w.oos_metrics.get("win_rate_pct", np.nan) for w in self.windows]
return float(np.nanmean(vals)) if vals else np.nan
@property
def oos_profit_factor_avg(self) -> float:
"""OOS 平均盈利因子 (总盈利/总亏损, >1 为正期望)
注意: 单窗口 PF 可能是 inf (只有盈利单无亏损单) 或 nan (无交易),
这两种都过滤掉, 只对有效窗口求平均。
"""
vals = [w.oos_metrics.get("profit_factor", np.nan) for w in self.windows]
vals = [v for v in vals if not np.isnan(v) and not np.isinf(v)]
return float(np.mean(vals)) if vals else np.nan
@property
def oos_expectancy_avg(self) -> float:
"""OOS 平均每笔期望值 (单位: 百分比, >0 为正期望)
注意: 无交易窗口的 expectancy 是 nan, 过滤掉。
"""
vals = [w.oos_metrics.get("expectancy", np.nan) for w in self.windows]
vals = [v for v in vals if not np.isnan(v) and not np.isinf(v)]
return float(np.mean(vals)) if vals else np.nan
@property
def decay_ratio(self) -> float:
"""OOS/IS 夏普衰减比, 越接近 1.0 越好, < 0.5 判定过拟合"""
is_val = self.is_sharpe_avg
oos_val = self.oos_sharpe_avg
if np.isnan(is_val) or np.isnan(oos_val) or is_val == 0:
return np.nan
return oos_val / is_val
@property
def is_overfit(self) -> bool:
"""衰减比 < 0.5 判定为过拟合"""
decay = self.decay_ratio
if np.isnan(decay):
return True
return decay < 0.5
@property
def param_stability(self) -> dict:
"""参数稳定性: 各参数值被选中的频次分布"""
stability = {}
for w in self.windows:
for k, v in w.best_params.items():
if k not in stability:
stability[k] = {}
key = str(v)
stability[k][key] = stability[k].get(key, 0) + 1
return stability
def export(self, path: str):
"""导出逐窗口明细到 CSV"""
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
rows = []
for w in self.windows:
row = {
"window": w.idx,
"train_start": w.train_start, "train_end": w.train_end,
"test_start": w.test_start, "test_end": w.test_end,
"best_params": str(w.best_params),
}
for k, v in w.is_metrics.items():
row[f"is_{k}"] = v
for k, v in w.oos_metrics.items():
row[f"oos_{k}"] = v
rows.append(row)
pd.DataFrame(rows).to_csv(path, index=False, encoding="utf-8-sig")
def summary(self) -> str:
decay_str = f"{self.decay_ratio:.2%}" if not np.isnan(self.decay_ratio) else "N/A"
pf_str = f"{self.oos_profit_factor_avg:.2f}" if not np.isnan(self.oos_profit_factor_avg) else "N/A"
exp_str = f"{self.oos_expectancy_avg:.4f}" if not np.isnan(self.oos_expectancy_avg) else "N/A"
lines = [
f"Walk-Forward 验证: {self.n_windows} 个窗口",
f" IS 平均夏普: {self.is_sharpe_avg:>8.4f}",
f" OOS 平均夏普: {self.oos_sharpe_avg:>8.4f}",
f" 衰减比: {decay_str:>8}",
f" 过拟合判定: {'是 ⚠️' if self.is_overfit else '否 ✓'}",
f" OOS 总交易数: {self.oos_trades_total:>8d}",
f" OOS 平均收益: {self.oos_return_avg:>8.2f}%",
f" OOS 盈利因子: {pf_str:>8}",
f" OOS 每笔期望: {exp_str:>8}",
f" OOS 平均回撤: {self.oos_max_drawdown_avg:>8.2f}%",
f" OOS 平均胜率: {self.oos_win_rate_avg:>8.1f}%",
f" 参数稳定性:",
]
for param, dist in self.param_stability.items():
dist_str = ", ".join(
f"{k}:{v}" for k, v in sorted(dist.items(), key=lambda x: -x[1])
)
lines.append(f" {param}: {dist_str}")
return "\n".join(lines)
def to_dict(self) -> dict:
"""转为可 JSON 序列化的字典 (不含完整 windows 明细, 仅汇总)"""
decay = self.decay_ratio
return {
"n_windows": self.n_windows,
"metric": self.metric,
"is_sharpe_avg": self.is_sharpe_avg,
"oos_sharpe_avg": self.oos_sharpe_avg,
"oos_return_avg": self.oos_return_avg,
"oos_max_drawdown_avg": self.oos_max_drawdown_avg,
"oos_trades_total": self.oos_trades_total,
"oos_win_rate_avg": self.oos_win_rate_avg,
"oos_profit_factor_avg": self.oos_profit_factor_avg,
"oos_expectancy_avg": self.oos_expectancy_avg,
"decay_ratio": None if np.isnan(decay) else decay,
"is_overfit": self.is_overfit,
"param_stability": self.param_stability,
"windows": [
{
"idx": w.idx,
"train_start": w.train_start, "train_end": w.train_end,
"test_start": w.test_start, "test_end": w.test_end,
"best_params": w.best_params,
"is_metrics": w.is_metrics,
"oos_metrics": w.oos_metrics,
}
for w in self.windows
],
}
class WalkForwardValidator:
"""
Walk-Forward 验证器
参数:
train_size: 训练窗口大小 (bars)
test_size: 测试窗口大小 (bars)
step: 滚动步长 (默认 = test_size, 即非重叠)
用法:
wf = WalkForwardValidator(train_size=300, test_size=100)
result = wf.validate(
strategy_class=SmaCrossStrategy,
df=df,
param_grid={"fast": [5,10], "slow": [20,30]},
)
"""
def __init__(self, train_size: int = 300, test_size: int = 100, step: int | None = None):
self.train_size = train_size
self.test_size = test_size
self.step = step or test_size
def validate(
self,
strategy_class: Type[Strategy],
df: pd.DataFrame,
param_grid: dict,
metric: str = "sharpe_ratio",
symbol: str = "WF",
verbose: bool = True,
) -> WalkForwardResult:
"""
执行 walk-forward 验证
返回:
WalkForwardResult
"""
n = len(df)
min_required = self.train_size + self.test_size
if n < min_required:
raise ValueError(
f"数据不足: {n} bars, 至少需要 {min_required} bars "
f"(train={self.train_size} + test={self.test_size})"
)
windows = []
start = 0
window_idx = 0
while start + min_required <= n:
train_start = start
train_end = start + self.train_size
test_start = train_end
test_end = min(train_end + self.test_size, n)
if verbose:
print(f"\n 窗口 {window_idx}: "
f"train=[{train_start}:{train_end}] "
f"test=[{test_start}:{test_end}]")
train_df = df.iloc[train_start:train_end].reset_index(drop=True)
test_df = df.iloc[test_start:test_end].reset_index(drop=True)
# Step 1: 在训练集上优化参数
optimizer = StrategyOptimizer(metric=metric, direction=None)
opt_result = optimizer.optimize(
strategy_class=strategy_class,
df=train_df,
param_grid=param_grid,
symbol=f"{symbol}_train",
verbose=False,
)
best_params = opt_result.best_params
if not best_params:
if verbose:
print(f" ⚠️ 训练集无有效参数, 跳过")
start += self.step
window_idx += 1
continue
if verbose:
params_str = ", ".join(f"{k}={v}" for k, v in best_params.items())
print(f" 最优参数: {params_str}")
# Step 2: 提取 IS 指标 (从优化结果)
is_metrics = self._extract_is_metrics(opt_result)
# Step 3: 在测试集上用最优参数回测
oos_metrics = self._run_with_params(
strategy_class, best_params, test_df, f"{symbol}_test"
)
if oos_metrics is None:
oos_metrics = {
"sharpe_ratio": np.nan, "total_return_pct": np.nan,
"max_drawdown_pct": np.nan, "total_trades": 0,
"win_rate_pct": np.nan, "profit_factor": np.nan,
"expectancy": np.nan,
}
if verbose:
is_sharpe = is_metrics.get("sharpe_ratio", np.nan)
oos_sharpe = oos_metrics.get("sharpe_ratio", np.nan)
print(f" IS 夏普: {is_sharpe:.4f} OOS 夏普: {oos_sharpe:.4f}")
windows.append(WalkForwardWindow(
window_idx, train_start, train_end, test_start, test_end,
best_params, is_metrics, oos_metrics,
))
start += self.step
window_idx += 1
return WalkForwardResult(windows, metric)
def _run_with_params(self, strategy_class, params: dict, df, symbol) -> dict | None:
"""用指定参数运行回测, 返回指标字典"""
try:
strategy = strategy_class(**params)
if strategy.warmup_bars() >= len(df):
return None
signals = strategy.generate_signals(df)
if int(signals.entries.sum()) == 0:
return None
config = strategy.build_config()
arr = Strategy.to_arrays(df)
result = raptorbt.run_single_backtest(
timestamps=arr["timestamps"],
open=arr["open"], high=arr["high"], low=arr["low"], close=arr["close"],
volume=arr["volume"],
entries=signals.entries, exits=signals.exits,
direction=signals.direction, weight=1.0, symbol=symbol,
config=config,
)
m = result.metrics
return {
"sharpe_ratio": m.sharpe_ratio,
"total_return_pct": m.total_return_pct,
"max_drawdown_pct": m.max_drawdown_pct,
"total_trades": m.total_trades,
"win_rate_pct": m.win_rate_pct,
"profit_factor": m.profit_factor,
"expectancy": getattr(m, "expectancy", np.nan),
}
except Exception as e:
print(f" ⚠️ OOS 回测失败: {e}")
return None
def _extract_is_metrics(self, opt_result) -> dict:
"""从优化结果中提取最优行的完整指标"""
if not opt_result.best_params:
return {opt_result.metric: np.nan}
results = opt_result.results
# 找到匹配最优参数的行
mask = pd.Series([True] * len(results))
for k, v in opt_result.best_params.items():
mask &= (results[k] == v)
matched = results[mask]
if len(matched) == 0:
return {opt_result.metric: opt_result.best_score}
row = matched.iloc[0]
return {
"sharpe_ratio": row.get("sharpe_ratio", np.nan),
"total_return_pct": row.get("total_return_pct", np.nan),
"max_drawdown_pct": row.get("max_drawdown_pct", np.nan),
"total_trades": row.get("total_trades", 0),
"win_rate_pct": row.get("win_rate_pct", np.nan),
"profit_factor": row.get("profit_factor", np.nan),
"expectancy": row.get("expectancy", np.nan),
}