""" 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), }