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

1"""Overfitting validator module. 

2 

3Validates indicators to prevent overfitting. 

4""" 

5 

6from __future__ import annotations 

7 

8from dataclasses import dataclass, field 

9from typing import TYPE_CHECKING 

10 

11import numpy as np 

12import pandas as pd 

13 

14from monte_neo.utils.logger import get_logger 

15 

16if TYPE_CHECKING: 

17 from monte_neo.indicators.base import BaseIndicator 

18 from monte_neo.metrics.calculator import MetricsCalculator 

19 

20logger = get_logger(__name__) 

21 

22 

23@dataclass 

24class ValidationResult: 

25 """Validation result.""" 

26 

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) 

33 

34 

35class OverfitValidator: 

36 """Validate indicators for overfitting.""" 

37 

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. 

45 

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 

54 

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 = [] 

64 

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!") 

68 

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:] 

73 

74 # In-sample metrics 

75 is_signals = indicator.generate_signals(in_sample) 

76 is_metrics = metrics_calc.calculate_all(in_sample, is_signals) 

77 

78 # Out-of-sample metrics 

79 oos_signals = indicator.generate_signals(out_sample) 

80 oos_metrics = metrics_calc.calculate_all(out_sample, oos_signals) 

81 

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 ) 

88 

89 if oos_metrics.get("trade_count", 0) < self.min_trades // 3: 

90 warnings.append("Insufficient out-of-sample trades") 

91 

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

96 

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 

101 

102 if cv_std > 0.3: 

103 warnings.append(f"High CV variance ({cv_std:.2f})") 

104 

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

109 

110 # Overall score 

111 overall_score = self._calculate_overall_score( 

112 oos_ratio, avg_cv, cv_std, meets_targets 

113 ) 

114 

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 ) 

123 

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. 

131 

132 Args: 

133 indicator: Indicator to check. 

134 data: OHLCV data. 

135 lookback: Number of candles to check for repainting. 

136 

137 Returns: 

138 True if non-repainting. 

139 """ 

140 if len(data) < lookback + 10: 

141 return True 

142 

143 # 1. Generate signals for the full dataset 

144 full_signals = indicator.generate_signals(data) 

145 

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 

149 

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) 

154 

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 

162 

163 return True 

164 

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" 

172 

173 is_val = is_metrics.get(key_metric, 0) 

174 oos_val = oos_metrics.get(key_metric, 0) 

175 

176 if is_val <= 0: 

177 return 0.0 

178 

179 return oos_val / is_val 

180 

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 

190 

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 

195 

196 pd.concat([data.iloc[:test_start], data.iloc[test_end:]]) 

197 test = data.iloc[test_start:test_end] 

198 

199 if len(test) < 10: 

200 continue 

201 

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)) 

206 

207 return scores 

208 

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 

218 

219 actual = metrics[name] 

220 

221 if name in ["max_drawdown", "consecutive_losses"]: 

222 if actual > target: 

223 return False 

224 else: 

225 if actual < target: 

226 return False 

227 

228 return True 

229 

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 

239 

240 # OOS ratio contribution (0-0.4) 

241 score += min(0.4, oos_ratio * 0.4) 

242 

243 # CV mean contribution (0-0.3) 

244 score += min(0.3, cv_mean * 0.1) 

245 

246 # CV stability contribution (0-0.2) 

247 stability = max(0, 1 - cv_std) 

248 score += stability * 0.2 

249 

250 # Meets targets contribution (0.1) 

251 if meets_targets: 

252 score += 0.1 

253 

254 return min(1.0, score) 

255 

256 def quick_check( 

257 self, 

258 indicator: BaseIndicator, 

259 data: pd.DataFrame, 

260 metrics_calc: MetricsCalculator, 

261 ) -> bool: 

262 """Quick validation check. 

263 

264 Args: 

265 indicator: Indicator to check. 

266 data: OHLCV data. 

267 metrics_calc: Metrics calculator. 

268 

269 Returns: 

270 True if passes quick checks. 

271 """ 

272 signals = indicator.generate_signals(data) 

273 metrics = metrics_calc.calculate_all(data, signals) 

274 

275 # Check minimum trades 

276 if metrics.get("trade_count", 0) < self.min_trades: 

277 return False 

278 

279 # Check profit factor 

280 if metrics.get("profit_factor", 0) < 1.0: 

281 return False 

282 

283 # Check max drawdown 

284 if metrics.get("max_drawdown", 1) > 0.5: 

285 return False 

286 

287 return True