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
« prev ^ index » next coverage.py v7.13.1, created at 2026-01-28 16:27 +0200
1"""Sensitivity analysis module.
3Tests indicator stability across parameter variations.
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 SensitivityResult:
25 """Sensitivity analysis result."""
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
35class SensitivityAnalyzer:
36 """Sensitivity analysis for parameter robustness."""
38 def __init__(
39 self,
40 variation_range: float = 0.10,
41 n_steps: int = 5,
42 ) -> None:
43 """Initialize sensitivity analyzer.
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
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.
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.
69 Returns:
70 SensitivityResult with analysis.
71 """
72 # Generate variations
73 variations = self._generate_variations(base_value)
74 metrics_by_variation = {}
76 for value in variations:
77 # Update indicator parameter
78 indicator.set_parameter(param_name, value)
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
85 # Reset to base value
86 indicator.set_parameter(param_name, base_value)
88 # Calculate stability score
89 stability_score = self._calculate_stability(metrics_by_variation)
90 is_stable = stability_score >= 0.7 # 70% stability threshold
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 )
101 logger.info(f"Sensitivity {param_name}: stability={stability_score:.2f}")
102 return result
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.
112 Args:
113 indicator: Indicator to test.
114 data: OHLCV data.
115 metrics_calc: Metrics calculator.
117 Returns:
118 List of SensitivityResults.
119 """
120 results = []
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)
129 return results
131 def get_stability_report(
132 self,
133 results: list[SensitivityResult],
134 ) -> dict:
135 """Generate stability report from results.
137 Args:
138 results: List of sensitivity results.
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
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 }
163 def _generate_variations(self, base_value: float) -> list[float]:
164 """Generate parameter variations.
166 Args:
167 base_value: Base parameter value.
169 Returns:
170 List of variation values.
171 """
172 variations = [base_value] # Include base
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
179 return sorted(set(variations))
181 def _calculate_stability(
182 self,
183 metrics_by_variation: dict[float, dict],
184 ) -> float:
185 """Calculate stability score across variations.
187 Args:
188 metrics_by_variation: Metrics for each parameter value.
190 Returns:
191 Stability score 0-1.
192 """
193 if len(metrics_by_variation) < 2:
194 return 1.0
196 # Key metrics for stability assessment
197 key_metrics = ["profit_factor", "sharpe_ratio", "max_drawdown"]
199 stability_scores = []
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])
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)
216 return float(np.mean(stability_scores)) if stability_scores else 1.0