1pub mod zero_optimizer;
13pub mod zero_stage1;
14pub mod zero_stage2;
15pub mod zero_stage3;
16pub mod zero_stage3_overlap;
17pub mod zero_utils;
18
19pub use zero_optimizer::{ZeROConfig, ZeROOptimizer, ZeROStage};
20pub use zero_stage1::ZeROStage1;
21pub use zero_stage2::ZeROStage2;
22pub use zero_stage3::ZeROStage3;
23pub use zero_utils::{
24 all_gather_gradients, gather_parameters, gather_shards, partition_gradients,
25 partition_parameters, reduce_scatter_gradients, shard_range, slice_flat, GradientBuffer,
26 ParameterGroup, ParameterPartition, ZeROState,
27};
28
29#[derive(Debug, Clone, Copy, PartialEq, Eq)]
31pub enum ZeROImplementationStage {
32 Stage1,
34 Stage2,
36 Stage3,
38}
39
40#[derive(Debug, Clone)]
42pub struct ZeROMemoryStats {
43 pub optimizer_memory_saved: usize,
45 pub gradient_memory_saved: usize,
47 pub parameter_memory_saved: usize,
49 pub total_memory_saved: usize,
51 pub communication_overhead: usize,
53}
54
55impl Default for ZeROMemoryStats {
56 fn default() -> Self {
57 Self::new()
58 }
59}
60
61impl ZeROMemoryStats {
62 pub fn new() -> Self {
63 Self {
64 optimizer_memory_saved: 0,
65 gradient_memory_saved: 0,
66 parameter_memory_saved: 0,
67 total_memory_saved: 0,
68 communication_overhead: 0,
69 }
70 }
71
72 pub fn update_totals(&mut self) {
73 self.total_memory_saved =
74 self.optimizer_memory_saved + self.gradient_memory_saved + self.parameter_memory_saved;
75 }
76}
77
78pub fn partition_optimizer_state(state: &[Vec<f32>], world_size: usize) -> Vec<Vec<Vec<f32>>> {
86 assert!(world_size > 0, "world_size must be > 0");
87 let mut result: Vec<Vec<Vec<f32>>> = vec![Vec::new(); world_size];
88 for param_state in state {
89 let total = param_state.len();
90 let chunk_size = total.div_ceil(world_size);
91 for rank in 0..world_size {
92 let start = rank * chunk_size;
93 let end = (start + chunk_size).min(total);
94 let shard = if start < total { param_state[start..end].to_vec() } else { Vec::new() };
95 result[rank].push(shard);
96 }
97 }
98 result
99}
100
101pub fn partition_gradients_flat(grads: &[Vec<f32>], world_size: usize) -> Vec<Vec<Vec<f32>>> {
105 assert!(world_size > 0, "world_size must be > 0");
106 let mut result: Vec<Vec<Vec<f32>>> = vec![Vec::new(); world_size];
107 for grad in grads {
108 let total = grad.len();
109 let chunk_size = total.div_ceil(world_size);
110 for rank in 0..world_size {
111 let start = rank * chunk_size;
112 let end = (start + chunk_size).min(total);
113 let shard = if start < total { grad[start..end].to_vec() } else { Vec::new() };
114 result[rank].push(shard);
115 }
116 }
117 result
118}
119
120pub fn partition_parameters_flat(params: &[Vec<f32>], world_size: usize) -> Vec<Vec<Vec<f32>>> {
124 assert!(world_size > 0, "world_size must be > 0");
125 let mut result: Vec<Vec<Vec<f32>>> = vec![Vec::new(); world_size];
126 for param in params {
127 let total = param.len();
128 let chunk_size = total.div_ceil(world_size);
129 for rank in 0..world_size {
130 let start = rank * chunk_size;
131 let end = (start + chunk_size).min(total);
132 let shard = if start < total { param[start..end].to_vec() } else { Vec::new() };
133 result[rank].push(shard);
134 }
135 }
136 result
137}
138
139pub fn gather_parameters_flat(partitioned: &[Vec<Vec<f32>>]) -> Vec<Vec<f32>> {
144 if partitioned.is_empty() {
145 return Vec::new();
146 }
147 let num_params = partitioned[0].len();
148 let mut result: Vec<Vec<f32>> = vec![Vec::new(); num_params];
149 for rank_data in partitioned {
150 for (param_idx, shard) in rank_data.iter().enumerate() {
151 if param_idx < result.len() {
152 result[param_idx].extend_from_slice(shard);
153 }
154 }
155 }
156 result
157}
158
159pub fn zero_stage_memory_reduction(
166 stage: u8,
167 world_size: usize,
168 param_bytes: usize,
169 grad_bytes: usize,
170 opt_bytes: usize,
171) -> f32 {
172 if world_size <= 1 {
173 return 0.0;
174 }
175 let total_bytes = (param_bytes + grad_bytes + opt_bytes) as f32;
176 if total_bytes == 0.0 {
177 return 0.0;
178 }
179 let ws = world_size as f32;
180 let save_fraction = (ws - 1.0) / ws;
181 let saved_bytes = match stage {
182 1 => opt_bytes as f32 * save_fraction,
183 2 => (opt_bytes + grad_bytes) as f32 * save_fraction,
184 3 => (param_bytes + grad_bytes + opt_bytes) as f32 * save_fraction,
185 _ => 0.0,
186 };
187 saved_bytes / total_bytes
188}
189
190#[derive(Debug, Clone)]
194pub struct ZeroConfig {
195 pub stage: u8,
197 pub world_size: usize,
199 pub overlap_comm: bool,
201 pub reduce_bucket_size: usize,
203}
204
205impl Default for ZeroConfig {
206 fn default() -> Self {
207 Self {
208 stage: 1,
209 world_size: 1,
210 overlap_comm: true,
211 reduce_bucket_size: 500_000_000,
212 }
213 }
214}
215
216impl ZeroConfig {
217 pub fn validate(&self) -> Result<(), String> {
223 if self.stage == 0 || self.stage > 3 {
224 return Err(format!("ZeRO stage must be 1, 2, or 3; got {}", self.stage));
225 }
226 if self.world_size == 0 {
227 return Err("world_size must be >= 1".to_string());
228 }
229 Ok(())
230 }
231}
232
233#[cfg(test)]
234mod tests {
235 use super::*;
236
237 fn make_params(n: usize) -> Vec<f32> {
240 (0..n).map(|i| i as f32).collect()
241 }
242
243 #[test]
246 fn test_partition_optimizer_state_basic() {
247 let state = vec![make_params(8)];
249 let partitioned = partition_optimizer_state(&state, 4);
250 assert_eq!(partitioned.len(), 4);
251 for rank in 0..4 {
252 assert_eq!(
253 partitioned[rank][0].len(),
254 2,
255 "rank {rank} should have 2 elements"
256 );
257 }
258 assert_eq!(partitioned[0][0], vec![0.0, 1.0]);
260 assert_eq!(partitioned[1][0], vec![2.0, 3.0]);
261 assert_eq!(partitioned[2][0], vec![4.0, 5.0]);
262 assert_eq!(partitioned[3][0], vec![6.0, 7.0]);
263 }
264
265 #[test]
266 fn test_partition_optimizer_state_uneven() {
267 let state = vec![make_params(7)];
269 let partitioned = partition_optimizer_state(&state, 3);
270 assert_eq!(partitioned.len(), 3);
271 assert_eq!(partitioned[0][0].len(), 3);
272 assert_eq!(partitioned[1][0].len(), 3);
273 assert_eq!(partitioned[2][0].len(), 1);
274 let total: usize = partitioned.iter().map(|r| r[0].len()).sum();
276 assert_eq!(total, 7);
277 }
278
279 #[test]
280 fn test_partition_optimizer_state_multiple_states() {
281 let state = vec![make_params(4), make_params(6), make_params(2)];
283 let partitioned = partition_optimizer_state(&state, 2);
284 assert_eq!(partitioned.len(), 2);
285 for rank_data in &partitioned {
286 assert_eq!(rank_data.len(), 3, "each rank should have 3 param states");
287 }
288 }
289
290 #[test]
291 fn test_partition_optimizer_state_rank_sizes_sum_to_original() {
292 let state = vec![make_params(10), make_params(7)];
293 let partitioned = partition_optimizer_state(&state, 4);
294 for param_idx in 0..2 {
295 let total: usize = partitioned.iter().map(|r| r[param_idx].len()).sum();
296 assert_eq!(total, state[param_idx].len());
297 }
298 }
299
300 #[test]
303 fn test_partition_gradients_basic() {
304 let grads = vec![make_params(16)];
305 let partitioned = partition_gradients_flat(&grads, 4);
306 assert_eq!(partitioned.len(), 4);
307 for rank in 0..4 {
308 assert_eq!(partitioned[rank][0].len(), 4);
309 }
310 }
311
312 #[test]
313 fn test_partition_gradients_multi() {
314 let grads = vec![make_params(8), make_params(4)];
315 let partitioned = partition_gradients_flat(&grads, 2);
316 assert_eq!(partitioned[0][0].len(), 4);
318 assert_eq!(partitioned[0][1].len(), 2);
319 }
320
321 #[test]
322 fn test_partition_gradients_size_check() {
323 let grads = vec![make_params(9), make_params(5)];
324 let partitioned = partition_gradients_flat(&grads, 3);
325 for (param_idx, original) in grads.iter().enumerate() {
326 let total: usize = partitioned.iter().map(|r| r[param_idx].len()).sum();
327 assert_eq!(total, original.len());
328 }
329 }
330
331 #[test]
334 fn test_partition_parameters_basic() {
335 let params = vec![make_params(12)];
336 let partitioned = partition_parameters_flat(¶ms, 4);
337 assert_eq!(partitioned.len(), 4);
338 for rank in 0..4 {
339 assert_eq!(partitioned[rank][0].len(), 3);
340 }
341 }
342
343 #[test]
344 fn test_partition_parameters_no_duplicate() {
345 let params = vec![make_params(20)];
347 let partitioned = partition_parameters_flat(¶ms, 4);
348 let total: usize = partitioned.iter().map(|r| r[0].len()).sum();
349 assert_eq!(total, 20);
350 }
351
352 #[test]
353 fn test_partition_parameters_world_size_1() {
354 let params = vec![make_params(10)];
355 let partitioned = partition_parameters_flat(¶ms, 1);
356 assert_eq!(partitioned.len(), 1);
357 assert_eq!(partitioned[0][0], make_params(10));
358 }
359
360 #[test]
363 fn test_gather_is_inverse_of_partition() {
364 let original = vec![make_params(12), make_params(8)];
365 let partitioned = partition_parameters_flat(&original, 4);
366 let gathered = gather_parameters_flat(&partitioned);
367 assert_eq!(gathered.len(), original.len());
368 for (idx, orig) in original.iter().enumerate() {
369 assert_eq!(&gathered[idx], orig, "param {idx} mismatch after gather");
370 }
371 }
372
373 #[test]
374 fn test_gather_inverse_uneven() {
375 let original = vec![make_params(7), make_params(11)];
376 let partitioned = partition_parameters_flat(&original, 3);
377 let gathered = gather_parameters_flat(&partitioned);
378 for (idx, orig) in original.iter().enumerate() {
379 assert_eq!(&gathered[idx], orig);
380 }
381 }
382
383 #[test]
384 fn test_gather_empty() {
385 let gathered = gather_parameters_flat(&[]);
386 assert!(gathered.is_empty());
387 }
388
389 #[test]
392 fn test_stage1_memory_reduction() {
393 let ratio = zero_stage_memory_reduction(1, 4, 1000, 1000, 1000);
396 let expected = (1000.0f32 * 0.75) / 3000.0;
398 assert!(
399 (ratio - expected).abs() < 1e-5,
400 "got {ratio}, expected {expected}"
401 );
402 }
403
404 #[test]
405 fn test_stage2_memory_reduction() {
406 let ratio = zero_stage_memory_reduction(2, 4, 1000, 1000, 1000);
407 let expected = (2000.0f32 * 0.75) / 3000.0;
409 assert!(
410 (ratio - expected).abs() < 1e-5,
411 "got {ratio}, expected {expected}"
412 );
413 }
414
415 #[test]
416 fn test_stage3_memory_reduction() {
417 let ratio = zero_stage_memory_reduction(3, 4, 1000, 1000, 1000);
418 let expected = 3000.0f32 * 0.75 / 3000.0;
420 assert!(
421 (ratio - expected).abs() < 1e-5,
422 "got {ratio}, expected {expected}"
423 );
424 }
425
426 #[test]
427 fn test_memory_reduction_world_size_1() {
428 let ratio = zero_stage_memory_reduction(3, 1, 1000, 1000, 1000);
429 assert_eq!(ratio, 0.0);
430 }
431
432 #[test]
433 fn test_memory_reduction_stage3_is_greater_than_stage1() {
434 let r1 = zero_stage_memory_reduction(1, 4, 1000, 1000, 1000);
435 let r3 = zero_stage_memory_reduction(3, 4, 1000, 1000, 1000);
436 assert!(r3 > r1, "stage3 should save more than stage1");
437 }
438
439 #[test]
442 fn test_zero_config_valid() {
443 let cfg = ZeroConfig {
444 stage: 2,
445 world_size: 4,
446 ..Default::default()
447 };
448 assert!(cfg.validate().is_ok());
449 }
450
451 #[test]
452 fn test_zero_config_invalid_stage_zero() {
453 let cfg = ZeroConfig {
454 stage: 0,
455 world_size: 4,
456 ..Default::default()
457 };
458 assert!(cfg.validate().is_err());
459 }
460
461 #[test]
462 fn test_zero_config_invalid_stage_four() {
463 let cfg = ZeroConfig {
464 stage: 4,
465 world_size: 4,
466 ..Default::default()
467 };
468 assert!(cfg.validate().is_err());
469 }
470
471 #[test]
472 fn test_zero_config_invalid_world_size() {
473 let cfg = ZeroConfig {
474 stage: 1,
475 world_size: 0,
476 ..Default::default()
477 };
478 assert!(cfg.validate().is_err());
479 }
480
481 #[test]
482 fn test_zero_config_all_stages_valid() {
483 for stage in 1u8..=3 {
484 let cfg = ZeroConfig {
485 stage,
486 world_size: 8,
487 ..Default::default()
488 };
489 assert!(cfg.validate().is_ok(), "stage {stage} should be valid");
490 }
491 }
492
493 #[test]
496 fn test_zero_memory_stats_new() {
497 let stats = ZeROMemoryStats::new();
498 assert_eq!(stats.optimizer_memory_saved, 0);
499 assert_eq!(stats.gradient_memory_saved, 0);
500 assert_eq!(stats.parameter_memory_saved, 0);
501 assert_eq!(stats.total_memory_saved, 0);
502 assert_eq!(stats.communication_overhead, 0);
503 }
504
505 #[test]
506 fn test_zero_memory_stats_update_totals() {
507 let mut stats = ZeROMemoryStats::new();
508 stats.optimizer_memory_saved = 100;
509 stats.gradient_memory_saved = 200;
510 stats.parameter_memory_saved = 300;
511 stats.update_totals();
512 assert_eq!(stats.total_memory_saved, 600);
513 }
514
515 #[test]
516 fn test_partition_large_vectors() {
517 let params: Vec<Vec<f32>> =
518 (0..5).map(|p| (0..1000).map(|i| (p * 1000 + i) as f32).collect()).collect();
519 let partitioned = partition_parameters_flat(¶ms, 8);
520 assert_eq!(partitioned.len(), 8);
521 assert_eq!(partitioned[0][0].len(), 125);
523 let gathered = gather_parameters_flat(&partitioned);
525 for (idx, orig) in params.iter().enumerate() {
526 assert_eq!(&gathered[idx], orig, "param {idx} mismatch");
527 }
528 }
529}