optirs_core/privacy/private_hyperparameter_optimization/
functions.rs1use crate::error::Result;
6use crate::privacy::PrivacyBudget;
7use scirs2_core::numeric::Float;
8use std::fmt::Debug;
9
10use super::types::{
11 HPOEvaluation, HPOResult, ParameterConfiguration, ParameterSpace, StatisticalTestResult,
12};
13
14pub type ObjectiveFn<T> = Box<dyn Fn(&ParameterConfiguration<T>) -> Result<f64> + Send + Sync>;
15pub type RuleFn<T> = Box<dyn Fn(&HPOResult<T>) -> bool + Send + Sync>;
16pub type TestFn<T> = Box<dyn Fn(&[HPOResult<T>]) -> StatisticalTestResult + Send + Sync>;
17pub trait NoisyOptimizer<T: Float + Debug + Send + Sync + 'static>: Send + Sync {
19 fn suggest_next(
21 &mut self,
22 parameterspace: &ParameterSpace<T>,
23 evaluation_history: &[HPOEvaluation<T>],
24 _privacy_budget: &PrivacyBudget,
25 ) -> Result<ParameterConfiguration<T>>;
26 fn update(
28 &mut self,
29 config: &ParameterConfiguration<T>,
30 result: &HPOResult<T>,
31 _privacy_budget: &PrivacyBudget,
32 ) -> Result<()>;
33 fn name(&self) -> &str;
35}
36#[cfg(test)]
37mod tests {
38 use super::super::budget_manager::HPOBudgetManager;
39 use super::super::types::{
40 BudgetAllocationStrategy, EarlyStoppingConfig, HyperparameterNoiseMechanism,
41 ParameterBounds, ParameterDefinition, ParameterPrior, ParameterTransformation,
42 ParameterType, PrivateHPOConfig, PrivateRandomSearch, SearchAlgorithm, SensitivityBounds,
43 SmoothSensitivityParams, ValidationStrategy,
44 };
45 use super::*;
46 use crate::privacy::DifferentialPrivacyConfig;
47 use std::collections::HashMap;
48 #[test]
49 fn test_private_hpoconfig() {
50 let config = PrivateHPOConfig {
51 base_privacyconfig: DifferentialPrivacyConfig::default(),
52 budget_allocation: BudgetAllocationStrategy::Equal,
53 search_algorithm: SearchAlgorithm::RandomSearch,
54 num_evaluations: 100,
55 cv_folds: 5,
56 early_stopping: EarlyStoppingConfig {
57 enabled: true,
58 patience: 10,
59 min_improvement: 0.01,
60 max_evaluations: 100,
61 },
62 noise_mechanism: HyperparameterNoiseMechanism::Gaussian,
63 sensitivity_bounds: SensitivityBounds {
64 global_sensitivity: HashMap::<String, f64>::new(),
65 local_sensitivity: HashMap::<String, (f64, f64)>::new(),
66 smooth_sensitivity: HashMap::<String, SmoothSensitivityParams<f64>>::new(),
67 },
68 private_model_selection: true,
69 validation_strategy: ValidationStrategy::KFoldCV,
70 };
71 assert_eq!(config.num_evaluations, 100);
72 assert_eq!(config.cv_folds, 5);
73 assert!(config.early_stopping.enabled);
74 }
75 #[test]
76 fn test_parameter_space() {
77 let mut parameters = HashMap::new();
78 parameters.insert(
79 "learning_rate".to_string(),
80 ParameterDefinition {
81 name: "learning_rate".to_string(),
82 param_type: ParameterType::Continuous,
83 bounds: ParameterBounds {
84 min: Some(0.001),
85 max: Some(0.1),
86 step: None,
87 valid_values: None,
88 },
89 prior: Some(ParameterPrior::LogNormal(-3.0, 1.0)),
90 transformation: Some(ParameterTransformation::Log),
91 },
92 );
93 let parameterspace = ParameterSpace {
94 parameters,
95 constraints: Vec::new(),
96 defaultconfig: None,
97 };
98 assert!(parameterspace.parameters.contains_key("learning_rate"));
99 }
100 #[test]
101 fn test_budget_manager() {
102 let baseconfig = DifferentialPrivacyConfig::default();
103 let budget_manager = HPOBudgetManager::new(baseconfig, BudgetAllocationStrategy::Equal, 10)
104 .expect("unwrap failed");
105 assert!(budget_manager
106 .has_budget_remaining()
107 .expect("unwrap failed"));
108 }
109 #[test]
110 fn test_private_random_search() {
111 let config = PrivateHPOConfig {
112 base_privacyconfig: DifferentialPrivacyConfig::default(),
113 budget_allocation: BudgetAllocationStrategy::Equal,
114 search_algorithm: SearchAlgorithm::RandomSearch,
115 num_evaluations: 10,
116 cv_folds: 3,
117 early_stopping: EarlyStoppingConfig {
118 enabled: false,
119 patience: 5,
120 min_improvement: 0.01,
121 max_evaluations: 10,
122 },
123 noise_mechanism: HyperparameterNoiseMechanism::Gaussian,
124 sensitivity_bounds: SensitivityBounds {
125 global_sensitivity: HashMap::<String, f64>::new(),
126 local_sensitivity: HashMap::<String, (f64, f64)>::new(),
127 smooth_sensitivity: HashMap::<String, SmoothSensitivityParams<f64>>::new(),
128 },
129 private_model_selection: false,
130 validation_strategy: ValidationStrategy::HoldOut,
131 };
132 let optimizer = PrivateRandomSearch::new(config).expect("unwrap failed");
133 assert_eq!(optimizer.name(), "PrivateRandomSearch");
134 }
135}