Skip to main content

scirs2_core/distributed/param_server/
server.rs

1//! Parameter server with multiple consistency models
2//!
3//! Supports BSP (Bulk Synchronous Parallel), ASP (Asynchronous Parallel),
4//! and SSP (Stale Synchronous Parallel) consistency models.
5
6use std::collections::HashMap;
7
8use crate::error::{CoreError, CoreResult, ErrorContext};
9
10use super::types::{
11    AggregationMethod, ConsistencyModel, ParamServerConfig, ParameterUpdate, WorkerState,
12};
13
14/// A checkpoint of the parameter server state
15#[derive(Debug, Clone)]
16pub struct ServerCheckpoint {
17    /// Snapshot of all parameters and their versions
18    pub parameters: HashMap<String, (Vec<f64>, u64)>,
19    /// Global version at checkpoint time
20    pub version: u64,
21    /// Snapshot of all worker states
22    pub worker_states: Vec<WorkerState>,
23}
24
25/// Parameter server supporting BSP, ASP, and SSP consistency models
26#[derive(Debug)]
27pub struct ParameterServer {
28    /// Configuration
29    config: ParamServerConfig,
30    /// In-memory parameter store: key -> (values, version)
31    parameters: HashMap<String, (Vec<f64>, u64)>,
32    /// Registered workers
33    workers: Vec<WorkerState>,
34    /// Global version counter
35    global_version: u64,
36    /// Buffered updates for BSP/SSP (key -> list of updates)
37    update_buffer: HashMap<String, Vec<ParameterUpdate>>,
38    /// Set of worker IDs that have pushed in the current BSP round
39    bsp_pushed_workers: Vec<bool>,
40}
41
42impl ParameterServer {
43    /// Create a new parameter server with the given configuration
44    #[must_use]
45    pub fn new(config: ParamServerConfig) -> Self {
46        let num_workers = config.num_workers;
47        Self {
48            config,
49            parameters: HashMap::new(),
50            workers: Vec::new(),
51            global_version: 0,
52            update_buffer: HashMap::new(),
53            bsp_pushed_workers: vec![false; num_workers],
54        }
55    }
56
57    /// Register a new worker and return its ID
58    pub fn register_worker(&mut self) -> usize {
59        let worker_id = self.workers.len();
60        self.workers.push(WorkerState::new(worker_id));
61        // Extend BSP tracking if needed
62        if self.bsp_pushed_workers.len() <= worker_id {
63            self.bsp_pushed_workers.resize(worker_id + 1, false);
64        }
65        worker_id
66    }
67
68    /// Initialize a parameter with the given key and values
69    pub fn init_parameter(&mut self, key: impl Into<String>, values: Vec<f64>) {
70        let key = key.into();
71        self.parameters.entry(key).or_insert((values, 0));
72    }
73
74    /// Push an update from a worker
75    ///
76    /// Behavior depends on the consistency model:
77    /// - BSP: buffers until all workers push, then applies via `barrier_sync`
78    /// - ASP: applies immediately
79    /// - SSP: applies if within staleness bound, otherwise buffers
80    pub fn push(&mut self, update: ParameterUpdate) -> CoreResult<()> {
81        let worker_id = update.worker_id;
82        if worker_id >= self.workers.len() {
83            return Err(CoreError::ValueError(ErrorContext::new(format!(
84                "Unknown worker ID: {worker_id}"
85            ))));
86        }
87        if !self.workers[worker_id].is_alive {
88            return Err(CoreError::ComputationError(ErrorContext::new(format!(
89                "Worker {worker_id} is not alive"
90            ))));
91        }
92
93        // Update worker version
94        self.workers[worker_id].version = update.version;
95
96        match &self.config.consistency {
97            ConsistencyModel::BSP => {
98                // Buffer the update
99                self.update_buffer
100                    .entry(update.key.clone())
101                    .or_default()
102                    .push(update);
103                self.bsp_pushed_workers[worker_id] = true;
104            }
105            ConsistencyModel::ASP => {
106                // Apply immediately
107                self.apply_single_update(&update)?;
108            }
109            ConsistencyModel::SSP { max_staleness } => {
110                let min_version = self
111                    .workers
112                    .iter()
113                    .filter(|w| w.is_alive)
114                    .map(|w| w.version)
115                    .min()
116                    .unwrap_or(0);
117                let staleness = update.version.saturating_sub(min_version) as usize;
118
119                if staleness <= *max_staleness {
120                    // Within staleness bound — apply immediately
121                    self.apply_single_update(&update)?;
122                } else {
123                    // Too stale — buffer until slower workers catch up
124                    self.update_buffer
125                        .entry(update.key.clone())
126                        .or_default()
127                        .push(update);
128                }
129            }
130        }
131        Ok(())
132    }
133
134    /// Pull current parameter values for a given key
135    pub fn pull(&self, key: &str, worker_id: usize) -> CoreResult<(Vec<f64>, u64)> {
136        if worker_id >= self.workers.len() {
137            return Err(CoreError::ValueError(ErrorContext::new(format!(
138                "Unknown worker ID: {worker_id}"
139            ))));
140        }
141        self.parameters
142            .get(key)
143            .cloned()
144            .ok_or_else(|| CoreError::ValueError(ErrorContext::new(format!("Unknown key: {key}"))))
145    }
146
147    /// Aggregate updates using the configured method
148    #[must_use]
149    pub fn aggregate_updates(updates: &[ParameterUpdate], method: &AggregationMethod) -> Vec<f64> {
150        if updates.is_empty() {
151            return Vec::new();
152        }
153        let dim = updates[0].values.len();
154
155        match method {
156            AggregationMethod::Mean => {
157                let mut sum = vec![0.0; dim];
158                for u in updates {
159                    for (s, v) in sum.iter_mut().zip(u.values.iter()) {
160                        *s += v;
161                    }
162                }
163                let n = updates.len() as f64;
164                sum.iter().map(|s| s / n).collect()
165            }
166            AggregationMethod::Sum => {
167                let mut sum = vec![0.0; dim];
168                for u in updates {
169                    for (s, v) in sum.iter_mut().zip(u.values.iter()) {
170                        *s += v;
171                    }
172                }
173                sum
174            }
175            AggregationMethod::WeightedMean { weights } => {
176                let mut weighted_sum = vec![0.0; dim];
177                let mut total_weight = 0.0;
178                for u in updates {
179                    let w = weights.get(u.worker_id).copied().unwrap_or(1.0);
180                    total_weight += w;
181                    for (s, v) in weighted_sum.iter_mut().zip(u.values.iter()) {
182                        *s += v * w;
183                    }
184                }
185                if total_weight.abs() < f64::EPSILON {
186                    weighted_sum
187                } else {
188                    weighted_sum.iter().map(|s| s / total_weight).collect()
189                }
190            }
191        }
192    }
193
194    /// BSP barrier synchronization
195    ///
196    /// Applies all buffered updates and increments the global version.
197    /// Returns an error if not all alive workers have pushed.
198    pub fn barrier_sync(&mut self) -> CoreResult<()> {
199        // Check that all alive workers have pushed
200        for w in &self.workers {
201            if w.is_alive
202                && !self
203                    .bsp_pushed_workers
204                    .get(w.worker_id)
205                    .copied()
206                    .unwrap_or(false)
207            {
208                return Err(CoreError::ComputationError(ErrorContext::new(format!(
209                    "BSP barrier: worker {} has not pushed yet",
210                    w.worker_id
211                ))));
212            }
213        }
214
215        // Apply all buffered updates
216        let keys: Vec<String> = self.update_buffer.keys().cloned().collect();
217        for key in &keys {
218            if let Some(updates) = self.update_buffer.get(key) {
219                let aggregated = Self::aggregate_updates(updates, &self.config.aggregation);
220                if !aggregated.is_empty() {
221                    let version = self.global_version + 1;
222                    self.parameters.insert(key.clone(), (aggregated, version));
223                }
224            }
225        }
226
227        // Clear buffer and reset BSP state
228        self.update_buffer.clear();
229        for flag in &mut self.bsp_pushed_workers {
230            *flag = false;
231        }
232        self.global_version += 1;
233
234        Ok(())
235    }
236
237    /// Create a checkpoint of the current server state
238    #[must_use]
239    pub fn checkpoint(&self) -> ServerCheckpoint {
240        ServerCheckpoint {
241            parameters: self.parameters.clone(),
242            version: self.global_version,
243            worker_states: self.workers.clone(),
244        }
245    }
246
247    /// Restore a parameter server from a checkpoint
248    pub fn restore(checkpoint: &ServerCheckpoint, config: ParamServerConfig) -> CoreResult<Self> {
249        let num_workers = config.num_workers;
250        Ok(Self {
251            config,
252            parameters: checkpoint.parameters.clone(),
253            workers: checkpoint.worker_states.clone(),
254            global_version: checkpoint.version,
255            update_buffer: HashMap::new(),
256            bsp_pushed_workers: vec![false; num_workers],
257        })
258    }
259
260    /// Get the current global version
261    #[must_use]
262    pub fn global_version(&self) -> u64 {
263        self.global_version
264    }
265
266    /// Get number of registered workers
267    #[must_use]
268    pub fn num_workers(&self) -> usize {
269        self.workers.len()
270    }
271
272    /// Get a reference to the worker states
273    #[must_use]
274    pub fn workers(&self) -> &[WorkerState] {
275        &self.workers
276    }
277
278    /// Get a mutable reference to the worker states
279    pub fn workers_mut(&mut self) -> &mut Vec<WorkerState> {
280        &mut self.workers
281    }
282
283    /// Get the configuration
284    #[must_use]
285    pub fn config(&self) -> &ParamServerConfig {
286        &self.config
287    }
288
289    /// Apply a single update directly to the parameter store
290    fn apply_single_update(&mut self, update: &ParameterUpdate) -> CoreResult<()> {
291        let entry = self
292            .parameters
293            .entry(update.key.clone())
294            .or_insert_with(|| (vec![0.0; update.values.len()], 0));
295
296        // For single-worker updates, just replace values and bump version
297        entry.0 = update.values.clone();
298        entry.1 = update.version;
299        Ok(())
300    }
301}
302
303#[cfg(test)]
304mod tests {
305    use super::*;
306
307    #[test]
308    fn test_register_workers() {
309        let config = ParamServerConfig::default();
310        let mut ps = ParameterServer::new(config);
311        let id0 = ps.register_worker();
312        let id1 = ps.register_worker();
313        assert_eq!(id0, 0);
314        assert_eq!(id1, 1);
315        assert_eq!(ps.num_workers(), 2);
316    }
317
318    #[test]
319    fn test_init_and_pull() {
320        let config = ParamServerConfig::default();
321        let mut ps = ParameterServer::new(config);
322        let wid = ps.register_worker();
323        ps.init_parameter("w1", vec![1.0, 2.0, 3.0]);
324        let (vals, ver) = ps.pull("w1", wid).expect("pull should succeed");
325        assert_eq!(vals, vec![1.0, 2.0, 3.0]);
326        assert_eq!(ver, 0);
327    }
328
329    #[test]
330    fn test_pull_unknown_key() {
331        let config = ParamServerConfig::default();
332        let mut ps = ParameterServer::new(config);
333        let wid = ps.register_worker();
334        let result = ps.pull("nonexistent", wid);
335        assert!(result.is_err());
336    }
337
338    #[test]
339    fn test_bsp_push_and_barrier() {
340        let config = ParamServerConfig {
341            num_workers: 2,
342            consistency: ConsistencyModel::BSP,
343            aggregation: AggregationMethod::Mean,
344            ..ParamServerConfig::default()
345        };
346        let mut ps = ParameterServer::new(config);
347        let w0 = ps.register_worker();
348        let w1 = ps.register_worker();
349        ps.init_parameter("w", vec![0.0, 0.0]);
350
351        // Worker 0 pushes
352        ps.push(ParameterUpdate {
353            key: "w".to_string(),
354            values: vec![2.0, 4.0],
355            worker_id: w0,
356            version: 1,
357        })
358        .expect("push w0");
359
360        // Barrier should fail (w1 hasn't pushed)
361        assert!(ps.barrier_sync().is_err());
362
363        // Worker 1 pushes
364        ps.push(ParameterUpdate {
365            key: "w".to_string(),
366            values: vec![4.0, 6.0],
367            worker_id: w1,
368            version: 1,
369        })
370        .expect("push w1");
371
372        // Barrier should succeed now
373        ps.barrier_sync().expect("barrier");
374
375        let (vals, ver) = ps.pull("w", w0).expect("pull after barrier");
376        // Mean of [2,4] and [4,6] = [3,5]
377        assert!((vals[0] - 3.0).abs() < f64::EPSILON);
378        assert!((vals[1] - 5.0).abs() < f64::EPSILON);
379        assert_eq!(ver, 1);
380    }
381
382    #[test]
383    fn test_asp_push() {
384        let config = ParamServerConfig {
385            num_workers: 1,
386            consistency: ConsistencyModel::ASP,
387            ..ParamServerConfig::default()
388        };
389        let mut ps = ParameterServer::new(config);
390        let w0 = ps.register_worker();
391        ps.init_parameter("p", vec![0.0]);
392
393        ps.push(ParameterUpdate {
394            key: "p".to_string(),
395            values: vec![42.0],
396            worker_id: w0,
397            version: 1,
398        })
399        .expect("asp push");
400
401        let (vals, _) = ps.pull("p", w0).expect("pull");
402        assert!((vals[0] - 42.0).abs() < f64::EPSILON);
403    }
404
405    #[test]
406    fn test_ssp_within_bound() {
407        let config = ParamServerConfig {
408            num_workers: 2,
409            consistency: ConsistencyModel::SSP { max_staleness: 2 },
410            ..ParamServerConfig::default()
411        };
412        let mut ps = ParameterServer::new(config);
413        let w0 = ps.register_worker();
414        let _w1 = ps.register_worker();
415
416        ps.push(ParameterUpdate {
417            key: "s".to_string(),
418            values: vec![10.0],
419            worker_id: w0,
420            version: 1,
421        })
422        .expect("ssp push within bound");
423
424        // Should be applied immediately since staleness = 1 - 0 = 1 <= 2
425        let (vals, _) = ps.pull("s", w0).expect("pull");
426        assert!((vals[0] - 10.0).abs() < f64::EPSILON);
427    }
428
429    #[test]
430    fn test_aggregate_sum() {
431        let updates = vec![
432            ParameterUpdate {
433                key: "k".into(),
434                values: vec![1.0, 2.0],
435                worker_id: 0,
436                version: 1,
437            },
438            ParameterUpdate {
439                key: "k".into(),
440                values: vec![3.0, 4.0],
441                worker_id: 1,
442                version: 1,
443            },
444        ];
445        let result = ParameterServer::aggregate_updates(&updates, &AggregationMethod::Sum);
446        assert!((result[0] - 4.0).abs() < f64::EPSILON);
447        assert!((result[1] - 6.0).abs() < f64::EPSILON);
448    }
449
450    #[test]
451    fn test_checkpoint_and_restore() {
452        let config = ParamServerConfig {
453            num_workers: 1,
454            consistency: ConsistencyModel::ASP,
455            ..ParamServerConfig::default()
456        };
457        let mut ps = ParameterServer::new(config.clone());
458        let w0 = ps.register_worker();
459        ps.init_parameter("x", vec![1.0, 2.0]);
460        ps.push(ParameterUpdate {
461            key: "x".to_string(),
462            values: vec![5.0, 6.0],
463            worker_id: w0,
464            version: 1,
465        })
466        .expect("push");
467
468        let cp = ps.checkpoint();
469        let restored = ParameterServer::restore(&cp, config).expect("restore");
470        let (vals, _) = restored.pull("x", w0).expect("pull from restored");
471        assert!((vals[0] - 5.0).abs() < f64::EPSILON);
472        assert!((vals[1] - 6.0).abs() < f64::EPSILON);
473    }
474}