Skip to main content

optirs_core/distributed/
averaging.rs

1use crate::error::{OptimError, Result};
2use scirs2_core::ndarray::{Array, Dimension, ScalarOperand, Zip};
3use scirs2_core::numeric::Float;
4use std::collections::HashMap;
5use std::fmt::Debug;
6
7/// Parameter averaging strategies for distributed training
8#[derive(Debug, Clone, Copy, PartialEq)]
9pub enum AveragingStrategy {
10    /// Simple arithmetic mean
11    Arithmetic,
12    /// Weighted average based on data sizes
13    WeightedByData,
14    /// Weighted average based on computation times
15    WeightedByTime,
16    /// Federated averaging (FedAvg)
17    Federated,
18    /// Momentum-based averaging
19    Momentum {
20        /// Momentum factor
21        momentum: f64,
22    },
23    /// Exponentially weighted moving average
24    ExponentialMovingAverage {
25        /// Decay factor
26        decay: f64,
27    },
28}
29
30/// Distributed parameter averager
31#[derive(Debug)]
32pub struct ParameterAverager<A: Float, D: Dimension> {
33    /// Current averaged parameters
34    averaged_params: Vec<Array<A, D>>,
35    /// Averaging strategy
36    strategy: AveragingStrategy,
37    /// Node weights for weighted averaging
38    node_weights: HashMap<usize, A>,
39    /// Number of participating nodes
40    numnodes: usize,
41    /// Momentum buffer for momentum-based averaging
42    momentum_buffer: Option<Vec<Array<A, D>>>,
43    /// Step count for EMA decay adjustment
44    step_count: usize,
45    /// Whether averager is initialized
46    initialized: bool,
47}
48
49impl<A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension + Send + Sync>
50    ParameterAverager<A, D>
51{
52    /// Create a new parameter averager
53    pub fn new(strategy: AveragingStrategy, numnodes: usize) -> Self {
54        Self {
55            averaged_params: Vec::new(),
56            strategy,
57            node_weights: HashMap::new(),
58            numnodes,
59            momentum_buffer: None,
60            step_count: 0,
61            initialized: false,
62        }
63    }
64
65    /// Initialize averager with parameter shapes
66    pub fn initialize(&mut self, params: &[Array<A, D>]) -> Result<()> {
67        if self.initialized {
68            return Err(OptimError::InvalidConfig(
69                "Parameter averager already initialized".to_string(),
70            ));
71        }
72
73        self.averaged_params = params.to_vec();
74
75        // Initialize momentum buffer if needed
76        if matches!(self.strategy, AveragingStrategy::Momentum { .. }) {
77            self.momentum_buffer = Some(params.iter().map(|p| Array::zeros(p.raw_dim())).collect());
78        }
79
80        // Initialize uniform weights
81        let numnodes_a = A::from(self.numnodes).ok_or_else(|| {
82            OptimError::InvalidConfig(format!(
83                "node count {} could not be represented in the parameter type",
84                self.numnodes
85            ))
86        })?;
87        let uniform_weight = A::one() / numnodes_a;
88        for nodeid in 0..self.numnodes {
89            self.node_weights.insert(nodeid, uniform_weight);
90        }
91
92        self.initialized = true;
93        Ok(())
94    }
95
96    /// Set weight for a specific node
97    pub fn set_node_weight(&mut self, nodeid: usize, weight: A) -> Result<()> {
98        if nodeid >= self.numnodes {
99            return Err(OptimError::InvalidConfig(format!(
100                "Node ID {} exceeds number of nodes {}",
101                nodeid, self.numnodes
102            )));
103        }
104        self.node_weights.insert(nodeid, weight);
105        Ok(())
106    }
107
108    /// Average parameters from multiple nodes
109    pub fn average_parameters(
110        &mut self,
111        nodeparameters: &[(usize, Vec<Array<A, D>>)],
112    ) -> Result<()> {
113        if !self.initialized {
114            if let Some((_, first_params)) = nodeparameters.first() {
115                self.initialize(first_params)?;
116            } else {
117                return Err(OptimError::InvalidConfig(
118                    "No _parameters provided for initialization".to_string(),
119                ));
120            }
121        }
122
123        // Validate input
124        for (nodeid, params) in nodeparameters {
125            if *nodeid >= self.numnodes {
126                return Err(OptimError::InvalidConfig(format!(
127                    "Node ID {} exceeds number of nodes {}",
128                    nodeid, self.numnodes
129                )));
130            }
131            if params.len() != self.averaged_params.len() {
132                return Err(OptimError::DimensionMismatch(format!(
133                    "Expected {} parameter arrays, got {}",
134                    self.averaged_params.len(),
135                    params.len()
136                )));
137            }
138        }
139
140        self.step_count += 1;
141
142        match self.strategy {
143            AveragingStrategy::Arithmetic => {
144                self.arithmetic_average(nodeparameters)?;
145            }
146            AveragingStrategy::WeightedByData | AveragingStrategy::WeightedByTime => {
147                self.weighted_average(nodeparameters)?;
148            }
149            AveragingStrategy::Federated => {
150                self.federated_average(nodeparameters)?;
151            }
152            AveragingStrategy::Momentum { momentum } => {
153                self.momentum_average(nodeparameters, momentum)?;
154            }
155            AveragingStrategy::ExponentialMovingAverage { decay } => {
156                self.ema_average(nodeparameters, decay)?;
157            }
158        }
159
160        Ok(())
161    }
162
163    /// Simple arithmetic averaging
164    fn arithmetic_average(&mut self, nodeparameters: &[(usize, Vec<Array<A, D>>)]) -> Result<()> {
165        // Reset averaged _parameters
166        for param in &mut self.averaged_params {
167            param.fill(A::zero());
168        }
169
170        let numnodes = A::from(nodeparameters.len()).ok_or_else(|| {
171            OptimError::InvalidConfig(format!(
172                "node count {} could not be represented in the parameter type",
173                nodeparameters.len()
174            ))
175        })?;
176
177        // Sum all _parameters
178        for (_node_id, params) in nodeparameters {
179            for (avg_param, param) in self.averaged_params.iter_mut().zip(params.iter()) {
180                Zip::from(avg_param).and(param).for_each(|avg, &p| {
181                    *avg = *avg + p;
182                });
183            }
184        }
185
186        // Divide by number of nodes
187        for param in &mut self.averaged_params {
188            param.mapv_inplace(|x| x / numnodes);
189        }
190
191        Ok(())
192    }
193
194    /// Weighted averaging using node weights
195    fn weighted_average(&mut self, nodeparameters: &[(usize, Vec<Array<A, D>>)]) -> Result<()> {
196        // Reset averaged _parameters
197        for param in &mut self.averaged_params {
198            param.fill(A::zero());
199        }
200
201        // Compute total weight
202        let total_weight: A = nodeparameters
203            .iter()
204            .map(|(nodeid, _)| self.node_weights.get(nodeid).copied().unwrap_or(A::zero()))
205            .fold(A::zero(), |acc, w| acc + w);
206
207        if total_weight <= A::zero() {
208            return Err(OptimError::InvalidConfig(
209                "Total node weights must be > 0".to_string(),
210            ));
211        }
212
213        // Weighted sum
214        for (nodeid, params) in nodeparameters {
215            let weight = self.node_weights.get(nodeid).copied().unwrap_or(A::zero()) / total_weight;
216
217            for (avg_param, param) in self.averaged_params.iter_mut().zip(params.iter()) {
218                Zip::from(avg_param).and(param).for_each(|avg, &p| {
219                    *avg = *avg + weight * p;
220                });
221            }
222        }
223
224        Ok(())
225    }
226
227    /// Federated averaging (FedAvg). This delegates to the same weighted-average
228    /// machinery as `WeightedByData`: the caller MUST call `set_node_weight` with
229    /// each node's local sample-size fraction before invoking `average_parameters`,
230    /// otherwise `initialize` seeds uniform weights and this degenerates to plain
231    /// `Arithmetic` averaging (not an error, but not FedAvg's defining property
232    /// either -- callers wanting genuine FedAvg with only local dataset sizes
233    /// available should prefer `distributed::fedprox::FedProxOptimizer`, which
234    /// accepts sample counts directly).
235    fn federated_average(&mut self, nodeparameters: &[(usize, Vec<Array<A, D>>)]) -> Result<()> {
236        self.weighted_average(nodeparameters)
237    }
238
239    /// Momentum-based averaging
240    fn momentum_average(
241        &mut self,
242        nodeparameters: &[(usize, Vec<Array<A, D>>)],
243        momentum: f64,
244    ) -> Result<()> {
245        if !(0.0..=1.0).contains(&momentum) {
246            return Err(OptimError::InvalidConfig(format!(
247                "momentum must be in [0, 1], got {momentum}"
248            )));
249        }
250        let momentum_factor = A::from(momentum).ok_or_else(|| {
251            OptimError::InvalidConfig(format!(
252                "momentum {momentum} could not be represented in the parameter type"
253            ))
254        })?;
255        let one_minus_momentum = A::one() - momentum_factor;
256
257        // First compute arithmetic average of incoming _parameters
258        let mut current_average: Vec<Array<A, D>> = self
259            .averaged_params
260            .iter()
261            .map(|param| Array::zeros(param.raw_dim()))
262            .collect();
263
264        let numnodes = A::from(nodeparameters.len()).ok_or_else(|| {
265            OptimError::InvalidConfig(format!(
266                "node count {} could not be represented in the parameter type",
267                nodeparameters.len()
268            ))
269        })?;
270        for (_node_id, params) in nodeparameters {
271            for (avg_param, param) in current_average.iter_mut().zip(params.iter()) {
272                Zip::from(avg_param).and(param).for_each(|avg, &p| {
273                    *avg = *avg + p / numnodes;
274                });
275            }
276        }
277
278        // Apply momentum update. The momentum buffer is only allocated by
279        // `initialize` when the strategy is `Momentum` *at that moment* -- if
280        // the caller switched strategies afterward (or never initialized this
281        // way), silently discarding the incoming update would corrupt training
282        // with no signal, so we fail loudly instead.
283        let momentum_buf = self.momentum_buffer.as_mut().ok_or_else(|| {
284            OptimError::InvalidState(
285                "Momentum averaging selected but the momentum buffer was never initialized; \
286                 call initialize() while the strategy is AveragingStrategy::Momentum"
287                    .to_string(),
288            )
289        })?;
290
291        for ((avg_param, current_param), momentum_param) in self
292            .averaged_params
293            .iter_mut()
294            .zip(current_average.iter())
295            .zip(momentum_buf.iter_mut())
296        {
297            // Update momentum buffer first
298            Zip::from(&mut *momentum_param)
299                .and(current_param)
300                .for_each(|mom, &curr| {
301                    *mom = momentum_factor * *mom + one_minus_momentum * curr;
302                });
303
304            // Copy momentum buffer to averaged params
305            avg_param.assign(&*momentum_param);
306        }
307
308        Ok(())
309    }
310
311    /// Exponential moving average
312    fn ema_average(
313        &mut self,
314        nodeparameters: &[(usize, Vec<Array<A, D>>)],
315        decay: f64,
316    ) -> Result<()> {
317        if !(0.0..=1.0).contains(&decay) {
318            return Err(OptimError::InvalidConfig(format!(
319                "decay must be in [0, 1], got {decay}"
320            )));
321        }
322        let decay_factor = A::from(decay).ok_or_else(|| {
323            OptimError::InvalidConfig(format!(
324                "decay {decay} could not be represented in the parameter type"
325            ))
326        })?;
327        let one_minus_decay = A::one() - decay_factor;
328
329        // First compute arithmetic average of incoming _parameters
330        let mut current_average: Vec<Array<A, D>> = self
331            .averaged_params
332            .iter()
333            .map(|param| Array::zeros(param.raw_dim()))
334            .collect();
335
336        let numnodes = A::from(nodeparameters.len()).ok_or_else(|| {
337            OptimError::InvalidConfig(format!(
338                "node count {} could not be represented in the parameter type",
339                nodeparameters.len()
340            ))
341        })?;
342        for (_node_id, params) in nodeparameters {
343            for (avg_param, param) in current_average.iter_mut().zip(params.iter()) {
344                Zip::from(avg_param).and(param).for_each(|avg, &p| {
345                    *avg = *avg + p / numnodes;
346                });
347            }
348        }
349
350        // Apply EMA update
351        for (avg_param, current_param) in
352            self.averaged_params.iter_mut().zip(current_average.iter())
353        {
354            Zip::from(avg_param)
355                .and(current_param)
356                .for_each(|avg, &curr| {
357                    *avg = decay_factor * *avg + one_minus_decay * curr;
358                });
359        }
360
361        Ok(())
362    }
363
364    /// Get current averaged parameters
365    pub fn get_averaged_parameters(&self) -> &[Array<A, D>] {
366        &self.averaged_params
367    }
368
369    /// Get cloned averaged parameters
370    pub fn get_averaged_parameters_cloned(&self) -> Vec<Array<A, D>> {
371        self.averaged_params.clone()
372    }
373
374    /// Reset averager state
375    pub fn reset(&mut self) {
376        self.step_count = 0;
377        for param in &mut self.averaged_params {
378            param.fill(A::zero());
379        }
380        if let Some(ref mut momentum_buf) = self.momentum_buffer {
381            for buf in momentum_buf {
382                buf.fill(A::zero());
383            }
384        }
385    }
386
387    /// Get step count
388    pub fn step_count(&self) -> usize {
389        self.step_count
390    }
391
392    /// Get number of nodes
393    pub fn numnodes(&self) -> usize {
394        self.numnodes
395    }
396
397    /// Get averaging strategy
398    pub fn strategy(&self) -> AveragingStrategy {
399        self.strategy
400    }
401
402    /// Check if initialized
403    pub fn is_initialized(&self) -> bool {
404        self.initialized
405    }
406}