trustformers_debug/simulation_tools/
what_if_analysis.rs1use super::types::*;
8use chrono::{DateTime, Utc};
9use serde::{Deserialize, Serialize};
10use std::collections::HashMap;
11
12#[derive(Debug, Clone, Serialize, Deserialize)]
14pub struct WhatIfAnalysisResult {
15 pub timestamp: DateTime<Utc>,
17 pub base_scenario: Scenario,
19 pub scenarios: Vec<Scenario>,
21 pub impact_analysis: ScenarioImpactAnalysis,
23 pub sensitivity_analysis: FeatureSensitivityAnalysis,
25 pub counterfactual_insights: Vec<CounterfactualInsight>,
27 pub decision_boundary_exploration: DecisionBoundaryExploration,
29}
30
31#[derive(Debug, Clone, Serialize, Deserialize)]
33pub struct Scenario {
34 pub id: String,
36 pub description: String,
38 pub features: HashMap<String, f64>,
40 pub prediction: f64,
42 pub confidence: Option<f64>,
49 pub changed_features: Vec<FeatureChange>,
51 pub distance_from_base: f64,
53 pub plausibility: f64,
55}
56
57#[derive(Debug, Clone, Serialize, Deserialize)]
59pub struct FeatureChange {
60 pub feature_name: String,
62 pub original_value: f64,
64 pub new_value: f64,
66 pub change_magnitude: f64,
68 pub change_direction: ChangeDirection,
70 pub change_type: ChangeType,
72}
73
74#[derive(Debug, Clone, Serialize, Deserialize)]
76pub struct ScenarioImpactAnalysis {
77 pub high_impact_scenarios: Vec<String>,
79 pub prediction_flip_scenarios: Vec<String>,
81 pub avg_prediction_change: f64,
83 pub max_prediction_change: f64,
85 pub stability_analysis: PredictionStabilityAnalysis,
87 pub feature_importance_ranking: Vec<FeatureImportanceRank>,
89}
90
91#[derive(Debug, Clone, Serialize, Deserialize)]
93pub struct PredictionStabilityAnalysis {
94 pub stability_score: f64,
96 pub prediction_variance: f64,
98 pub prediction_flips: usize,
100 pub stability_by_magnitude: HashMap<String, f64>,
102}
103
104#[derive(Debug, Clone, Serialize, Deserialize)]
106pub struct FeatureImportanceRank {
107 pub feature_name: String,
109 pub importance_score: f64,
111 pub rank: usize,
113 pub avg_impact: f64,
115 pub change_frequency: usize,
117}
118
119#[derive(Debug, Clone, Serialize, Deserialize)]
121pub struct FeatureSensitivityAnalysis {
122 pub feature_sensitivities: HashMap<String, f64>,
124 pub most_sensitive_features: Vec<String>,
126 pub least_sensitive_features: Vec<String>,
128 pub non_linear_features: Vec<String>,
130 pub interaction_sensitivities: Vec<FeatureInteractionSensitivity>,
132}
133
134#[derive(Debug, Clone, Serialize, Deserialize)]
136pub struct FeatureInteractionSensitivity {
137 pub feature1: String,
139 pub feature2: String,
141 pub sensitivity_score: f64,
143 pub interaction_type: InteractionType,
145}
146
147#[derive(Debug, Clone, Serialize, Deserialize)]
149pub struct CounterfactualInsight {
150 pub description: String,
152 pub required_changes: Vec<FeatureChange>,
154 pub predicted_outcome: f64,
156 pub confidence: Option<f64>,
159 pub feasibility: ImplementationFeasibility,
161}
162
163#[derive(Debug, Clone, Serialize, Deserialize)]
165pub struct DecisionBoundaryExploration {
166 pub boundary_points: Vec<BoundaryPoint>,
168 pub boundary_complexity: BoundaryComplexity,
170 pub local_linearity: LocalLinearityAnalysis,
172 pub crossing_analysis: BoundaryCrossingAnalysis,
174}
175
176#[derive(Debug, Clone, Serialize, Deserialize)]
178pub struct BoundaryPoint {
179 pub coordinates: HashMap<String, f64>,
181 pub distance_to_boundary: f64,
183 pub prediction: f64,
185 pub gradient_direction: HashMap<String, f64>,
187}
188
189#[derive(Debug, Clone, Serialize, Deserialize)]
191pub struct BoundaryComplexity {
192 pub complexity_score: f64,
194 pub curvature: f64,
196 pub inflection_points: usize,
198 pub complexity_class: ComplexityClass,
200}
201
202#[derive(Debug, Clone, Serialize, Deserialize)]
204pub struct LocalLinearityAnalysis {
205 pub avg_linearity: f64,
207 pub linearity_by_region: HashMap<String, f64>,
209 pub most_linear_regions: Vec<String>,
211 pub most_nonlinear_regions: Vec<String>,
213}
214
215#[derive(Debug, Clone, Serialize, Deserialize)]
217pub struct BoundaryCrossingAnalysis {
218 pub crossing_count: usize,
220 pub avg_crossing_distance: f64,
222 pub crossing_directions: Vec<CrossingDirection>,
224 pub common_crossing_features: Vec<String>,
226}
227
228#[derive(Debug, Clone, Serialize, Deserialize)]
230pub struct CrossingDirection {
231 pub direction: HashMap<String, f64>,
233 pub magnitude: f64,
235 pub frequency: usize,
237}