1use crate::{OptimizerError, OptimizerResult};
11use serde::{Deserialize, Serialize};
12use std::collections::HashMap;
13use std::ops::Add;
14use torsh_tensor::Tensor;
15
16use scirs2_core::random::prelude::*;
18
19#[derive(Debug, Clone, Serialize, Deserialize)]
21pub struct DPConfig {
22 pub target_epsilon: f64,
24 pub target_delta: f64,
26 pub noise_multiplier: f64,
28 pub l2_norm_clip: f32,
30 pub num_training_steps: usize,
32 pub sampling_rate: f64,
34 pub adaptive_clipping: bool,
36 pub clipping_quantile: f64,
38}
39
40impl DPConfig {
41 pub fn new(target_epsilon: f64, target_delta: f64, num_training_steps: usize) -> Self {
42 Self {
43 target_epsilon,
44 target_delta,
45 noise_multiplier: 1.0,
46 l2_norm_clip: 1.0,
47 num_training_steps,
48 sampling_rate: 0.01,
49 adaptive_clipping: false,
50 clipping_quantile: 0.5,
51 }
52 }
53
54 pub fn calculate_noise_multiplier(&self) -> f64 {
56 let steps = self.num_training_steps as f64;
59 let q = self.sampling_rate;
60
61 let sigma = (q * steps * (self.target_epsilon.exp() - 1.0) / self.target_delta).sqrt();
63 sigma.max(0.1) }
65
66 pub fn privacy_spent(&self, step: usize) -> (f64, f64) {
68 let steps_ratio = step as f64 / self.num_training_steps as f64;
69 let epsilon_spent = self.target_epsilon * steps_ratio;
70 let delta_spent = self.target_delta * steps_ratio;
71 (epsilon_spent, delta_spent)
72 }
73}
74
75#[derive(Debug, Clone)]
77pub struct DPState {
78 pub current_step: usize,
80 pub epsilon_spent: f64,
82 pub delta_spent: f64,
83 pub gradient_norms: Vec<f32>,
85 pub clipping_stats: ClippingStats,
87}
88
89impl Default for DPState {
90 fn default() -> Self {
91 Self {
92 current_step: 0,
93 epsilon_spent: 0.0,
94 delta_spent: 0.0,
95 gradient_norms: Vec::new(),
96 clipping_stats: ClippingStats::default(),
97 }
98 }
99}
100
101impl DPState {
102 pub fn new() -> Self {
103 Self::default()
104 }
105
106 pub fn update_step(&mut self, config: &DPConfig) {
107 self.current_step += 1;
108 let (eps, delta) = config.privacy_spent(self.current_step);
109 self.epsilon_spent = eps;
110 self.delta_spent = delta;
111 }
112}
113
114#[derive(Debug, Clone)]
116pub struct ClippingStats {
117 pub total_gradients: usize,
118 pub clipped_gradients: usize,
119 pub average_norm_before: f32,
120 pub average_norm_after: f32,
121 pub max_norm_observed: f32,
122}
123
124impl Default for ClippingStats {
125 fn default() -> Self {
126 Self {
127 total_gradients: 0,
128 clipped_gradients: 0,
129 average_norm_before: 0.0,
130 average_norm_after: 0.0,
131 max_norm_observed: 0.0,
132 }
133 }
134}
135
136impl ClippingStats {
137 pub fn new() -> Self {
138 Self::default()
139 }
140
141 pub fn clipping_rate(&self) -> f32 {
142 if self.total_gradients == 0 {
143 0.0
144 } else {
145 self.clipped_gradients as f32 / self.total_gradients as f32
146 }
147 }
148}
149
150pub struct DPManager {
152 config: DPConfig,
153 state: DPState,
154 rng: Random,
155}
156
157impl DPManager {
158 pub fn new(config: DPConfig) -> Self {
159 Self {
160 config,
161 state: DPState::new(),
162 rng: Random::default(),
163 }
164 }
165
166 pub fn privatize_gradients(
168 &mut self,
169 gradients: &HashMap<String, Tensor>,
170 ) -> OptimizerResult<HashMap<String, Tensor>> {
171 let mut private_gradients = HashMap::new();
172
173 for (param_name, gradient) in gradients {
174 let private_grad = self.apply_dp_to_gradient(gradient)?;
175 private_gradients.insert(param_name.clone(), private_grad);
176 }
177
178 self.state.update_step(&self.config);
179 Ok(private_gradients)
180 }
181
182 fn apply_dp_to_gradient(&mut self, gradient: &Tensor) -> OptimizerResult<Tensor> {
184 let clipped_grad = self.clip_gradient(gradient)?;
186
187 let noisy_grad = self.add_noise(&clipped_grad)?;
189
190 Ok(noisy_grad)
191 }
192
193 fn clip_gradient(&mut self, gradient: &Tensor) -> OptimizerResult<Tensor> {
195 let grad_norm = gradient.norm()?.item()?;
196 self.state.clipping_stats.total_gradients += 1;
197 self.state.clipping_stats.max_norm_observed =
198 self.state.clipping_stats.max_norm_observed.max(grad_norm);
199
200 let clip_bound = if self.config.adaptive_clipping {
201 self.get_adaptive_clip_bound()
202 } else {
203 self.config.l2_norm_clip
204 };
205
206 self.state.clipping_stats.average_norm_before =
207 (self.state.clipping_stats.average_norm_before
208 * (self.state.clipping_stats.total_gradients - 1) as f32
209 + grad_norm)
210 / self.state.clipping_stats.total_gradients as f32;
211
212 let clipped_grad = if grad_norm > clip_bound {
213 self.state.clipping_stats.clipped_gradients += 1;
214 gradient.mul_scalar(clip_bound / grad_norm)?
215 } else {
216 gradient.clone()
217 };
218
219 let clipped_norm = if grad_norm > clip_bound {
220 clip_bound
221 } else {
222 grad_norm
223 };
224 self.state.clipping_stats.average_norm_after =
225 (self.state.clipping_stats.average_norm_after
226 * (self.state.clipping_stats.total_gradients - 1) as f32
227 + clipped_norm)
228 / self.state.clipping_stats.total_gradients as f32;
229
230 Ok(clipped_grad)
231 }
232
233 fn add_noise(&mut self, gradient: &Tensor) -> OptimizerResult<Tensor> {
235 let noise_scale = self.config.noise_multiplier * self.config.l2_norm_clip as f64;
236 let noise_tensor =
237 self.generate_gaussian_noise(gradient.shape().dims(), noise_scale as f32)?;
238
239 Ok(gradient.add(&noise_tensor)?)
240 }
241
242 fn generate_gaussian_noise(
244 &mut self,
245 shape: &[usize],
246 std_dev: f32,
247 ) -> OptimizerResult<Tensor> {
248 let total_elements: usize = shape.iter().product();
250 let mut rng = thread_rng();
251 let normal = Normal::new(0.0, std_dev).map_err(|e| {
252 OptimizerError::InvalidParameter(format!("Invalid noise parameters: {e}"))
253 })?;
254 let noise_data: Vec<f32> = (0..total_elements).map(|_| rng.sample(normal)).collect();
255
256 Tensor::from_data(
257 noise_data,
258 shape.to_vec(),
259 torsh_core::device::DeviceType::Cpu,
260 )
261 .map_err(OptimizerError::TensorError)
262 }
263
264 fn get_adaptive_clip_bound(&self) -> f32 {
266 if self.state.gradient_norms.is_empty() {
267 return self.config.l2_norm_clip;
268 }
269
270 let mut sorted_norms = self.state.gradient_norms.clone();
271 sorted_norms.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
272
273 let quantile_idx = (sorted_norms.len() as f64 * self.config.clipping_quantile) as usize;
274 sorted_norms[quantile_idx.min(sorted_norms.len() - 1)]
275 }
276
277 pub fn is_privacy_exhausted(&self) -> bool {
279 self.state.epsilon_spent >= self.config.target_epsilon
280 || self.state.delta_spent >= self.config.target_delta
281 }
282
283 pub fn remaining_budget(&self) -> (f64, f64) {
285 let remaining_epsilon = (self.config.target_epsilon - self.state.epsilon_spent).max(0.0);
286 let remaining_delta = (self.config.target_delta - self.state.delta_spent).max(0.0);
287 (remaining_epsilon, remaining_delta)
288 }
289
290 pub fn get_privacy_summary(&self) -> PrivacySummary {
292 PrivacySummary {
293 epsilon_spent: self.state.epsilon_spent,
294 delta_spent: self.state.delta_spent,
295 epsilon_remaining: self.config.target_epsilon - self.state.epsilon_spent,
296 delta_remaining: self.config.target_delta - self.state.delta_spent,
297 steps_completed: self.state.current_step,
298 clipping_rate: self.state.clipping_stats.clipping_rate(),
299 noise_multiplier: self.config.noise_multiplier,
300 }
301 }
302
303 pub fn update_gradient_norms(&mut self, norms: &[f32]) {
305 self.state.gradient_norms.extend_from_slice(norms);
306
307 const MAX_HISTORY: usize = 1000;
309 if self.state.gradient_norms.len() > MAX_HISTORY {
310 let start = self.state.gradient_norms.len() - MAX_HISTORY;
311 self.state.gradient_norms = self.state.gradient_norms[start..].to_vec();
312 }
313 }
314}
315
316#[derive(Debug, Clone, Serialize, Deserialize)]
318pub struct PrivacySummary {
319 pub epsilon_spent: f64,
320 pub delta_spent: f64,
321 pub epsilon_remaining: f64,
322 pub delta_remaining: f64,
323 pub steps_completed: usize,
324 pub clipping_rate: f32,
325 pub noise_multiplier: f64,
326}
327
328pub trait DifferentiallyPrivateOptimizer {
330 fn dp_step(&mut self, dp_manager: &mut DPManager) -> OptimizerResult<()>;
332
333 fn privacy_status(&self, dp_manager: &DPManager) -> PrivacySummary;
335
336 fn should_stop_for_privacy(&self, dp_manager: &DPManager) -> bool {
338 dp_manager.is_privacy_exhausted()
339 }
340}
341
342#[cfg(test)]
343mod tests {
344 use super::*;
345 use torsh_tensor::creation::randn;
346
347 #[test]
348 fn test_dp_config_creation() {
349 let config = DPConfig::new(1.0, 1e-5, 1000);
350 assert_eq!(config.target_epsilon, 1.0);
351 assert_eq!(config.target_delta, 1e-5);
352 assert_eq!(config.num_training_steps, 1000);
353 }
354
355 #[test]
356 fn test_noise_multiplier_calculation() {
357 let config = DPConfig::new(1.0, 1e-5, 1000);
358 let noise_multiplier = config.calculate_noise_multiplier();
359 assert!(noise_multiplier > 0.0);
360 }
361
362 #[test]
363 fn test_privacy_budget_tracking() {
364 let config = DPConfig::new(1.0, 1e-5, 1000);
365 let (eps, delta) = config.privacy_spent(500);
366 assert!((eps - 0.5).abs() < 1e-6);
367 assert!((delta - 5e-6).abs() < 1e-9);
368 }
369
370 #[test]
371 fn test_dp_manager_creation() {
372 let config = DPConfig::new(1.0, 1e-5, 1000);
373 let manager = DPManager::new(config);
374 assert_eq!(manager.state.current_step, 0);
375 assert_eq!(manager.state.epsilon_spent, 0.0);
376 }
377
378 #[test]
379 fn test_gradient_clipping() -> OptimizerResult<()> {
380 let config = DPConfig::new(1.0, 1e-5, 1000);
381 let mut manager = DPManager::new(config);
382
383 let gradient = randn::<f32>(&[2, 2])?.mul_scalar(10.0)?; let clipped = manager.clip_gradient(&gradient)?;
385
386 let clipped_norm = clipped.norm()?.item()?;
387 assert!(clipped_norm <= manager.config.l2_norm_clip + 1e-6);
388 Ok(())
389 }
390
391 #[test]
392 fn test_noise_generation() {
393 let config = DPConfig::new(1.0, 1e-5, 1000);
394 let mut manager = DPManager::new(config);
395
396 let noise = manager.generate_gaussian_noise(&[2, 2], 1.0).unwrap();
397 assert_eq!(noise.shape().dims(), &[2, 2]);
398 }
399
400 #[test]
401 fn test_privacy_exhaustion() {
402 let config = DPConfig::new(1.0, 1e-5, 10);
403 let mut state = DPState::new();
404
405 for _ in 0..20 {
407 state.update_step(&config);
408 }
409
410 let manager = DPManager {
411 config,
412 state,
413 rng: Random::default(),
414 };
415 assert!(manager.is_privacy_exhausted());
416 }
417
418 #[test]
419 fn test_clipping_statistics() {
420 let mut stats = ClippingStats::new();
421 assert_eq!(stats.clipping_rate(), 0.0);
422
423 stats.total_gradients = 10;
424 stats.clipped_gradients = 3;
425 assert!((stats.clipping_rate() - 0.3).abs() < 1e-6);
426 }
427}