Skip to main content

trustformers_debug/simulation_tools/
what_if_analysis.rs

1//! What-If Analysis for Model Behavior Exploration
2//!
3//! This module provides comprehensive what-if analysis capabilities including
4//! scenario generation, feature sensitivity analysis, counterfactual insights,
5//! and decision boundary exploration.
6
7use super::types::*;
8use chrono::{DateTime, Utc};
9use serde::{Deserialize, Serialize};
10use std::collections::HashMap;
11
12/// What-if analysis result
13#[derive(Debug, Clone, Serialize, Deserialize)]
14pub struct WhatIfAnalysisResult {
15    /// Analysis timestamp
16    pub timestamp: DateTime<Utc>,
17    /// Base scenario (original input)
18    pub base_scenario: Scenario,
19    /// Generated what-if scenarios
20    pub scenarios: Vec<Scenario>,
21    /// Scenario impact analysis
22    pub impact_analysis: ScenarioImpactAnalysis,
23    /// Feature sensitivity analysis
24    pub sensitivity_analysis: FeatureSensitivityAnalysis,
25    /// Counterfactual insights
26    pub counterfactual_insights: Vec<CounterfactualInsight>,
27    /// Decision boundary exploration
28    pub decision_boundary_exploration: DecisionBoundaryExploration,
29}
30
31/// Individual scenario for what-if analysis
32#[derive(Debug, Clone, Serialize, Deserialize)]
33pub struct Scenario {
34    /// Scenario ID
35    pub id: String,
36    /// Scenario description
37    pub description: String,
38    /// Input features for this scenario
39    pub features: HashMap<String, f64>,
40    /// Model prediction for this scenario
41    pub prediction: f64,
42    /// Prediction confidence.
43    ///
44    /// `None` from [`super::analyzer::SimulationAnalyzer`]: the model function
45    /// it drives returns a bare scalar prediction with no uncertainty, and each
46    /// scenario is evaluated once. It used to be a flat `0.8` for every
47    /// scenario.
48    pub confidence: Option<f64>,
49    /// Changed features from base scenario
50    pub changed_features: Vec<FeatureChange>,
51    /// Distance from base scenario
52    pub distance_from_base: f64,
53    /// Scenario plausibility
54    pub plausibility: f64,
55}
56
57/// Change made to a feature in a scenario
58#[derive(Debug, Clone, Serialize, Deserialize)]
59pub struct FeatureChange {
60    /// Feature name
61    pub feature_name: String,
62    /// Original value
63    pub original_value: f64,
64    /// New value
65    pub new_value: f64,
66    /// Change magnitude
67    pub change_magnitude: f64,
68    /// Change direction
69    pub change_direction: ChangeDirection,
70    /// Change type
71    pub change_type: ChangeType,
72}
73
74/// Analysis of scenario impacts
75#[derive(Debug, Clone, Serialize, Deserialize)]
76pub struct ScenarioImpactAnalysis {
77    /// Most impactful scenarios
78    pub high_impact_scenarios: Vec<String>,
79    /// Scenarios with prediction flips
80    pub prediction_flip_scenarios: Vec<String>,
81    /// Average prediction change
82    pub avg_prediction_change: f64,
83    /// Maximum prediction change
84    pub max_prediction_change: f64,
85    /// Prediction stability analysis
86    pub stability_analysis: PredictionStabilityAnalysis,
87    /// Feature importance ranking from scenarios
88    pub feature_importance_ranking: Vec<FeatureImportanceRank>,
89}
90
91/// Analysis of prediction stability across scenarios
92#[derive(Debug, Clone, Serialize, Deserialize)]
93pub struct PredictionStabilityAnalysis {
94    /// Stability score (0-1, higher is more stable)
95    pub stability_score: f64,
96    /// Prediction variance across scenarios
97    pub prediction_variance: f64,
98    /// Number of prediction flips
99    pub prediction_flips: usize,
100    /// Stability by feature change magnitude
101    pub stability_by_magnitude: HashMap<String, f64>,
102}
103
104/// Feature importance rank from scenario analysis
105#[derive(Debug, Clone, Serialize, Deserialize)]
106pub struct FeatureImportanceRank {
107    /// Feature name
108    pub feature_name: String,
109    /// Importance score
110    pub importance_score: f64,
111    /// Rank (1 = most important)
112    pub rank: usize,
113    /// Average impact when changed
114    pub avg_impact: f64,
115    /// Number of scenarios where this feature was changed
116    pub change_frequency: usize,
117}
118
119/// Feature sensitivity analysis from what-if scenarios
120#[derive(Debug, Clone, Serialize, Deserialize)]
121pub struct FeatureSensitivityAnalysis {
122    /// Sensitivity scores for each feature
123    pub feature_sensitivities: HashMap<String, f64>,
124    /// Most sensitive features
125    pub most_sensitive_features: Vec<String>,
126    /// Least sensitive features
127    pub least_sensitive_features: Vec<String>,
128    /// Non-linear sensitivity detection
129    pub non_linear_features: Vec<String>,
130    /// Feature interaction sensitivities
131    pub interaction_sensitivities: Vec<FeatureInteractionSensitivity>,
132}
133
134/// Sensitivity of feature interactions
135#[derive(Debug, Clone, Serialize, Deserialize)]
136pub struct FeatureInteractionSensitivity {
137    /// First feature
138    pub feature1: String,
139    /// Second feature
140    pub feature2: String,
141    /// Interaction sensitivity score
142    pub sensitivity_score: f64,
143    /// Interaction type
144    pub interaction_type: InteractionType,
145}
146
147/// Counterfactual insight from what-if analysis
148#[derive(Debug, Clone, Serialize, Deserialize)]
149pub struct CounterfactualInsight {
150    /// Insight description
151    pub description: String,
152    /// Required feature changes for desired outcome
153    pub required_changes: Vec<FeatureChange>,
154    /// Predicted outcome
155    pub predicted_outcome: f64,
156    /// Confidence in the counterfactual, inherited from the scenario it came
157    /// from; `None` when the scenario carried none.
158    pub confidence: Option<f64>,
159    /// Implementation feasibility
160    pub feasibility: ImplementationFeasibility,
161}
162
163/// Decision boundary exploration from what-if analysis
164#[derive(Debug, Clone, Serialize, Deserialize)]
165pub struct DecisionBoundaryExploration {
166    /// Points near decision boundary
167    pub boundary_points: Vec<BoundaryPoint>,
168    /// Boundary complexity assessment
169    pub boundary_complexity: BoundaryComplexity,
170    /// Local linearity analysis
171    pub local_linearity: LocalLinearityAnalysis,
172    /// Boundary crossing analysis
173    pub crossing_analysis: BoundaryCrossingAnalysis,
174}
175
176/// Point near decision boundary
177#[derive(Debug, Clone, Serialize, Deserialize)]
178pub struct BoundaryPoint {
179    /// Point coordinates
180    pub coordinates: HashMap<String, f64>,
181    /// Distance to boundary
182    pub distance_to_boundary: f64,
183    /// Prediction at this point
184    pub prediction: f64,
185    /// Gradient direction
186    pub gradient_direction: HashMap<String, f64>,
187}
188
189/// Assessment of decision boundary complexity
190#[derive(Debug, Clone, Serialize, Deserialize)]
191pub struct BoundaryComplexity {
192    /// Complexity score
193    pub complexity_score: f64,
194    /// Boundary curvature
195    pub curvature: f64,
196    /// Number of inflection points
197    pub inflection_points: usize,
198    /// Complexity classification
199    pub complexity_class: ComplexityClass,
200}
201
202/// Local linearity analysis of decision boundary
203#[derive(Debug, Clone, Serialize, Deserialize)]
204pub struct LocalLinearityAnalysis {
205    /// Average local linearity score
206    pub avg_linearity: f64,
207    /// Linearity by region
208    pub linearity_by_region: HashMap<String, f64>,
209    /// Most linear regions
210    pub most_linear_regions: Vec<String>,
211    /// Most non-linear regions
212    pub most_nonlinear_regions: Vec<String>,
213}
214
215/// Analysis of decision boundary crossings
216#[derive(Debug, Clone, Serialize, Deserialize)]
217pub struct BoundaryCrossingAnalysis {
218    /// Number of boundary crossings found
219    pub crossing_count: usize,
220    /// Average crossing distance
221    pub avg_crossing_distance: f64,
222    /// Crossing directions
223    pub crossing_directions: Vec<CrossingDirection>,
224    /// Most common crossing features
225    pub common_crossing_features: Vec<String>,
226}
227
228/// Direction of boundary crossing
229#[derive(Debug, Clone, Serialize, Deserialize)]
230pub struct CrossingDirection {
231    /// Direction vector
232    pub direction: HashMap<String, f64>,
233    /// Direction magnitude
234    pub magnitude: f64,
235    /// Frequency of this direction
236    pub frequency: usize,
237}