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#[derive(Debug, Clone, Copy, PartialEq)]
9pub enum AveragingStrategy {
10 Arithmetic,
12 WeightedByData,
14 WeightedByTime,
16 Federated,
18 Momentum {
20 momentum: f64,
22 },
23 ExponentialMovingAverage {
25 decay: f64,
27 },
28}
29
30#[derive(Debug)]
32pub struct ParameterAverager<A: Float, D: Dimension> {
33 averaged_params: Vec<Array<A, D>>,
35 strategy: AveragingStrategy,
37 node_weights: HashMap<usize, A>,
39 numnodes: usize,
41 momentum_buffer: Option<Vec<Array<A, D>>>,
43 step_count: usize,
45 initialized: bool,
47}
48
49impl<A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension + Send + Sync>
50 ParameterAverager<A, D>
51{
52 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 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 if matches!(self.strategy, AveragingStrategy::Momentum { .. }) {
77 self.momentum_buffer = Some(params.iter().map(|p| Array::zeros(p.raw_dim())).collect());
78 }
79
80 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 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 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 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 fn arithmetic_average(&mut self, nodeparameters: &[(usize, Vec<Array<A, D>>)]) -> Result<()> {
165 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 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 for param in &mut self.averaged_params {
188 param.mapv_inplace(|x| x / numnodes);
189 }
190
191 Ok(())
192 }
193
194 fn weighted_average(&mut self, nodeparameters: &[(usize, Vec<Array<A, D>>)]) -> Result<()> {
196 for param in &mut self.averaged_params {
198 param.fill(A::zero());
199 }
200
201 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 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 fn federated_average(&mut self, nodeparameters: &[(usize, Vec<Array<A, D>>)]) -> Result<()> {
236 self.weighted_average(nodeparameters)
237 }
238
239 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 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 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 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 avg_param.assign(&*momentum_param);
306 }
307
308 Ok(())
309 }
310
311 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 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 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 pub fn get_averaged_parameters(&self) -> &[Array<A, D>] {
366 &self.averaged_params
367 }
368
369 pub fn get_averaged_parameters_cloned(&self) -> Vec<Array<A, D>> {
371 self.averaged_params.clone()
372 }
373
374 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 pub fn step_count(&self) -> usize {
389 self.step_count
390 }
391
392 pub fn numnodes(&self) -> usize {
394 self.numnodes
395 }
396
397 pub fn strategy(&self) -> AveragingStrategy {
399 self.strategy
400 }
401
402 pub fn is_initialized(&self) -> bool {
404 self.initialized
405 }
406}