1use scirs2_core::ndarray::{Array, Array1, Dimension, ScalarOperand};
18use scirs2_core::numeric::Float;
19use scirs2_core::parallel_ops::*;
20use std::fmt::Debug;
21
22use crate::error::Result;
23use crate::optimizers::Optimizer;
24
25#[derive(Debug)]
66pub struct ParallelOptimizer<O, A, D>
67where
68 O: Optimizer<A, D> + Clone + Send + Sync,
69 A: Float + ScalarOperand + Debug + Send + Sync,
70 D: Dimension,
71{
72 base_optimizer: O,
73 group_optimizers: Vec<O>,
75 _phantom_a: std::marker::PhantomData<A>,
76 _phantom_d: std::marker::PhantomData<D>,
77}
78
79impl<O, A, D> ParallelOptimizer<O, A, D>
80where
81 O: Optimizer<A, D> + Clone + Send + Sync,
82 A: Float + ScalarOperand + Debug + Send + Sync,
83 D: Dimension,
84{
85 pub fn new(base_optimizer: O) -> Self {
91 Self {
92 base_optimizer,
93 group_optimizers: Vec::new(),
94 _phantom_a: std::marker::PhantomData,
95 _phantom_d: std::marker::PhantomData,
96 }
97 }
98
99 pub fn num_groups(&self) -> usize {
101 self.group_optimizers.len()
102 }
103
104 pub fn group_optimizer(&self, group: usize) -> Option<&O> {
106 self.group_optimizers.get(group)
107 }
108
109 pub fn group_optimizer_mut(&mut self, group: usize) -> Option<&mut O> {
111 self.group_optimizers.get_mut(group)
112 }
113
114 pub fn reset_group_state(&mut self) {
116 self.group_optimizers.clear();
117 }
118
119 fn ensure_group_optimizers(&mut self, count: usize) {
121 while self.group_optimizers.len() < count {
122 let fresh = self.base_optimizer.clone();
123 self.group_optimizers.push(fresh);
124 }
125 }
126
127 pub fn step_parallel_groups(
141 &mut self,
142 params_list: &[Array<A, D>],
143 grads_list: &[Array<A, D>],
144 ) -> Result<Vec<Array<A, D>>>
145 where
146 Array<A, D>: Clone + Send + Sync,
147 {
148 if params_list.len() != grads_list.len() {
149 return Err(crate::error::OptimError::InvalidConfig(format!(
150 "Parameter groups ({}) and gradient groups ({}) must have same length",
151 params_list.len(),
152 grads_list.len()
153 )));
154 }
155
156 let num_groups = params_list.len();
159 self.ensure_group_optimizers(num_groups);
160
161 let results: Vec<Result<Array<A, D>>> = self.group_optimizers[..num_groups]
164 .par_iter_mut()
165 .zip(params_list.par_iter())
166 .zip(grads_list.par_iter())
167 .map(|((optimizer, params), grads)| optimizer.step(params, grads))
168 .collect();
169
170 let mut updated_params = Vec::with_capacity(results.len());
172 for result in results {
173 updated_params.push(result?);
174 }
175
176 Ok(updated_params)
177 }
178
179 pub fn inner(&self) -> &O {
181 &self.base_optimizer
182 }
183
184 pub fn inner_mut(&mut self) -> &mut O {
186 &mut self.base_optimizer
187 }
188
189 pub fn get_learning_rate(&self) -> A {
191 self.base_optimizer.get_learning_rate()
192 }
193
194 pub fn set_learning_rate(&mut self, learning_rate: A) {
196 self.base_optimizer.set_learning_rate(learning_rate);
197 for optimizer in self.group_optimizers.iter_mut() {
198 optimizer.set_learning_rate(learning_rate);
199 }
200 }
201}
202
203pub struct ParallelBatchProcessor {
208 min_chunk_size: usize,
210 num_threads: Option<usize>,
212}
213
214impl ParallelBatchProcessor {
215 pub fn new(min_chunk_size: usize) -> Self {
221 Self {
222 min_chunk_size,
223 num_threads: None,
224 }
225 }
226
227 pub fn with_threads(mut self, num_threads: Option<usize>) -> Self {
233 self.num_threads = num_threads;
234 self
235 }
236
237 pub fn should_use_parallel(&self, size: usize) -> bool {
247 let num_cores = num_cpus::get();
248 size >= self.min_chunk_size * num_cores
249 }
250
251 pub fn optimal_chunk_size(&self, total_size: usize) -> usize {
261 let num_cores = self.num_threads.unwrap_or_else(num_cpus::get);
262 let chunk_size = total_size / num_cores;
263 chunk_size.max(self.min_chunk_size)
264 }
265}
266
267impl Default for ParallelBatchProcessor {
268 fn default() -> Self {
269 Self::new(1024)
270 }
271}
272
273pub fn parallel_step<O, A, D>(
300 optimizer: &mut O,
301 params_list: &[Array<A, D>],
302 grads_list: &[Array<A, D>],
303) -> Result<Vec<Array<A, D>>>
304where
305 O: Optimizer<A, D> + Clone + Send + Sync,
306 A: Float + ScalarOperand + Debug + Send + Sync,
307 D: Dimension,
308 Array<A, D>: Clone + Send + Sync,
309{
310 if params_list.len() != grads_list.len() {
311 return Err(crate::error::OptimError::InvalidConfig(format!(
312 "Parameter groups ({}) and gradient groups ({}) must have same length",
313 params_list.len(),
314 grads_list.len()
315 )));
316 }
317
318 let params_refs: Vec<&Array<A, D>> = params_list.iter().collect();
319 let grads_refs: Vec<&Array<A, D>> = grads_list.iter().collect();
320 optimizer.step_list(¶ms_refs, &grads_refs)
321}
322
323pub fn parallel_step_array1<O, A>(
327 optimizer: &mut O,
328 params_list: &[Array1<A>],
329 grads_list: &[Array1<A>],
330) -> Result<Vec<Array1<A>>>
331where
332 O: Optimizer<A, scirs2_core::ndarray::Ix1> + Clone + Send + Sync,
333 A: Float + ScalarOperand + Debug + Send + Sync,
334{
335 if params_list.len() != grads_list.len() {
336 return Err(crate::error::OptimError::InvalidConfig(format!(
337 "Parameter groups ({}) and gradient groups ({}) must have same length",
338 params_list.len(),
339 grads_list.len()
340 )));
341 }
342
343 let params_refs: Vec<&Array1<A>> = params_list.iter().collect();
344 let grads_refs: Vec<&Array1<A>> = grads_list.iter().collect();
345 optimizer.step_list(¶ms_refs, &grads_refs)
346}
347
348#[cfg(test)]
349mod tests {
350 use super::*;
351 use crate::optimizers::{Adam, SGD};
352 use approx::assert_relative_eq;
353
354 #[test]
355 fn test_parallel_optimizer_basic() {
356 let optimizer = SGD::new(0.1);
357 let mut parallel_opt = ParallelOptimizer::new(optimizer);
358
359 let params_list = vec![
360 Array1::from_vec(vec![1.0f32, 2.0, 3.0]),
361 Array1::from_vec(vec![4.0, 5.0, 6.0]),
362 ];
363 let grads_list = vec![
364 Array1::from_vec(vec![0.1, 0.2, 0.3]),
365 Array1::from_vec(vec![0.1, 0.2, 0.3]),
366 ];
367
368 let results = parallel_opt
369 .step_parallel_groups(¶ms_list, &grads_list)
370 .expect("step_parallel_groups succeeds in test_parallel_optimizer_basic");
371
372 assert_eq!(results.len(), 2);
373 assert_relative_eq!(results[0][0], 0.99, epsilon = 1e-6);
374 assert_relative_eq!(results[1][0], 3.99, epsilon = 1e-6);
375 }
376
377 #[test]
378 fn test_parallel_optimizer_multiple_groups() {
379 let optimizer = SGD::new(0.01);
380 let mut parallel_opt = ParallelOptimizer::new(optimizer);
381
382 let params_list: Vec<Array1<f32>> =
384 (0..10).map(|i| Array1::from_elem(100, i as f32)).collect();
385 let grads_list: Vec<Array1<f32>> = (0..10).map(|_| Array1::from_elem(100, 0.1)).collect();
386
387 let results = parallel_opt
388 .step_parallel_groups(¶ms_list, &grads_list)
389 .expect("step_parallel_groups succeeds in test_parallel_optimizer_multiple_groups");
390
391 assert_eq!(results.len(), 10);
392 assert_relative_eq!(results[0][0], 0.0 - 0.01 * 0.1, epsilon = 1e-6);
394 }
395
396 #[test]
397 fn test_parallel_step_function() {
398 let mut optimizer = SGD::new(0.1);
399
400 let params_list = vec![
401 Array1::from_vec(vec![1.0f32, 2.0]),
402 Array1::from_vec(vec![3.0, 4.0]),
403 ];
404 let grads_list = vec![
405 Array1::from_vec(vec![0.1, 0.2]),
406 Array1::from_vec(vec![0.3, 0.4]),
407 ];
408
409 let results = parallel_step(&mut optimizer, ¶ms_list, &grads_list)
410 .expect("parallel_step succeeds in test_parallel_step_function");
411
412 assert_eq!(results.len(), 2);
413 assert_relative_eq!(results[0][0], 0.99, epsilon = 1e-6);
414 assert_relative_eq!(results[1][0], 2.97, epsilon = 1e-6);
415 }
416
417 #[test]
418 fn test_parallel_batch_processor() {
419 let processor = ParallelBatchProcessor::new(1024);
420
421 assert!(!processor.should_use_parallel(100));
423
424 let num_cores = num_cpus::get();
426 assert!(processor.should_use_parallel(1024 * num_cores * 2));
427
428 let chunk_size = processor.optimal_chunk_size(10000);
430 assert!(chunk_size >= 1024);
431 }
432
433 #[test]
434 fn test_parallel_batch_processor_threads() {
435 let processor = ParallelBatchProcessor::new(1024).with_threads(Some(4));
436
437 let chunk_size = processor.optimal_chunk_size(10000);
438 assert!(chunk_size >= 1024);
440 assert!(chunk_size <= 10000);
441 }
442
443 #[test]
444 fn test_parallel_optimizer_learning_rate() {
445 let optimizer = SGD::new(0.1);
446 let mut parallel_opt: ParallelOptimizer<_, f64, scirs2_core::ndarray::Ix1> =
447 ParallelOptimizer::new(optimizer);
448
449 assert_relative_eq!(parallel_opt.get_learning_rate(), 0.1, epsilon = 1e-6);
450
451 parallel_opt.set_learning_rate(0.2);
452 assert_relative_eq!(parallel_opt.get_learning_rate(), 0.2, epsilon = 1e-6);
453 }
454
455 #[test]
461 fn test_parallel_optimizer_preserves_adam_state() {
462 let optimizer = Adam::new(0.1f64);
463 let mut parallel_opt = ParallelOptimizer::new(optimizer);
464
465 let params = vec![Array1::from_vec(vec![0.0f64])];
466 let grads_first = vec![Array1::from_vec(vec![1.0f64])];
467 let grads_second = vec![Array1::from_vec(vec![0.0f64])];
468
469 let after_first = parallel_opt
470 .step_parallel_groups(¶ms, &grads_first)
471 .expect("first parallel step failed");
472 assert_relative_eq!(after_first[0][0], -0.1, epsilon = 1e-9);
474
475 let after_second = parallel_opt
476 .step_parallel_groups(&after_first, &grads_second)
477 .expect("second parallel step failed");
478
479 assert!(
481 (after_second[0][0] - after_first[0][0]).abs() > 1e-3,
482 "optimizer state was discarded between calls: {} vs {}",
483 after_second[0][0],
484 after_first[0][0]
485 );
486
487 assert_relative_eq!(after_second[0][0], -0.167014, epsilon = 1e-5);
489 assert_eq!(parallel_opt.num_groups(), 1);
490 }
491
492 #[test]
494 fn test_parallel_optimizer_groups_have_independent_state() {
495 let optimizer = Adam::new(0.1f64);
496 let mut parallel_opt = ParallelOptimizer::new(optimizer);
497
498 let params = vec![
499 Array1::from_vec(vec![0.0f64]),
500 Array1::from_vec(vec![0.0f64, 0.0]),
501 ];
502 let grads = vec![
503 Array1::from_vec(vec![1.0f64]),
504 Array1::from_vec(vec![1.0f64, 1.0]),
505 ];
506
507 let first = parallel_opt
508 .step_parallel_groups(¶ms, &grads)
509 .expect("first parallel step failed");
510 assert_eq!(parallel_opt.num_groups(), 2);
511 assert_eq!(first[0].len(), 1);
512 assert_eq!(first[1].len(), 2);
513
514 let second = parallel_opt
515 .step_parallel_groups(&first, &grads)
516 .expect("second parallel step failed");
517
518 assert_relative_eq!(second[0][0], second[1][0], epsilon = 1e-12);
520 assert_relative_eq!(second[1][0], second[1][1], epsilon = 1e-12);
521
522 let first_delta = (first[0][0] - 0.0).abs();
524 let second_delta = (second[0][0] - first[0][0]).abs();
525 assert!(second_delta < first_delta);
526 }
527
528 #[test]
530 fn test_parallel_step_array1_preserves_state() {
531 let mut optimizer = Adam::new(0.1f64);
532
533 let params = vec![Array1::from_vec(vec![0.0f64])];
534 let grads_first = vec![Array1::from_vec(vec![1.0f64])];
535 let grads_second = vec![Array1::from_vec(vec![0.0f64])];
536
537 let first =
538 parallel_step_array1(&mut optimizer, ¶ms, &grads_first).expect("first step failed");
539 let second = parallel_step_array1(&mut optimizer, &first, &grads_second)
540 .expect("second step failed");
541
542 assert_relative_eq!(first[0][0], -0.1, epsilon = 1e-9);
543 assert_relative_eq!(second[0][0], -0.167014, epsilon = 1e-5);
544 }
545
546 #[test]
547 fn test_parallel_optimizer_learning_rate_propagates_to_groups() {
548 let optimizer = SGD::new(0.1f64);
549 let mut parallel_opt = ParallelOptimizer::new(optimizer);
550
551 let params = vec![Array1::from_vec(vec![1.0f64])];
552 let grads = vec![Array1::from_vec(vec![1.0f64])];
553
554 let _ = parallel_opt
555 .step_parallel_groups(¶ms, &grads)
556 .expect("step failed");
557 parallel_opt.set_learning_rate(0.5);
558
559 let updated = parallel_opt
560 .step_parallel_groups(¶ms, &grads)
561 .expect("step failed");
562 assert_relative_eq!(updated[0][0], 0.5, epsilon = 1e-9);
563 }
564
565 #[test]
566 fn test_parallel_step_array1() {
567 let mut optimizer = SGD::new(0.1);
568
569 let params_list = vec![
570 Array1::from_vec(vec![1.0f32, 2.0, 3.0]),
571 Array1::from_vec(vec![4.0, 5.0, 6.0]),
572 ];
573 let grads_list = vec![
574 Array1::from_vec(vec![0.1, 0.2, 0.3]),
575 Array1::from_vec(vec![0.1, 0.2, 0.3]),
576 ];
577
578 let results = parallel_step_array1(&mut optimizer, ¶ms_list, &grads_list)
579 .expect("parallel_step_array1 succeeds in test_parallel_step_array1");
580
581 assert_eq!(results.len(), 2);
582 assert_relative_eq!(results[0][0], 0.99, epsilon = 1e-6);
583 assert_relative_eq!(results[1][0], 3.99, epsilon = 1e-6);
584 }
585}