Skip to main content

optirs_core/
parallel_optimizer.rs

1//! Parallel optimizer operations using scirs2_core
2//!
3//! This module provides parallel processing capabilities for optimizers,
4//! enabling efficient multi-core utilization for large-scale optimization.
5//!
6//! # Features
7//!
8//! - Parallel parameter group processing
9//! - Parallel batch updates
10//! - Automatic work distribution across CPU cores
11//! - Zero-copy parameter handling
12//!
13//! # Performance
14//!
15//! Expected speedup: 4-8x on multi-core systems for multiple parameter groups
16
17use 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/// Parallel optimizer wrapper for processing multiple parameter groups
26///
27/// This wrapper enables parallel processing of multiple parameter groups,
28/// providing significant speedup on multi-core systems.
29///
30/// # Examples
31///
32/// ```
33/// use scirs2_core::ndarray::Array1;
34/// use optirs_core::optimizers::{SGD, Optimizer};
35/// use optirs_core::parallel_optimizer::ParallelOptimizer;
36///
37/// // Create base optimizer
38/// let optimizer = SGD::new(0.01);
39///
40/// // Wrap in parallel optimizer
41/// let mut parallel_opt = ParallelOptimizer::new(optimizer);
42///
43/// // Process multiple parameter groups in parallel
44/// let params_list = vec![
45///     Array1::zeros(1000),
46///     Array1::zeros(2000),
47///     Array1::zeros(1500),
48/// ];
49/// let grads_list = vec![
50///     Array1::from_elem(1000, 0.1),
51///     Array1::from_elem(2000, 0.1),
52///     Array1::from_elem(1500, 0.1),
53/// ];
54///
55/// let updated = parallel_opt.step_parallel_groups(&params_list, &grads_list).expect("parallel_opt.step_parallel_groups succeeds");
56/// ```
57///
58/// # State handling
59///
60/// Stateful optimizers (Adam, AdamW, LAMB, RAdam, Lion, ...) keep per-parameter
61/// moment estimates. A dedicated optimizer instance is therefore materialized and
62/// **retained** for every parameter group, and each group's instance is mutated in
63/// place across calls. Without this, every call would optimize with a freshly reset
64/// clone and, for example, Adam would degenerate into sign-SGD.
65#[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    /// Persistent per-group optimizer instances (index == parameter-group index)
74    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    /// Creates a new parallel optimizer wrapper
86    ///
87    /// # Arguments
88    ///
89    /// * `base_optimizer` - The base optimizer to parallelize
90    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    /// Number of parameter groups for which persistent state is currently held
100    pub fn num_groups(&self) -> usize {
101        self.group_optimizers.len()
102    }
103
104    /// Access the persistent optimizer instance of a parameter group
105    pub fn group_optimizer(&self, group: usize) -> Option<&O> {
106        self.group_optimizers.get(group)
107    }
108
109    /// Access the persistent optimizer instance of a parameter group mutably
110    pub fn group_optimizer_mut(&mut self, group: usize) -> Option<&mut O> {
111        self.group_optimizers.get_mut(group)
112    }
113
114    /// Drop all per-group state, restarting every group from the base optimizer
115    pub fn reset_group_state(&mut self) {
116        self.group_optimizers.clear();
117    }
118
119    /// Grow the per-group optimizer pool so that `count` groups have persistent state
120    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    /// Process multiple parameter groups in parallel
128    ///
129    /// This method distributes parameter groups across available CPU cores
130    /// for parallel processing.
131    ///
132    /// # Arguments
133    ///
134    /// * `params_list` - List of parameter arrays
135    /// * `grads_list` - List of gradient arrays
136    ///
137    /// # Returns
138    ///
139    /// Updated parameter arrays processed in parallel
140    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        // Materialize (once) and then reuse a persistent optimizer instance per group so
157        // that momentum / second-moment state survives across calls.
158        let num_groups = params_list.len();
159        self.ensure_group_optimizers(num_groups);
160
161        // Use parallel iterator from scirs2_core. Each group mutates its own optimizer,
162        // so the borrows are disjoint and no state is discarded.
163        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        // Collect results and handle errors
171        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    /// Get the underlying optimizer
180    pub fn inner(&self) -> &O {
181        &self.base_optimizer
182    }
183
184    /// Get mutable reference to underlying optimizer
185    pub fn inner_mut(&mut self) -> &mut O {
186        &mut self.base_optimizer
187    }
188
189    /// Get the current learning rate from the base optimizer
190    pub fn get_learning_rate(&self) -> A {
191        self.base_optimizer.get_learning_rate()
192    }
193
194    /// Set the learning rate on the base optimizer and on every per-group instance
195    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
203/// Parallel batch processor for large parameter arrays
204///
205/// This processor splits large parameter arrays into chunks and processes
206/// them in parallel, providing speedup for very large models.
207pub struct ParallelBatchProcessor {
208    /// Minimum chunk size for parallel processing
209    min_chunk_size: usize,
210    /// Number of threads to use (None = automatic)
211    num_threads: Option<usize>,
212}
213
214impl ParallelBatchProcessor {
215    /// Creates a new parallel batch processor
216    ///
217    /// # Arguments
218    ///
219    /// * `min_chunk_size` - Minimum size of each chunk (default: 1024)
220    pub fn new(min_chunk_size: usize) -> Self {
221        Self {
222            min_chunk_size,
223            num_threads: None,
224        }
225    }
226
227    /// Set the number of threads to use
228    ///
229    /// # Arguments
230    ///
231    /// * `num_threads` - Number of threads (None for automatic)
232    pub fn with_threads(mut self, num_threads: Option<usize>) -> Self {
233        self.num_threads = num_threads;
234        self
235    }
236
237    /// Determine if parallel processing should be used
238    ///
239    /// # Arguments
240    ///
241    /// * `size` - Size of the parameter array
242    ///
243    /// # Returns
244    ///
245    /// True if parallel processing would be beneficial
246    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    /// Get optimal chunk size for parallel processing
252    ///
253    /// # Arguments
254    ///
255    /// * `total_size` - Total size of the array
256    ///
257    /// # Returns
258    ///
259    /// Optimal chunk size for parallel processing
260    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
273/// Helper function to update several parameter groups with a single optimizer
274///
275/// This is a convenience function for one-off multi-group processing without
276/// creating a [`ParallelOptimizer`] instance.
277///
278/// The update is delegated to [`Optimizer::step_list`], which keeps an independent
279/// state slot per parameter-group index. Consequently the optimizer state is
280/// preserved across calls and groups never share moments.
281///
282/// # Note on parallelism
283///
284/// A single `&mut O` cannot be mutated from several threads at once, so this helper
285/// walks the groups sequentially. Use [`ParallelOptimizer::step_parallel_groups`]
286/// when you want the groups themselves processed in parallel: it keeps one
287/// optimizer instance per group and therefore both parallelizes *and* preserves
288/// state.
289///
290/// # Arguments
291///
292/// * `optimizer` - The optimizer to use (its per-index state is updated in place)
293/// * `params_list` - List of parameter arrays
294/// * `grads_list` - List of gradient arrays
295///
296/// # Returns
297///
298/// Updated parameter arrays
299pub 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(&params_refs, &grads_refs)
321}
322
323/// Multi-group processing for `Array1` specifically (optimized path)
324///
325/// See [`parallel_step`] for the state and parallelism semantics.
326pub 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(&params_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(&params_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        // Create 10 parameter groups
383        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(&params_list, &grads_list)
389            .expect("step_parallel_groups succeeds in test_parallel_optimizer_multiple_groups");
390
391        assert_eq!(results.len(), 10);
392        // Verify first group was updated correctly
393        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, &params_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        // Small array - should not use parallel
422        assert!(!processor.should_use_parallel(100));
423
424        // Large array - should use parallel
425        let num_cores = num_cpus::get();
426        assert!(processor.should_use_parallel(1024 * num_cores * 2));
427
428        // Test optimal chunk size calculation
429        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        // With 4 threads, chunk size should be around 10000/4 = 2500
439        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    /// Regression test for the "parallel wrapper discards optimizer state" bug.
456    ///
457    /// The wrapper used to clone the base optimizer on every call and throw the clone
458    /// away, so Adam always ran at t=1 with zero moments, i.e. it degenerated into
459    /// sign-SGD (`lr * sign(g)`), and a zero gradient produced no movement at all.
460    #[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(&params, &grads_first)
471            .expect("first parallel step failed");
472        // t = 1 with a unit gradient is exactly -lr for Adam.
473        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        // A stateless (bugged) optimizer sees m = v = 0 and does not move at all.
480        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        // With retained state: m = 0.09, v = 0.000999 => step ~= 0.067014
488        assert_relative_eq!(after_second[0][0], -0.167014, epsilon = 1e-5);
489        assert_eq!(parallel_opt.num_groups(), 1);
490    }
491
492    /// Two groups must never share optimizer state.
493    #[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(&params, &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        // Both groups are at t = 2 with identical gradients, so both must agree.
519        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        // And the second step must be smaller than the first (bias correction at t=2).
523        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    /// `parallel_step_array1` keeps per-group state in the optimizer it is given.
529    #[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, &params, &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(&params, &grads)
556            .expect("step failed");
557        parallel_opt.set_learning_rate(0.5);
558
559        let updated = parallel_opt
560            .step_parallel_groups(&params, &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, &params_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}