Coverage for src / monte_neo / monte_carlo / sensitivity.py: 94%

69 statements  

« prev     ^ index     » next       coverage.py v7.13.1, created at 2026-01-28 16:27 +0200

1"""Sensitivity analysis module. 

2 

3Tests indicator stability across parameter variations. 

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

25 """Sensitivity analysis result.""" 

26 

27 parameter_name: str 

28 base_value: float 

29 variations: list[float] = field(default_factory=list) 

30 metrics_by_variation: dict = field(default_factory=dict) 

31 stability_score: float = 0.0 

32 is_stable: bool = False 

33 

34 

35class SensitivityAnalyzer: 

36 """Sensitivity analysis for parameter robustness.""" 

37 

38 def __init__( 

39 self, 

40 variation_range: float = 0.10, 

41 n_steps: int = 5, 

42 ) -> None: 

43 """Initialize sensitivity analyzer. 

44 

45 Args: 

46 variation_range: Variation range (e.g., 0.10 for ±10%). 

47 n_steps: Number of steps in each direction. 

48 """ 

49 self.variation_range = variation_range 

50 self.n_steps = n_steps 

51 

52 def analyze_parameter( 

53 self, 

54 indicator: BaseIndicator, 

55 param_name: str, 

56 base_value: float, 

57 data: pd.DataFrame, 

58 metrics_calc: MetricsCalculator, 

59 ) -> SensitivityResult: 

60 """Analyze sensitivity to a single parameter. 

61 

62 Args: 

63 indicator: Indicator to test. 

64 param_name: Parameter name to vary. 

65 base_value: Base parameter value. 

66 data: OHLCV data. 

67 metrics_calc: Metrics calculator. 

68 

69 Returns: 

70 SensitivityResult with analysis. 

71 """ 

72 # Generate variations 

73 variations = self._generate_variations(base_value) 

74 metrics_by_variation = {} 

75 

76 for value in variations: 

77 # Update indicator parameter 

78 indicator.set_parameter(param_name, value) 

79 

80 # Generate signals and calculate metrics 

81 signals = indicator.generate_signals(data) 

82 metrics = metrics_calc.calculate_all(data, signals) 

83 metrics_by_variation[value] = metrics 

84 

85 # Reset to base value 

86 indicator.set_parameter(param_name, base_value) 

87 

88 # Calculate stability score 

89 stability_score = self._calculate_stability(metrics_by_variation) 

90 is_stable = stability_score >= 0.7 # 70% stability threshold 

91 

92 result = SensitivityResult( 

93 parameter_name=param_name, 

94 base_value=base_value, 

95 variations=variations, 

96 metrics_by_variation=metrics_by_variation, 

97 stability_score=stability_score, 

98 is_stable=is_stable, 

99 ) 

100 

101 logger.info(f"Sensitivity {param_name}: stability={stability_score:.2f}") 

102 return result 

103 

104 def analyze_all_parameters( 

105 self, 

106 indicator: BaseIndicator, 

107 data: pd.DataFrame, 

108 metrics_calc: MetricsCalculator, 

109 ) -> list[SensitivityResult]: 

110 """Analyze sensitivity to all indicator parameters. 

111 

112 Args: 

113 indicator: Indicator to test. 

114 data: OHLCV data. 

115 metrics_calc: Metrics calculator. 

116 

117 Returns: 

118 List of SensitivityResults. 

119 """ 

120 results = [] 

121 

122 for param_name, param_value in indicator.get_parameters().items(): 

123 if isinstance(param_value, (int, float)): 

124 result = self.analyze_parameter( 

125 indicator, param_name, float(param_value), data, metrics_calc 

126 ) 

127 results.append(result) 

128 

129 return results 

130 

131 def get_stability_report( 

132 self, 

133 results: list[SensitivityResult], 

134 ) -> dict: 

135 """Generate stability report from results. 

136 

137 Args: 

138 results: List of sensitivity results. 

139 

140 Returns: 

141 Report dictionary. 

142 """ 

143 total_params = len(results) 

144 stable_params = sum(1 for r in results if r.is_stable) 

145 avg_stability = np.mean([r.stability_score for r in results]) if results else 0 

146 

147 return { 

148 "total_parameters": total_params, 

149 "stable_parameters": stable_params, 

150 "unstable_parameters": total_params - stable_params, 

151 "average_stability": float(avg_stability), 

152 "overall_stable": stable_params == total_params, 

153 "parameter_details": { 

154 r.parameter_name: { 

155 "base_value": r.base_value, 

156 "stability_score": r.stability_score, 

157 "is_stable": r.is_stable, 

158 } 

159 for r in results 

160 }, 

161 } 

162 

163 def _generate_variations(self, base_value: float) -> list[float]: 

164 """Generate parameter variations. 

165 

166 Args: 

167 base_value: Base parameter value. 

168 

169 Returns: 

170 List of variation values. 

171 """ 

172 variations = [base_value] # Include base 

173 

174 for i in range(1, self.n_steps + 1): 

175 factor = i * self.variation_range / self.n_steps 

176 variations.append(base_value * (1 - factor)) # Lower 

177 variations.append(base_value * (1 + factor)) # Higher 

178 

179 return sorted(set(variations)) 

180 

181 def _calculate_stability( 

182 self, 

183 metrics_by_variation: dict[float, dict], 

184 ) -> float: 

185 """Calculate stability score across variations. 

186 

187 Args: 

188 metrics_by_variation: Metrics for each parameter value. 

189 

190 Returns: 

191 Stability score 0-1. 

192 """ 

193 if len(metrics_by_variation) < 2: 

194 return 1.0 

195 

196 # Key metrics for stability assessment 

197 key_metrics = ["profit_factor", "sharpe_ratio", "max_drawdown"] 

198 

199 stability_scores = [] 

200 

201 for metric in key_metrics: 

202 values = [] 

203 for variation_metrics in metrics_by_variation.values(): 

204 if metric in variation_metrics: 

205 values.append(variation_metrics[metric]) 

206 

207 if len(values) >= 2: 

208 # Coefficient of variation (lower = more stable) 

209 mean_val = np.mean(values) 

210 if mean_val != 0: 

211 cv = np.std(values) / abs(mean_val) 

212 # Convert to stability score (1 = stable, 0 = unstable) 

213 stability = max(0, 1 - cv) 

214 stability_scores.append(stability) 

215 

216 return float(np.mean(stability_scores)) if stability_scores else 1.0