Coverage for src / monte_neo / monte_carlo / walk_forward.py: 87%
119 statements
« prev ^ index » next coverage.py v7.13.1, created at 2026-01-28 16:27 +0200
« prev ^ index » next coverage.py v7.13.1, created at 2026-01-28 16:27 +0200
1"""Walk-forward analysis module.
3Implements rolling walk-forward testing for robust out-of-sample validation.
4"""
6from __future__ import annotations
8from dataclasses import dataclass, field
9from typing import TYPE_CHECKING
11import numpy as np
12import pandas as pd
14from monte_neo.utils.logger import get_logger
16if TYPE_CHECKING:
17 from monte_neo.indicators.base import BaseIndicator
18 from monte_neo.metrics.calculator import MetricsCalculator
20logger = get_logger(__name__)
23@dataclass
24class WalkForwardWindow:
25 """Single walk-forward window result."""
27 window_idx: int
28 train_start: int
29 train_end: int
30 test_start: int
31 test_end: int
32 train_metrics: dict = field(default_factory=dict)
33 test_metrics: dict = field(default_factory=dict)
34 passed: bool = False
37@dataclass
38class WalkForwardResult:
39 """Complete walk-forward analysis result."""
41 windows: list[WalkForwardWindow] = field(default_factory=list)
42 overall_passed: bool = False
43 pass_rate: float = 0.0
44 avg_test_performance: dict = field(default_factory=dict)
45 efficiency_ratio: float = 0.0 # test_performance / train_performance
48class WalkForwardAnalyzer:
49 """Walk-forward analysis for out-of-sample validation."""
51 def __init__(
52 self,
53 n_splits: int = 5,
54 train_pct: float = 0.7,
55 anchored: bool = False,
56 ) -> None:
57 """Initialize walk-forward analyzer.
59 Args:
60 n_splits: Number of walk-forward windows.
61 train_pct: Training data percentage per window.
62 anchored: If True, training always starts from beginning.
63 """
64 self.n_splits = n_splits
65 self.train_pct = train_pct
66 self.anchored = anchored
68 def generate_scenarios(
69 self,
70 data: pd.DataFrame,
71 n_splits: int | None = None,
72 ) -> list[pd.DataFrame]:
73 """Generate walk-forward test scenarios (out-of-sample segments).
75 Args:
76 data: OHLCV data.
77 n_splits: Number of splits override.
79 Returns:
80 List of DataFrames representing out-of-sample periods.
81 """
82 if n_splits:
83 self.n_splits = n_splits
85 windows = self._generate_windows(len(data))
86 scenarios = []
88 for window in windows:
89 test_data = data.iloc[window.test_start : window.test_end].copy()
90 if not test_data.empty:
91 scenarios.append(test_data)
93 return scenarios
95 def analyze(
96 self,
97 indicator: BaseIndicator,
98 data: pd.DataFrame,
99 metrics_calc: MetricsCalculator,
100 target_metrics: dict[str, float],
101 ) -> WalkForwardResult:
102 """Run walk-forward analysis.
104 Args:
105 indicator: Indicator to test.
106 data: OHLCV data.
107 metrics_calc: Metrics calculator.
108 target_metrics: Target metrics for passing.
110 Returns:
111 WalkForwardResult with all windows.
112 """
113 windows = self._generate_windows(len(data))
114 results = []
116 for window in windows:
117 # Get train and test data
118 train_data = data.iloc[window.train_start : window.train_end]
119 test_data = data.iloc[window.test_start : window.test_end]
120 if train_data.empty or test_data.empty:
121 continue
123 # Optimize on training data (simplified - just calculate metrics)
124 train_signals = indicator.generate_signals(train_data)
125 train_metrics = metrics_calc.calculate_all(train_data, train_signals)
126 window.train_metrics = train_metrics
128 # Test on out-of-sample data
129 test_signals = indicator.generate_signals(test_data)
130 test_metrics = metrics_calc.calculate_all(test_data, test_signals)
131 window.test_metrics = test_metrics
133 # Check if test results meet targets
134 window.passed = self._check_targets(test_metrics, target_metrics)
135 results.append(window)
137 # Aggregate results
138 pass_rate = sum(1 for w in results if w.passed) / len(results) if results else 0
139 avg_test = self._aggregate_metrics([w.test_metrics for w in results])
140 efficiency = self._calculate_efficiency(results)
142 result = WalkForwardResult(
143 windows=results,
144 overall_passed=pass_rate >= 0.6, # 60% of windows must pass
145 pass_rate=pass_rate,
146 avg_test_performance=avg_test,
147 efficiency_ratio=efficiency,
148 )
150 logger.info(
151 f"Walk-forward: {pass_rate:.1%} pass rate, efficiency={efficiency:.2f}"
152 )
153 return result
155 def _generate_windows(self, n_samples: int) -> list[WalkForwardWindow]:
156 """Generate walk-forward windows.
158 Args:
159 n_samples: Total number of samples.
161 Returns:
162 List of WalkForwardWindow objects.
163 """
164 if n_samples <= 0 or self.n_splits <= 0:
165 return []
167 windows = []
168 window_size = n_samples // self.n_splits
169 if window_size <= 1:
170 return []
172 train_size = int(window_size * self.train_pct)
173 test_size = window_size - train_size
174 if train_size <= 0 or test_size <= 0:
175 return []
177 for i in range(self.n_splits):
178 if self.anchored:
179 train_start = 0
180 train_end = train_size + (i * test_size)
181 test_start = train_end
182 test_end = min(test_start + test_size, n_samples)
183 else:
184 train_start = i * test_size
185 train_end = train_start + train_size
186 test_start = train_end
187 test_end = min(test_start + test_size, n_samples)
189 if test_start < test_end:
190 windows.append(
191 WalkForwardWindow(
192 window_idx=i,
193 train_start=train_start,
194 train_end=train_end,
195 test_start=test_start,
196 test_end=test_end,
197 )
198 )
200 return windows
202 def _check_targets(
203 self,
204 metrics: dict[str, float],
205 targets: dict[str, float],
206 ) -> bool:
207 """Check if metrics meet targets.
209 Args:
210 metrics: Calculated metrics.
211 targets: Target values.
213 Returns:
214 True if all targets are met.
215 """
216 for metric_name, target_value in targets.items():
217 if metric_name not in metrics:
218 continue
220 actual = metrics[metric_name]
221 if not np.isfinite(actual):
222 return False
224 # Handle metrics that should be less than target
225 if metric_name in ["max_drawdown", "consecutive_losses"]:
226 if actual > target_value:
227 return False
228 else:
229 if actual < target_value:
230 return False
232 return True
234 def _aggregate_metrics(self, metrics_list: list[dict]) -> dict:
235 """Aggregate metrics across windows.
237 Args:
238 metrics_list: List of metrics dictionaries.
240 Returns:
241 Aggregated metrics.
242 """
243 if not metrics_list:
244 return {}
246 aggregated = {}
247 all_keys: set[str] = set()
248 for m in metrics_list:
249 all_keys.update(m.keys())
251 for key in all_keys:
252 raw_values = [m.get(key) for m in metrics_list if key in m]
253 values = [
254 v for v in raw_values
255 if isinstance(v, (int, float)) and np.isfinite(v)
256 ]
257 if values:
258 aggregated[key] = {
259 "mean": float(np.mean(values)),
260 "std": float(np.std(values)),
261 "min": float(np.min(values)),
262 "max": float(np.max(values)),
263 }
265 return aggregated
267 def _calculate_efficiency(
268 self,
269 windows: list[WalkForwardWindow],
270 ) -> float:
271 """Calculate walk-forward efficiency ratio.
273 Efficiency = out-of-sample performance / in-sample performance.
275 Args:
276 windows: List of walk-forward windows.
278 Returns:
279 Efficiency ratio (ideally close to 1.0).
280 """
281 if not windows:
282 return 0.0
284 train_pf = []
285 test_pf = []
287 for w in windows:
288 if "profit_factor" in w.train_metrics and "profit_factor" in w.test_metrics:
289 train_value = w.train_metrics["profit_factor"]
290 test_value = w.test_metrics["profit_factor"]
291 if np.isfinite(train_value) and np.isfinite(test_value):
292 train_pf.append(train_value)
293 test_pf.append(test_value)
295 if not train_pf or np.mean(train_pf) == 0:
296 return 0.0
298 return float(np.mean(test_pf) / np.mean(train_pf))