Skip to main content

trustformers_optim/zero/
mod.rs

1//! ZeRO (Zero Redundancy Optimizer) Implementation for TrustformeRS
2//!
3//! ZeRO is a memory-efficient training technique that partitions optimizer states,
4//! gradients, and parameters across devices to reduce memory usage while maintaining
5//! training efficiency.
6//!
7//! Implements three stages:
8//! - Stage 1: Partition optimizer states
9//! - Stage 2: Partition optimizer states + gradients
10//! - Stage 3: Partition optimizer states + gradients + parameters
11
12pub 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/// ZeRO optimization stages
30#[derive(Debug, Clone, Copy, PartialEq, Eq)]
31pub enum ZeROImplementationStage {
32    /// Stage 1: Partition optimizer states only
33    Stage1,
34    /// Stage 2: Partition optimizer states + gradients
35    Stage2,
36    /// Stage 3: Partition optimizer states + gradients + parameters
37    Stage3,
38}
39
40/// Memory statistics for ZeRO optimization
41#[derive(Debug, Clone)]
42pub struct ZeROMemoryStats {
43    /// Memory saved by partitioning optimizer states
44    pub optimizer_memory_saved: usize,
45    /// Memory saved by partitioning gradients
46    pub gradient_memory_saved: usize,
47    /// Memory saved by partitioning parameters
48    pub parameter_memory_saved: usize,
49    /// Total memory saved
50    pub total_memory_saved: usize,
51    /// Memory overhead from communication buffers
52    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
78// ─── Flat partition helpers (no Tensor dependency) ───────────────────────────
79
80/// Partition optimizer states across `world_size` ranks.
81///
82/// `state` is a slice of optimizer state vectors (e.g., Adam m, v per parameter).
83/// Returns `Vec<Vec<Vec<f32>>>` where `result[rank][param_idx]` is that rank's slice
84/// of the optimizer state for parameter `param_idx`.
85pub 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
101/// Partition gradients across `world_size` ranks.
102///
103/// Returns `result[rank][param_idx]` = that rank's slice of grad for param `param_idx`.
104pub 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
120/// Partition parameters across `world_size` ranks.
121///
122/// Returns `result[rank][param_idx]` = that rank's slice of the parameter.
123pub 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
139/// Gather (reconstruct) parameters from their partitioned shards.
140///
141/// `partitioned[rank][param_idx]` = rank's shard for that param.
142/// Returns `Vec<Vec<f32>>` indexed by param_idx with full concatenated values.
143pub 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
159/// Calculate memory reduction ratio for a given ZeRO stage.
160///
161/// Returns the fraction of total baseline memory that is saved (0.0 to 1.0).
162/// - Stage 1: saves optimizer_bytes * (world_size - 1) / world_size
163/// - Stage 2: saves (optimizer_bytes + grad_bytes) * (world_size - 1) / world_size
164/// - Stage 3: saves all bytes * (world_size - 1) / world_size
165pub 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// ─── ZeroConfig ─────────────────────────────────────────────────────────────
191
192/// Simple configuration struct for ZeRO stage selection and validation.
193#[derive(Debug, Clone)]
194pub struct ZeroConfig {
195    /// ZeRO stage: 1, 2, or 3
196    pub stage: u8,
197    /// Number of distributed ranks
198    pub world_size: usize,
199    /// Overlap communication with computation
200    pub overlap_comm: bool,
201    /// Number of gradient elements per reduce bucket
202    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    /// Validate the configuration.
218    ///
219    /// Returns `Err` with a descriptive message if:
220    /// - `stage` is not in 1..=3
221    /// - `world_size` is 0
222    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    // ── Helper ───────────────────────────────────────────────────────────────
238
239    fn make_params(n: usize) -> Vec<f32> {
240        (0..n).map(|i| i as f32).collect()
241    }
242
243    // ── partition_optimizer_state ─────────────────────────────────────────────
244
245    #[test]
246    fn test_partition_optimizer_state_basic() {
247        // 1 state vec of 8 floats, world_size=4 → each rank gets 2
248        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        // verify values: rank 0 = [0,1], rank 1 = [2,3], ...
259        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        // 7 elements, world_size=3 → chunks of ceil(7/3)=3, ranks get [3, 3, 1]
268        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        // total elements = 7
275        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        // world_size=2, 3 state vecs of different lengths
282        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    // ── partition_gradients_flat ──────────────────────────────────────────────
301
302    #[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        // rank 0: first 4 of param0, first 2 of param1
317        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    // ── partition_parameters_flat ─────────────────────────────────────────────
332
333    #[test]
334    fn test_partition_parameters_basic() {
335        let params = vec![make_params(12)];
336        let partitioned = partition_parameters_flat(&params, 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        // Total elements across all ranks must equal original count
346        let params = vec![make_params(20)];
347        let partitioned = partition_parameters_flat(&params, 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(&params, 1);
356        assert_eq!(partitioned.len(), 1);
357        assert_eq!(partitioned[0][0], make_params(10));
358    }
359
360    // ── gather_parameters_flat ────────────────────────────────────────────────
361
362    #[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    // ── zero_stage_memory_reduction ───────────────────────────────────────────
390
391    #[test]
392    fn test_stage1_memory_reduction() {
393        // Stage 1 saves opt_bytes * (ws-1)/ws
394        // world_size=4: saves 3/4 of opt_bytes
395        let ratio = zero_stage_memory_reduction(1, 4, 1000, 1000, 1000);
396        // saved = 1000 * 0.75 = 750, total = 3000, ratio = 750/3000 = 0.25
397        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        // saved = (opt+grad) * 0.75 = 2000 * 0.75 = 1500, total=3000, ratio=0.5
408        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        // saved = 3000 * 0.75 = 2250, total=3000, ratio=0.75
419        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    // ── ZeroConfig validation ─────────────────────────────────────────────────
440
441    #[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    // ── ZeROMemoryStats ───────────────────────────────────────────────────────
494
495    #[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(&params, 8);
520        assert_eq!(partitioned.len(), 8);
521        // Each rank should hold 1000/8 = ceil(1000/8) = 125 elements per param
522        assert_eq!(partitioned[0][0].len(), 125);
523        // Gather should recover original
524        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}