Coverage for src / monte_neo / core / validator.py: 87%
109 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"""Overfitting validator module.
3Validates indicators to prevent overfitting.
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 ValidationResult:
25 """Validation result."""
27 passed: bool
28 overall_score: float
29 in_sample_metrics: dict = field(default_factory=dict)
30 out_sample_metrics: dict = field(default_factory=dict)
31 cross_val_scores: list = field(default_factory=list)
32 warnings: list = field(default_factory=list)
35class OverfitValidator:
36 """Validate indicators for overfitting."""
38 def __init__(
39 self,
40 min_trades: int = 30,
41 min_oos_ratio: float = 0.6,
42 n_folds: int = 5,
43 ) -> None:
44 """Initialize validator.
46 Args:
47 min_trades: Minimum trades required.
48 min_oos_ratio: Minimum OOS/IS performance ratio.
49 n_folds: Number of cross-validation folds.
50 """
51 self.min_trades = min_trades
52 self.min_oos_ratio = min_oos_ratio
53 self.n_folds = n_folds
55 def validate(
56 self,
57 indicator: BaseIndicator,
58 data: pd.DataFrame,
59 metrics_calc: MetricsCalculator,
60 target_metrics: dict[str, float],
61 ) -> ValidationResult:
62 """Full validation check."""
63 warnings = []
65 # 1. Non-repainting check
66 if not self.check_non_repainting(indicator, data):
67 warnings.append("INDICATOR REPAINTS: Signals change when new data arrives!")
69 # In-sample / Out-of-sample split
70 split_idx = int(len(data) * 0.7)
71 in_sample = data.iloc[:split_idx]
72 out_sample = data.iloc[split_idx:]
74 # In-sample metrics
75 is_signals = indicator.generate_signals(in_sample)
76 is_metrics = metrics_calc.calculate_all(in_sample, is_signals)
78 # Out-of-sample metrics
79 oos_signals = indicator.generate_signals(out_sample)
80 oos_metrics = metrics_calc.calculate_all(out_sample, oos_signals)
82 # Check minimum trades
83 if is_metrics.get("trade_count", 0) < self.min_trades:
84 warnings.append(
85 f"In-sample trades "
86 f"({is_metrics.get('trade_count', 0)}) < {self.min_trades}"
87 )
89 if oos_metrics.get("trade_count", 0) < self.min_trades // 3:
90 warnings.append("Insufficient out-of-sample trades")
92 # Check OOS/IS ratio
93 oos_ratio = self._calculate_oos_ratio(is_metrics, oos_metrics)
94 if oos_ratio < self.min_oos_ratio:
95 warnings.append(f"OOS/IS ratio ({oos_ratio:.2f}) < {self.min_oos_ratio}")
97 # Cross-validation
98 cv_scores = self._cross_validate(indicator, data, metrics_calc)
99 avg_cv = np.mean(cv_scores) if cv_scores else 0
100 cv_std = np.std(cv_scores) if cv_scores else 0
102 if cv_std > 0.3:
103 warnings.append(f"High CV variance ({cv_std:.2f})")
105 # Check target metrics on OOS
106 meets_targets = self._check_targets(oos_metrics, target_metrics)
107 if not meets_targets:
108 warnings.append("OOS metrics don't meet targets")
110 # Overall score
111 overall_score = self._calculate_overall_score(
112 oos_ratio, avg_cv, cv_std, meets_targets
113 )
115 return ValidationResult(
116 passed=len(warnings) == 0 and overall_score >= 0.7,
117 overall_score=overall_score,
118 in_sample_metrics=is_metrics,
119 out_sample_metrics=oos_metrics,
120 cross_val_scores=cv_scores,
121 warnings=warnings,
122 )
124 def check_non_repainting(
125 self,
126 indicator: BaseIndicator,
127 data: pd.DataFrame,
128 lookback: int = 50,
129 ) -> bool:
130 """Check if indicator repaints by simulating real-time data arrival.
132 Args:
133 indicator: Indicator to check.
134 data: OHLCV data.
135 lookback: Number of candles to check for repainting.
137 Returns:
138 True if non-repainting.
139 """
140 if len(data) < lookback + 10:
141 return True
143 # 1. Generate signals for the full dataset
144 full_signals = indicator.generate_signals(data)
146 # 2. Simulate incremental data arrival and check if previous signals change
147 # We check the last 'lookback' points
148 test_start = len(data) - lookback
150 for i in range(test_start, len(data)):
151 # Partial data up to index i
152 partial_data = data.iloc[:i+1]
153 partial_signals = indicator.generate_signals(partial_data)
155 # Check if the signal at index i is the same as in the full dataset
156 # (Signal at current candle is allowed to change until candle closes,
157 # but we are checking closed candles here)
158 if not np.array_equal(full_signals[:i+1], partial_signals):
159 # Repainting detected!
160 logger.warning(f"Repainting detected at index {i}")
161 return False
163 return True
165 def _calculate_oos_ratio(
166 self,
167 is_metrics: dict,
168 oos_metrics: dict,
169 ) -> float:
170 """Calculate out-of-sample to in-sample ratio."""
171 key_metric = "sharpe_ratio"
173 is_val = is_metrics.get(key_metric, 0)
174 oos_val = oos_metrics.get(key_metric, 0)
176 if is_val <= 0:
177 return 0.0
179 return oos_val / is_val
181 def _cross_validate(
182 self,
183 indicator: BaseIndicator,
184 data: pd.DataFrame,
185 metrics_calc: MetricsCalculator,
186 ) -> list[float]:
187 """Perform k-fold cross-validation."""
188 scores = []
189 fold_size = len(data) // self.n_folds
191 for i in range(self.n_folds):
192 # Create test fold
193 test_start = i * fold_size
194 test_end = (i + 1) * fold_size
196 pd.concat([data.iloc[:test_start], data.iloc[test_end:]])
197 test = data.iloc[test_start:test_end]
199 if len(test) < 10:
200 continue
202 # Test on fold
203 signals = indicator.generate_signals(test)
204 metrics = metrics_calc.calculate_all(test, signals)
205 scores.append(metrics.get("sharpe_ratio", 0))
207 return scores
209 def _check_targets(
210 self,
211 metrics: dict,
212 targets: dict,
213 ) -> bool:
214 """Check if metrics meet targets."""
215 for name, target in targets.items():
216 if name not in metrics:
217 continue
219 actual = metrics[name]
221 if name in ["max_drawdown", "consecutive_losses"]:
222 if actual > target:
223 return False
224 else:
225 if actual < target:
226 return False
228 return True
230 def _calculate_overall_score(
231 self,
232 oos_ratio: float,
233 cv_mean: float,
234 cv_std: float,
235 meets_targets: bool,
236 ) -> float:
237 """Calculate overall validation score."""
238 score = 0.0
240 # OOS ratio contribution (0-0.4)
241 score += min(0.4, oos_ratio * 0.4)
243 # CV mean contribution (0-0.3)
244 score += min(0.3, cv_mean * 0.1)
246 # CV stability contribution (0-0.2)
247 stability = max(0, 1 - cv_std)
248 score += stability * 0.2
250 # Meets targets contribution (0.1)
251 if meets_targets:
252 score += 0.1
254 return min(1.0, score)
256 def quick_check(
257 self,
258 indicator: BaseIndicator,
259 data: pd.DataFrame,
260 metrics_calc: MetricsCalculator,
261 ) -> bool:
262 """Quick validation check.
264 Args:
265 indicator: Indicator to check.
266 data: OHLCV data.
267 metrics_calc: Metrics calculator.
269 Returns:
270 True if passes quick checks.
271 """
272 signals = indicator.generate_signals(data)
273 metrics = metrics_calc.calculate_all(data, signals)
275 # Check minimum trades
276 if metrics.get("trade_count", 0) < self.min_trades:
277 return False
279 # Check profit factor
280 if metrics.get("profit_factor", 0) < 1.0:
281 return False
283 # Check max drawdown
284 if metrics.get("max_drawdown", 1) > 0.5:
285 return False
287 return True