scirs2_core/distributed/param_server/
server.rs1use std::collections::HashMap;
7
8use crate::error::{CoreError, CoreResult, ErrorContext};
9
10use super::types::{
11 AggregationMethod, ConsistencyModel, ParamServerConfig, ParameterUpdate, WorkerState,
12};
13
14#[derive(Debug, Clone)]
16pub struct ServerCheckpoint {
17 pub parameters: HashMap<String, (Vec<f64>, u64)>,
19 pub version: u64,
21 pub worker_states: Vec<WorkerState>,
23}
24
25#[derive(Debug)]
27pub struct ParameterServer {
28 config: ParamServerConfig,
30 parameters: HashMap<String, (Vec<f64>, u64)>,
32 workers: Vec<WorkerState>,
34 global_version: u64,
36 update_buffer: HashMap<String, Vec<ParameterUpdate>>,
38 bsp_pushed_workers: Vec<bool>,
40}
41
42impl ParameterServer {
43 #[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 pub fn register_worker(&mut self) -> usize {
59 let worker_id = self.workers.len();
60 self.workers.push(WorkerState::new(worker_id));
61 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 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 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 self.workers[worker_id].version = update.version;
95
96 match &self.config.consistency {
97 ConsistencyModel::BSP => {
98 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 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 self.apply_single_update(&update)?;
122 } else {
123 self.update_buffer
125 .entry(update.key.clone())
126 .or_default()
127 .push(update);
128 }
129 }
130 }
131 Ok(())
132 }
133
134 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 #[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 pub fn barrier_sync(&mut self) -> CoreResult<()> {
199 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 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 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 #[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 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 #[must_use]
262 pub fn global_version(&self) -> u64 {
263 self.global_version
264 }
265
266 #[must_use]
268 pub fn num_workers(&self) -> usize {
269 self.workers.len()
270 }
271
272 #[must_use]
274 pub fn workers(&self) -> &[WorkerState] {
275 &self.workers
276 }
277
278 pub fn workers_mut(&mut self) -> &mut Vec<WorkerState> {
280 &mut self.workers
281 }
282
283 #[must_use]
285 pub fn config(&self) -> &ParamServerConfig {
286 &self.config
287 }
288
289 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 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 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 assert!(ps.barrier_sync().is_err());
362
363 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 ps.barrier_sync().expect("barrier");
374
375 let (vals, ver) = ps.pull("w", w0).expect("pull after barrier");
376 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 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}