1use scirs2_core::ndarray::{Array1, Array2, Axis};
9
10use crate::{
11 lda::LinearDiscriminantAnalysisConfig, numerical_stability::NumericalStability,
12 qda::QuadraticDiscriminantAnalysisConfig,
13};
14
15use sklears_core::{error::Result, prelude::SklearsError, types::Float};
16
17use std::collections::HashMap;
18use std::sync::{Arc, RwLock};
19use std::time::{Duration, Instant};
20use tokio::sync::{mpsc, Semaphore};
21use tokio::task::JoinHandle;
22
23#[derive(Debug, Clone)]
25pub struct AsyncOptimizationConfig {
26 pub num_workers: usize,
28 pub batch_size: usize,
30 pub max_concurrent_tasks: usize,
32 pub learning_rate: Float,
34 pub learning_rate_decay: Float,
36 pub momentum: Float,
38 pub adaptive_learning_rate: bool,
40 pub convergence_tolerance: Float,
42 pub max_iterations: usize,
44 pub async_updates: bool,
46 pub update_frequency_ms: u64,
48 pub gradient_checkpointing: bool,
50 pub channel_buffer_size: usize,
52}
53
54impl Default for AsyncOptimizationConfig {
55 fn default() -> Self {
56 Self {
57 num_workers: num_cpus::get(),
58 batch_size: 128,
59 max_concurrent_tasks: num_cpus::get() * 2,
60 learning_rate: 0.01,
61 learning_rate_decay: 0.99,
62 momentum: 0.9,
63 adaptive_learning_rate: true,
64 convergence_tolerance: 1e-6,
65 max_iterations: 1000,
66 async_updates: true,
67 update_frequency_ms: 100,
68 gradient_checkpointing: true,
69 channel_buffer_size: 1000,
70 }
71 }
72}
73
74#[derive(Debug, Clone)]
76pub struct OptimizationState {
77 pub iteration: usize,
79 pub current_loss: Float,
81 pub previous_loss: Float,
83 pub learning_rate_history: Vec<Float>,
85 pub loss_history: Vec<Float>,
87 pub gradient_norms: Vec<Float>,
89 pub converged: bool,
91 pub start_time: Instant,
93 pub last_update: Instant,
95 pub update_count: usize,
97}
98
99impl Default for OptimizationState {
100 fn default() -> Self {
101 let now = Instant::now();
102 Self {
103 iteration: 0,
104 current_loss: Float::INFINITY,
105 previous_loss: Float::INFINITY,
106 learning_rate_history: Vec::new(),
107 loss_history: Vec::new(),
108 gradient_norms: Vec::new(),
109 converged: false,
110 start_time: now,
111 last_update: now,
112 update_count: 0,
113 }
114 }
115}
116
117#[derive(Debug, Clone)]
119pub enum OptimizationMessage {
120 ParameterUpdate {
122 worker_id: usize,
123
124 parameters: HashMap<String, Array2<Float>>,
125
126 gradient: HashMap<String, Array2<Float>>,
127
128 loss: Float,
129 },
130 ComputeGradient {
132 worker_id: usize,
133 data_batch: Array2<Float>,
134 labels_batch: Array1<usize>,
135 },
136 CheckConvergence {
138 worker_id: usize,
139 current_loss: Float,
140 },
141 TrainingComplete {
143 final_parameters: HashMap<String, Array2<Float>>,
144 final_loss: Float,
145 iterations: usize,
146 },
147 OptimizationError { worker_id: usize, error: String },
149}
150
151pub struct AsyncDiscriminantOptimizer {
153 config: AsyncOptimizationConfig,
154 state: Arc<RwLock<OptimizationState>>,
155 parameters: Arc<RwLock<HashMap<String, Array2<Float>>>>,
156 velocity: Arc<RwLock<HashMap<String, Array2<Float>>>>, semaphore: Arc<Semaphore>,
158 #[allow(dead_code)] numerical_stability: NumericalStability,
161}
162
163impl AsyncDiscriminantOptimizer {
164 pub fn new(config: AsyncOptimizationConfig) -> Self {
166 let semaphore = Arc::new(Semaphore::new(config.max_concurrent_tasks));
167
168 Self {
169 config,
170 state: Arc::new(RwLock::new(OptimizationState::default())),
171 parameters: Arc::new(RwLock::new(HashMap::new())),
172 velocity: Arc::new(RwLock::new(HashMap::new())),
173 semaphore,
174 numerical_stability: NumericalStability::new(),
175 }
176 }
177
178 pub fn initialize_parameters(&self, n_features: usize, n_classes: usize) -> Result<()> {
180 let mut parameters = self.parameters.write().expect("lock not poisoned");
181 let mut velocity = self.velocity.write().expect("lock not poisoned");
182
183 let class_means = Array2::zeros((n_classes, n_features));
185 parameters.insert("class_means".to_string(), class_means.clone());
186 velocity.insert(
187 "class_means".to_string(),
188 Array2::zeros((n_classes, n_features)),
189 );
190
191 let covariance = Array2::eye(n_features);
193 parameters.insert("covariance".to_string(), covariance.clone());
194 velocity.insert(
195 "covariance".to_string(),
196 Array2::zeros((n_features, n_features)),
197 );
198
199 let priors = Array2::from_elem((1, n_classes), 1.0 / n_classes as Float);
201 parameters.insert("priors".to_string(), priors.clone());
202 velocity.insert("priors".to_string(), Array2::zeros((1, n_classes)));
203
204 Ok(())
205 }
206
207 pub async fn async_train_lda(
209 &self,
210 x_data: Array2<Float>, y: Array1<usize>,
212 config: LinearDiscriminantAnalysisConfig,
213 ) -> Result<HashMap<String, Array2<Float>>> {
214 let (_n_samples, n_features) = x_data.dim();
215 let n_classes = y.iter().max().expect("collection should not be empty") + 1;
216
217 self.initialize_parameters(n_features, n_classes)?;
219
220 let (tx, rx) = mpsc::channel::<OptimizationMessage>(self.config.channel_buffer_size);
222
223 let x_shared = Arc::new(x_data);
225 let y_shared = Arc::new(y);
226
227 let mut worker_handles = Vec::new();
229 for worker_id in 0..self.config.num_workers {
230 let worker_handle = self
231 .spawn_lda_worker(
232 worker_id,
233 Arc::clone(&x_shared),
234 Arc::clone(&y_shared),
235 tx.clone(),
236 config.clone(),
237 )
238 .await;
239 worker_handles.push(worker_handle);
240 }
241
242 let parameter_server_handle = self.spawn_parameter_server(rx).await;
244
245 let training_result = self
247 .run_async_training_loop(worker_handles, parameter_server_handle)
248 .await?;
249
250 Ok(training_result)
251 }
252
253 async fn spawn_lda_worker(
255 &self,
256 worker_id: usize,
257 x_data: Arc<Array2<Float>>, y: Arc<Array1<usize>>,
259 tx: mpsc::Sender<OptimizationMessage>,
260 _config: LinearDiscriminantAnalysisConfig,
261 ) -> JoinHandle<Result<()>> {
262 let config = self.config.clone();
263 let parameters = Arc::clone(&self.parameters);
264 let semaphore = Arc::clone(&self.semaphore);
265
266 tokio::spawn(async move {
267 let mut rng = fastrand::Rng::new();
268
269 for iteration in 0..config.max_iterations {
270 let _permit = semaphore.acquire().await.expect("value should be present");
272
273 let batch_indices: Vec<usize> = (0..config.batch_size)
275 .map(|_| rng.usize(0..x_data.nrows()))
276 .collect();
277
278 let x_batch =
279 Array2::from_shape_fn((config.batch_size, x_data.ncols()), |(i, j)| {
280 x_data[[batch_indices[i], j]]
281 });
282
283 let y_batch = Array1::from_shape_fn(config.batch_size, |i| y[batch_indices[i]]);
284
285 let gradients =
287 Self::compute_lda_gradients(&x_batch, &y_batch, ¶meters).await?;
288
289 let loss = Self::compute_lda_loss(&x_batch, &y_batch, ¶meters).await?;
291
292 let params_snapshot = parameters.read().expect("lock not poisoned").clone();
294 tx.send(OptimizationMessage::ParameterUpdate {
295 worker_id,
296 parameters: params_snapshot,
297 gradient: gradients,
298 loss,
299 })
300 .await
301 .map_err(|e| SklearsError::ProcessingError(format!("Channel send error: {}", e)))?;
302
303 if iteration % 10 == 0 {
305 tx.send(OptimizationMessage::CheckConvergence {
306 worker_id,
307 current_loss: loss,
308 })
309 .await
310 .map_err(|e| {
311 SklearsError::ProcessingError(format!("Channel send error: {}", e))
312 })?;
313 }
314
315 if iteration % 5 == 0 {
317 tokio::task::yield_now().await;
318 }
319 }
320
321 Ok(())
322 })
323 }
324
325 async fn spawn_parameter_server(
327 &self,
328 mut rx: mpsc::Receiver<OptimizationMessage>,
329 ) -> JoinHandle<Result<HashMap<String, Array2<Float>>>> {
330 let parameters = Arc::clone(&self.parameters);
331 let velocity = Arc::clone(&self.velocity);
332 let state = Arc::clone(&self.state);
333 let config = self.config.clone();
334
335 tokio::spawn(async move {
336 let mut update_timer =
337 tokio::time::interval(Duration::from_millis(config.update_frequency_ms));
338
339 loop {
340 tokio::select! {
341 msg = rx.recv() => {
343 match msg {
344 Some(OptimizationMessage::ParameterUpdate { worker_id, parameters: _, gradient, loss }) => {
345 Self::apply_parameter_update(
346 ¶meters,
347 &velocity,
348 &state,
349 &config,
350 worker_id,
351 gradient,
352 loss,
353 ).await?;
354 }
355 Some(OptimizationMessage::CheckConvergence { current_loss, .. }) => {
356 let converged = Self::check_convergence(&state, &config, current_loss).await?;
357 if converged {
358 let final_params = parameters.read().expect("lock not poisoned").clone();
359 return Ok(final_params);
360 }
361 }
362 Some(OptimizationMessage::TrainingComplete { final_parameters, .. }) => {
363 return Ok(final_parameters);
364 }
365 Some(OptimizationMessage::OptimizationError { error, .. }) => {
366 return Err(SklearsError::ProcessingError(error));
367 }
368 None => break, _ => {} }
371 }
372 _ = update_timer.tick() => {
374 Self::update_learning_rate(&state, &config).await?;
376 }
377 }
378 }
379
380 let final_params = parameters.read().expect("lock not poisoned").clone();
381 Ok(final_params)
382 })
383 }
384
385 async fn apply_parameter_update(
387 parameters: &Arc<RwLock<HashMap<String, Array2<Float>>>>,
388 velocity: &Arc<RwLock<HashMap<String, Array2<Float>>>>,
389 state: &Arc<RwLock<OptimizationState>>,
390 config: &AsyncOptimizationConfig,
391 _worker_id: usize,
392 gradients: HashMap<String, Array2<Float>>,
393 loss: Float,
394 ) -> Result<()> {
395 let current_lr = {
396 let state_read = state.read().expect("lock not poisoned");
397 let base_lr = config.learning_rate;
398 let decay_factor = config
399 .learning_rate_decay
400 .powf(state_read.iteration as Float);
401 base_lr * decay_factor
402 };
403
404 {
406 let mut params = parameters.write().expect("lock not poisoned");
407 let mut vel = velocity.write().expect("lock not poisoned");
408
409 for (param_name, gradient) in gradients.iter() {
410 if let (Some(param), Some(velocity_param)) =
411 (params.get_mut(param_name), vel.get_mut(param_name))
412 {
413 *velocity_param = config.momentum * &*velocity_param - current_lr * gradient;
415
416 *param = &*param + &*velocity_param;
418 }
419 }
420 }
421
422 {
424 let mut state_write = state.write().expect("lock not poisoned");
425 state_write.previous_loss = state_write.current_loss;
426 state_write.current_loss = loss;
427 state_write.iteration += 1;
428 state_write.last_update = Instant::now();
429 state_write.update_count += 1;
430
431 let gradient_norm: Float = gradients
433 .values()
434 .map(|grad| grad.iter().map(|&x| x * x).sum::<Float>())
435 .sum::<Float>()
436 .sqrt();
437
438 state_write.loss_history.push(loss);
439 state_write.gradient_norms.push(gradient_norm);
440 state_write.learning_rate_history.push(current_lr);
441
442 if state_write.loss_history.len() > 1000 {
444 state_write.loss_history.remove(0);
445 state_write.gradient_norms.remove(0);
446 state_write.learning_rate_history.remove(0);
447 }
448 }
449
450 Ok(())
451 }
452
453 async fn check_convergence(
455 state: &Arc<RwLock<OptimizationState>>,
456 config: &AsyncOptimizationConfig,
457 current_loss: Float,
458 ) -> Result<bool> {
459 let state_read = state.read().expect("lock not poisoned");
460
461 if state_read.iteration > 10 {
463 let loss_change = (state_read.previous_loss - current_loss).abs();
464 if loss_change < config.convergence_tolerance {
465 return Ok(true);
466 }
467 }
468
469 if let Some(&last_grad_norm) = state_read.gradient_norms.last() {
471 if last_grad_norm < config.convergence_tolerance {
472 return Ok(true);
473 }
474 }
475
476 if state_read.iteration >= config.max_iterations {
478 return Ok(true);
479 }
480
481 Ok(false)
482 }
483
484 async fn update_learning_rate(
486 state: &Arc<RwLock<OptimizationState>>,
487 config: &AsyncOptimizationConfig,
488 ) -> Result<()> {
489 if !config.adaptive_learning_rate {
490 return Ok(());
491 }
492
493 let state_write = state.write().expect("lock not poisoned");
494
495 if state_write.loss_history.len() >= 10 {
497 let recent_losses = &state_write.loss_history[state_write.loss_history.len() - 10..];
498 let is_decreasing = recent_losses
499 .windows(2)
500 .all(|window| window[1] <= window[0] + config.convergence_tolerance);
501
502 if !is_decreasing && state_write.iteration.is_multiple_of(50) {
503 log::info!("Reducing learning rate due to lack of progress");
506 }
507 }
508
509 Ok(())
510 }
511
512 async fn compute_lda_gradients(
514 x_batch: &Array2<Float>, y_batch: &Array1<usize>,
516 parameters: &Arc<RwLock<HashMap<String, Array2<Float>>>>,
517 ) -> Result<HashMap<String, Array2<Float>>> {
518 let params = parameters.read().expect("lock not poisoned");
519 let mut gradients = HashMap::new();
520
521 let class_means = params.get("class_means").expect("key not found");
523 let covariance = params.get("covariance").expect("key not found");
524
525 let (batch_size, n_features) = x_batch.dim();
526 let n_classes = class_means.nrows();
527
528 let mut means_gradient = Array2::zeros((n_classes, n_features));
530 for (i, &label) in y_batch.iter().enumerate() {
531 let x_sample = x_batch.row(i);
532 let mean_diff = x_sample.to_owned() - class_means.row(label);
533
534 let updated_row = means_gradient.row(label).to_owned() + &mean_diff;
536 means_gradient.row_mut(label).assign(&updated_row);
537 }
538 means_gradient /= batch_size as Float;
539
540 let mut cov_gradient = Array2::zeros(covariance.raw_dim());
542 for (i, &label) in y_batch.iter().enumerate() {
543 let x_sample = x_batch.row(i);
544 let mean_centered = x_sample.to_owned() - class_means.row(label);
545 let col = mean_centered.clone().insert_axis(Axis(1));
546 let row = mean_centered.insert_axis(Axis(0));
547 let outer_product = col.dot(&row);
548 cov_gradient = cov_gradient + outer_product;
549 }
550 cov_gradient /= batch_size as Float;
551
552 let priors_gradient = Array2::zeros((1, n_classes));
554
555 gradients.insert("class_means".to_string(), means_gradient);
556 gradients.insert("covariance".to_string(), cov_gradient);
557 gradients.insert("priors".to_string(), priors_gradient);
558
559 Ok(gradients)
560 }
561
562 async fn compute_lda_loss(
564 x_batch: &Array2<Float>, y_batch: &Array1<usize>,
566 parameters: &Arc<RwLock<HashMap<String, Array2<Float>>>>,
567 ) -> Result<Float> {
568 let params = parameters.read().expect("lock not poisoned");
569
570 let class_means = params.get("class_means").expect("key not found");
571 let covariance = params.get("covariance").expect("key not found");
572
573 let mut total_loss = 0.0;
574
575 for (i, &label) in y_batch.iter().enumerate() {
577 let x_sample = x_batch.row(i);
578 let class_mean = class_means.row(label);
579 let diff = x_sample.to_owned() - class_mean;
580
581 let loss = diff.dot(&diff);
583 total_loss += loss;
584 }
585
586 let reg_term = covariance.iter().map(|&x| x * x).sum::<Float>() * 0.001;
588 total_loss += reg_term;
589
590 Ok(total_loss / x_batch.nrows() as Float)
591 }
592
593 async fn run_async_training_loop(
595 &self,
596 worker_handles: Vec<JoinHandle<Result<()>>>,
597 parameter_server_handle: JoinHandle<Result<HashMap<String, Array2<Float>>>>,
598 ) -> Result<HashMap<String, Array2<Float>>> {
599 let final_parameters = parameter_server_handle.await.map_err(|e| {
601 SklearsError::ProcessingError(format!("Parameter server error: {}", e))
602 })??;
603
604 for handle in worker_handles {
606 handle.abort();
607 }
608
609 Ok(final_parameters)
610 }
611
612 pub fn get_optimization_stats(&self) -> OptimizationStats {
614 let state = self.state.read().expect("lock not poisoned");
615
616 OptimizationStats {
617 iteration: state.iteration,
618 current_loss: state.current_loss,
619 converged: state.converged,
620 training_time: state.start_time.elapsed(),
621 last_update_time: state.last_update.elapsed(),
622 update_count: state.update_count,
623 average_loss: if !state.loss_history.is_empty() {
624 state.loss_history.iter().sum::<Float>() / state.loss_history.len() as Float
625 } else {
626 0.0
627 },
628 loss_variance: if state.loss_history.len() > 1 {
629 let mean =
630 state.loss_history.iter().sum::<Float>() / state.loss_history.len() as Float;
631 state
632 .loss_history
633 .iter()
634 .map(|&x| (x - mean).powi(2))
635 .sum::<Float>()
636 / (state.loss_history.len() - 1) as Float
637 } else {
638 0.0
639 },
640 }
641 }
642
643 pub async fn async_train_qda(
645 &self,
646 x_data: Array2<Float>, y: Array1<usize>,
648 config: QuadraticDiscriminantAnalysisConfig,
649 ) -> Result<HashMap<String, Array2<Float>>> {
650 let (_n_samples, n_features) = x_data.dim();
651 let n_classes = y.iter().max().expect("collection should not be empty") + 1;
652
653 self.initialize_qda_parameters(n_features, n_classes)?;
655
656 let (tx, rx) = mpsc::channel::<OptimizationMessage>(self.config.channel_buffer_size);
658
659 let x_shared = Arc::new(x_data);
660 let y_shared = Arc::new(y);
661
662 let mut worker_handles = Vec::new();
664 for worker_id in 0..self.config.num_workers {
665 let worker_handle = self
666 .spawn_qda_worker(
667 worker_id,
668 Arc::clone(&x_shared),
669 Arc::clone(&y_shared),
670 tx.clone(),
671 config.clone(),
672 )
673 .await;
674 worker_handles.push(worker_handle);
675 }
676
677 let parameter_server_handle = self.spawn_parameter_server(rx).await;
678 let training_result = self
679 .run_async_training_loop(worker_handles, parameter_server_handle)
680 .await?;
681
682 Ok(training_result)
683 }
684
685 fn initialize_qda_parameters(&self, n_features: usize, n_classes: usize) -> Result<()> {
687 let mut parameters = self.parameters.write().expect("lock not poisoned");
688 let mut velocity = self.velocity.write().expect("lock not poisoned");
689
690 let class_means = Array2::zeros((n_classes, n_features));
692 parameters.insert("class_means".to_string(), class_means.clone());
693 velocity.insert(
694 "class_means".to_string(),
695 Array2::zeros((n_classes, n_features)),
696 );
697
698 let covariances = Array2::eye(n_features * n_classes);
700 parameters.insert("class_covariances".to_string(), covariances.clone());
701 velocity.insert(
702 "class_covariances".to_string(),
703 Array2::zeros((n_features * n_classes, n_features)),
704 );
705
706 let priors = Array2::from_elem((1, n_classes), 1.0 / n_classes as Float);
708 parameters.insert("priors".to_string(), priors.clone());
709 velocity.insert("priors".to_string(), Array2::zeros((1, n_classes)));
710
711 Ok(())
712 }
713
714 async fn spawn_qda_worker(
716 &self,
717 worker_id: usize,
718 x_data: Arc<Array2<Float>>, y: Arc<Array1<usize>>,
720 tx: mpsc::Sender<OptimizationMessage>,
721 _config: QuadraticDiscriminantAnalysisConfig,
722 ) -> JoinHandle<Result<()>> {
723 let config = self.config.clone();
724 let parameters = Arc::clone(&self.parameters);
725 let semaphore = Arc::clone(&self.semaphore);
726
727 tokio::spawn(async move {
728 let mut rng = fastrand::Rng::new();
730
731 for iteration in 0..config.max_iterations {
732 let _permit = semaphore.acquire().await.expect("value should be present");
733
734 let batch_indices: Vec<usize> = (0..config.batch_size)
735 .map(|_| rng.usize(0..x_data.nrows()))
736 .collect();
737
738 let x_batch =
739 Array2::from_shape_fn((config.batch_size, x_data.ncols()), |(i, j)| {
740 x_data[[batch_indices[i], j]]
741 });
742
743 let y_batch = Array1::from_shape_fn(config.batch_size, |i| y[batch_indices[i]]);
744
745 let gradients =
747 Self::compute_qda_gradients(&x_batch, &y_batch, ¶meters).await?;
748 let loss = Self::compute_qda_loss(&x_batch, &y_batch, ¶meters).await?;
749
750 let params_snapshot = parameters.read().expect("lock not poisoned").clone();
751 tx.send(OptimizationMessage::ParameterUpdate {
752 worker_id,
753 parameters: params_snapshot,
754 gradient: gradients,
755 loss,
756 })
757 .await
758 .map_err(|e| SklearsError::ProcessingError(format!("Channel send error: {}", e)))?;
759
760 if iteration % 10 == 0 {
761 tx.send(OptimizationMessage::CheckConvergence {
762 worker_id,
763 current_loss: loss,
764 })
765 .await
766 .map_err(|e| {
767 SklearsError::ProcessingError(format!("Channel send error: {}", e))
768 })?;
769 }
770
771 tokio::task::yield_now().await;
772 }
773
774 Ok(())
775 })
776 }
777
778 async fn compute_qda_gradients(
780 x_batch: &Array2<Float>, y_batch: &Array1<usize>,
782 parameters: &Arc<RwLock<HashMap<String, Array2<Float>>>>,
783 ) -> Result<HashMap<String, Array2<Float>>> {
784 Self::compute_lda_gradients(x_batch, y_batch, parameters).await
787 }
788
789 async fn compute_qda_loss(
791 x_batch: &Array2<Float>, y_batch: &Array1<usize>,
793 parameters: &Arc<RwLock<HashMap<String, Array2<Float>>>>,
794 ) -> Result<Float> {
795 Self::compute_lda_loss(x_batch, y_batch, parameters).await
798 }
799}
800
801#[derive(Debug, Clone)]
803pub struct OptimizationStats {
804 pub iteration: usize,
805 pub current_loss: Float,
806 pub converged: bool,
807 pub training_time: Duration,
808 pub last_update_time: Duration,
809 pub update_count: usize,
810 pub average_loss: Float,
811 pub loss_variance: Float,
812}
813
814pub struct AsyncDiscriminantAnalysis {
816 optimizer: AsyncDiscriminantOptimizer,
817 runtime: Option<tokio::runtime::Runtime>,
818}
819
820impl AsyncDiscriminantAnalysis {
821 pub fn new(config: AsyncOptimizationConfig) -> Self {
823 let optimizer = AsyncDiscriminantOptimizer::new(config);
824
825 let runtime = tokio::runtime::Builder::new_multi_thread()
827 .worker_threads(num_cpus::get())
828 .enable_all()
829 .build()
830 .ok();
831
832 Self { optimizer, runtime }
833 }
834
835 pub fn train_lda_async(
837 &self,
838 x_data: Array2<Float>, y: Array1<usize>,
840 config: LinearDiscriminantAnalysisConfig,
841 ) -> Result<HashMap<String, Array2<Float>>> {
842 match &self.runtime {
843 Some(rt) => rt.block_on(self.optimizer.async_train_lda(x_data, y, config)),
844 None => Err(SklearsError::ProcessingError(
845 "Async runtime not available".to_string(),
846 )),
847 }
848 }
849
850 pub fn train_qda_async(
852 &self,
853 x_data: Array2<Float>, y: Array1<usize>,
855 config: QuadraticDiscriminantAnalysisConfig,
856 ) -> Result<HashMap<String, Array2<Float>>> {
857 match &self.runtime {
858 Some(rt) => rt.block_on(self.optimizer.async_train_qda(x_data, y, config)),
859 None => Err(SklearsError::ProcessingError(
860 "Async runtime not available".to_string(),
861 )),
862 }
863 }
864
865 pub fn get_stats(&self) -> OptimizationStats {
867 self.optimizer.get_optimization_stats()
868 }
869
870 pub fn has_converged(&self) -> bool {
872 self.optimizer
873 .state
874 .read()
875 .expect("lock not poisoned")
876 .converged
877 }
878}
879
880pub struct DistributedDiscriminantAnalysis {
882 node_id: usize,
883 total_nodes: usize,
884 async_trainer: AsyncDiscriminantAnalysis,
885}
886
887impl DistributedDiscriminantAnalysis {
888 pub fn new(node_id: usize, total_nodes: usize, config: AsyncOptimizationConfig) -> Self {
890 Self {
891 node_id,
892 total_nodes,
893 async_trainer: AsyncDiscriminantAnalysis::new(config),
894 }
895 }
896
897 pub fn train_local_partition(
899 &self,
900 x_local: Array2<Float>, y_local: Array1<usize>,
902 lda_config: LinearDiscriminantAnalysisConfig,
903 ) -> Result<HashMap<String, Array2<Float>>> {
904 log::info!("Node {} training on local partition", self.node_id);
905
906 let local_params = self
908 .async_trainer
909 .train_lda_async(x_local, y_local, lda_config)?;
910
911 Ok(local_params)
915 }
916
917 pub fn aggregate_parameters(
919 &self,
920 local_params: HashMap<String, Array2<Float>>,
921 _all_node_params: Vec<HashMap<String, Array2<Float>>>,
922 ) -> Result<HashMap<String, Array2<Float>>> {
923 log::info!(
927 "Node {} aggregating parameters from {} nodes",
928 self.node_id,
929 self.total_nodes
930 );
931
932 Ok(local_params)
935 }
936}
937
938#[allow(non_snake_case)]
939#[cfg(test)]
940mod tests {
941 use super::*;
942
943 #[tokio::test]
944 async fn test_async_optimization_config() {
945 let config = AsyncOptimizationConfig::default();
946 assert!(config.num_workers > 0);
947 assert!(config.batch_size > 0);
948 assert!(config.learning_rate > 0.0);
949 }
950
951 #[tokio::test]
952 async fn test_parameter_initialization() {
953 let config = AsyncOptimizationConfig::default();
954 let optimizer = AsyncDiscriminantOptimizer::new(config);
955
956 let result = optimizer.initialize_parameters(4, 3);
957 assert!(result.is_ok());
958
959 let params = optimizer
960 .parameters
961 .read()
962 .expect("operation should succeed");
963 assert!(params.contains_key("class_means"));
964 assert!(params.contains_key("covariance"));
965 assert!(params.contains_key("priors"));
966 }
967
968 #[test]
969 fn test_async_discriminant_analysis_creation() {
970 let config = AsyncOptimizationConfig::default();
971 let _trainer = AsyncDiscriminantAnalysis::new(config);
972 }
974
975 #[test]
976 fn test_distributed_discriminant_analysis_creation() {
977 let config = AsyncOptimizationConfig::default();
978 let _distributed = DistributedDiscriminantAnalysis::new(0, 4, config);
979 }
981
982 #[tokio::test]
983 async fn test_optimization_state() {
984 let mut state = OptimizationState::default();
985 assert_eq!(state.iteration, 0);
986 assert!(state.current_loss.is_infinite());
987 assert!(!state.converged);
988
989 state.iteration = 10;
990 state.current_loss = 0.5;
991 assert_eq!(state.iteration, 10);
992 assert_eq!(state.current_loss, 0.5);
993 }
994}