Skip to main content

torsh_optim/
mixed_precision.rs

1//! Mixed precision support for optimizers
2//!
3//! Mixed precision training uses 16-bit floating point (fp16) for forward and backward passes
4//! while maintaining 32-bit floating point (fp32) master weights in the optimizer.
5//! This reduces memory usage and can improve training speed while maintaining model accuracy.
6
7use crate::{Optimizer, OptimizerResult, OptimizerState};
8use parking_lot::RwLock;
9use std::collections::HashMap;
10use std::sync::Arc;
11use torsh_core::dtype::DType;
12use torsh_core::error::Result;
13use torsh_core::DeviceType;
14use torsh_tensor::Tensor;
15
16/// Mixed precision configuration
17#[derive(Debug, Clone)]
18pub struct MixedPrecisionConfig {
19    /// Enable mixed precision training
20    pub enabled: bool,
21    /// Loss scaling factor to prevent gradient underflow
22    pub loss_scale: f32,
23    /// Whether to use dynamic loss scaling
24    pub dynamic_scale: bool,
25    /// Initial loss scale for dynamic scaling
26    pub init_scale: f32,
27    /// Factor to increase loss scale when no overflow is detected
28    pub scale_growth_factor: f32,
29    /// Number of steps to wait before increasing loss scale
30    pub scale_growth_interval: u32,
31    /// Factor to decrease loss scale when overflow is detected
32    pub backoff_factor: f32,
33    /// Maximum loss scale value
34    pub max_scale: f32,
35    /// Minimum loss scale value
36    pub min_scale: f32,
37}
38
39impl Default for MixedPrecisionConfig {
40    fn default() -> Self {
41        Self {
42            enabled: false,
43            loss_scale: 65536.0,
44            dynamic_scale: true,
45            init_scale: 65536.0,
46            scale_growth_factor: 2.0,
47            scale_growth_interval: 2000,
48            backoff_factor: 0.5,
49            max_scale: 2.0_f32.powi(24),
50            min_scale: 1.0,
51        }
52    }
53}
54
55/// Mixed precision optimizer wrapper
56pub struct MixedPrecisionOptimizer<O: Optimizer> {
57    optimizer: O,
58    config: MixedPrecisionConfig,
59    master_weights: HashMap<String, Tensor>,
60    loss_scaler: LossScaler,
61    overflow_detected: bool,
62}
63
64impl<O: Optimizer> MixedPrecisionOptimizer<O> {
65    /// Create a new mixed precision optimizer wrapper
66    pub fn new(optimizer: O, config: MixedPrecisionConfig) -> Self {
67        let loss_scaler = if config.dynamic_scale {
68            LossScaler::Dynamic(DynamicLossScaler::new(
69                config.init_scale,
70                config.scale_growth_factor,
71                config.scale_growth_interval,
72                config.backoff_factor,
73            ))
74        } else {
75            LossScaler::Static(StaticLossScaler::new(config.loss_scale))
76        };
77
78        Self {
79            optimizer,
80            config,
81            master_weights: HashMap::new(),
82            loss_scaler,
83            overflow_detected: false,
84        }
85    }
86
87    /// Create a mixed precision optimizer with default configuration
88    pub fn with_defaults(optimizer: O) -> Self {
89        let mut config = MixedPrecisionConfig::default();
90        config.enabled = true;
91        Self::new(optimizer, config)
92    }
93
94    /// Get the current loss scale
95    pub fn get_loss_scale(&self) -> f32 {
96        self.loss_scaler.get_scale()
97    }
98
99    /// Check if mixed precision is enabled
100    pub fn is_enabled(&self) -> bool {
101        self.config.enabled
102    }
103
104    /// Initialize master weights for mixed precision training
105    pub fn initialize_master_weights(&mut self, params: &[Arc<RwLock<Tensor>>]) -> Result<()> {
106        for param_arc in params {
107            let param = param_arc.read();
108            let param_id = format!("{:p}", param_arc.as_ref());
109
110            // Create fp32 master weight if parameter is fp16
111            if param.dtype() == DType::F16 {
112                let master_weight = param.to_dtype(DType::F32)?;
113                self.master_weights.insert(param_id, master_weight);
114            }
115        }
116        Ok(())
117    }
118
119    /// Scale loss for backward pass
120    pub fn scale_loss(&mut self, loss: &mut Tensor) -> Result<()> {
121        if self.config.enabled {
122            let scale = self.loss_scaler.get_scale();
123            loss.mul_scalar_(scale)?;
124        }
125        Ok(())
126    }
127
128    /// Unscale gradients before optimizer step
129    pub fn unscale_gradients(&mut self, params: &[Arc<RwLock<Tensor>>]) -> Result<bool> {
130        if !self.config.enabled {
131            return Ok(false);
132        }
133
134        let scale = self.loss_scaler.get_scale();
135        let inv_scale = 1.0 / scale;
136        let mut overflow_detected = false;
137
138        for param_arc in params {
139            let mut param = param_arc.write();
140            if let Some(grad) = param.grad_mut() {
141                // Check for inf/nan before unscaling
142                if self.has_inf_or_nan(grad)? {
143                    overflow_detected = true;
144                    break;
145                }
146
147                // Unscale gradient
148                grad.mul_scalar_(inv_scale)?;
149
150                // Check for inf/nan after unscaling
151                if self.has_inf_or_nan(grad)? {
152                    overflow_detected = true;
153                    break;
154                }
155            }
156        }
157
158        self.overflow_detected = overflow_detected;
159
160        if overflow_detected {
161            // Zero out gradients to prevent parameter updates
162            for param_arc in params {
163                let mut param = param_arc.write();
164                param.zero_grad();
165            }
166
167            // Update loss scaler
168            self.loss_scaler.on_overflow_detected();
169        } else {
170            self.loss_scaler.on_successful_step();
171        }
172
173        Ok(overflow_detected)
174    }
175
176    /// Update master weights from fp16 parameters
177    pub fn update_master_weights(&mut self, params: &[Arc<RwLock<Tensor>>]) -> Result<()> {
178        if !self.config.enabled {
179            return Ok(());
180        }
181
182        for param_arc in params {
183            let param = param_arc.read();
184            let param_id = format!("{:p}", param_arc.as_ref());
185
186            if let Some(master_weight) = self.master_weights.get_mut(&param_id) {
187                // Copy fp16 parameter to fp32 master weight
188                let param_fp32 = param.to_dtype(DType::F32)?;
189                *master_weight = param_fp32;
190            }
191        }
192
193        Ok(())
194    }
195
196    /// Copy master weights back to fp16 parameters
197    pub fn copy_master_to_params(&mut self, params: &[Arc<RwLock<Tensor>>]) -> Result<()> {
198        if !self.config.enabled {
199            return Ok(());
200        }
201
202        for param_arc in params {
203            let mut param = param_arc.write();
204            let param_id = format!("{:p}", param_arc.as_ref());
205
206            if let Some(master_weight) = self.master_weights.get(&param_id) {
207                // Copy fp32 master weight back to fp16 parameter
208                let param_fp16 = master_weight.to_dtype(param.dtype())?;
209                crate::param_update::assign(&mut param, &param_fp16)?;
210            }
211        }
212
213        Ok(())
214    }
215
216    /// Check if tensor contains inf or nan values
217    fn has_inf_or_nan(&self, tensor: &Tensor) -> Result<bool> {
218        // This is a simplified check - in a real implementation, this would
219        // use optimized kernels to check for inf/nan values
220        let data = tensor.to_vec()?;
221        Ok(data.iter().any(|&x| x.is_infinite() || x.is_nan()))
222    }
223
224    /// Get the underlying optimizer
225    pub fn inner(&self) -> &O {
226        &self.optimizer
227    }
228
229    /// Get the underlying optimizer mutably
230    pub fn inner_mut(&mut self) -> &mut O {
231        &mut self.optimizer
232    }
233}
234
235impl<O: Optimizer> Optimizer for MixedPrecisionOptimizer<O> {
236    fn step(&mut self) -> OptimizerResult<()> {
237        if self.overflow_detected {
238            // Skip optimizer step if overflow was detected
239            self.overflow_detected = false;
240            return Ok(());
241        }
242
243        self.optimizer.step()
244    }
245
246    fn zero_grad(&mut self) {
247        self.optimizer.zero_grad();
248    }
249
250    fn get_lr(&self) -> Vec<f32> {
251        self.optimizer.get_lr()
252    }
253
254    fn set_lr(&mut self, lr: f32) {
255        self.optimizer.set_lr(lr);
256    }
257
258    fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
259        self.optimizer.add_param_group(params, options);
260    }
261
262    fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
263        self.optimizer.parameters()
264    }
265
266    fn state_dict(&self) -> OptimizerResult<OptimizerState> {
267        self.optimizer.state_dict()
268    }
269
270    fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
271        self.optimizer.load_state_dict(state)
272    }
273}
274
275/// Loss scaling strategies
276enum LossScaler {
277    Static(StaticLossScaler),
278    Dynamic(DynamicLossScaler),
279}
280
281impl LossScaler {
282    fn get_scale(&self) -> f32 {
283        match self {
284            LossScaler::Static(scaler) => scaler.get_scale(),
285            LossScaler::Dynamic(scaler) => scaler.get_scale(),
286        }
287    }
288
289    fn on_overflow_detected(&mut self) {
290        match self {
291            LossScaler::Static(_) => {} // Static scaler doesn't change
292            LossScaler::Dynamic(scaler) => scaler.on_overflow_detected(),
293        }
294    }
295
296    fn on_successful_step(&mut self) {
297        match self {
298            LossScaler::Static(_) => {} // Static scaler doesn't change
299            LossScaler::Dynamic(scaler) => scaler.on_successful_step(),
300        }
301    }
302}
303
304/// Static loss scaler with fixed scale value
305struct StaticLossScaler {
306    scale: f32,
307}
308
309impl StaticLossScaler {
310    fn new(scale: f32) -> Self {
311        Self { scale }
312    }
313
314    fn get_scale(&self) -> f32 {
315        self.scale
316    }
317}
318
319/// Dynamic loss scaler that adjusts scale based on overflow detection
320struct DynamicLossScaler {
321    scale: f32,
322    growth_factor: f32,
323    growth_interval: u32,
324    backoff_factor: f32,
325    growth_tracker: u32,
326}
327
328impl DynamicLossScaler {
329    fn new(init_scale: f32, growth_factor: f32, growth_interval: u32, backoff_factor: f32) -> Self {
330        Self {
331            scale: init_scale,
332            growth_factor,
333            growth_interval,
334            backoff_factor,
335            growth_tracker: 0,
336        }
337    }
338
339    fn get_scale(&self) -> f32 {
340        self.scale
341    }
342
343    fn on_overflow_detected(&mut self) {
344        // Decrease scale and reset growth tracker
345        self.scale *= self.backoff_factor;
346        self.scale = self.scale.max(1.0); // Minimum scale of 1.0
347        self.growth_tracker = 0;
348    }
349
350    fn on_successful_step(&mut self) {
351        // Increment growth tracker
352        self.growth_tracker += 1;
353
354        // Increase scale if we've had enough successful steps
355        if self.growth_tracker >= self.growth_interval {
356            self.scale *= self.growth_factor;
357            self.scale = self.scale.min(2.0_f32.powi(24)); // Maximum scale
358            self.growth_tracker = 0;
359        }
360    }
361}
362
363/// Utilities for mixed precision training
364pub mod utils {
365    use super::*;
366
367    /// Check if the device supports mixed precision
368    pub fn supports_mixed_precision(device: &DeviceType) -> bool {
369        match device {
370            DeviceType::Cuda(_) => true,  // CUDA has good fp16 support
371            DeviceType::Metal(_) => true, // Metal supports fp16
372            DeviceType::Cpu => false,     // CPU fp16 support is limited
373            DeviceType::Wgpu(_) => false, // WebGPU fp16 support varies
374        }
375    }
376
377    /// Convert model parameters to fp16 for mixed precision training
378    pub fn convert_to_fp16(params: &[Arc<RwLock<Tensor>>]) -> Result<()> {
379        for param_arc in params {
380            let mut param = param_arc.write();
381            if param.dtype() == DType::F32 {
382                let param_fp16 = param.to_dtype(DType::F16)?;
383                *param = param_fp16;
384            }
385        }
386        Ok(())
387    }
388
389    /// Convert model parameters back to fp32
390    pub fn convert_to_fp32(params: &[Arc<RwLock<Tensor>>]) -> Result<()> {
391        for param_arc in params {
392            let mut param = param_arc.write();
393            if param.dtype() == DType::F16 {
394                let param_fp32 = param.to_dtype(DType::F32)?;
395                *param = param_fp32;
396            }
397        }
398        Ok(())
399    }
400
401    /// Estimate memory savings from using fp16
402    pub fn estimate_memory_savings(params: &[Arc<RwLock<Tensor>>]) -> (usize, usize, f64) {
403        let mut fp32_size = 0;
404        let mut fp16_size = 0;
405
406        for param_arc in params {
407            let param = param_arc.read();
408            let num_elements = param.numel();
409            fp32_size += num_elements * 4; // 4 bytes per fp32
410            fp16_size += num_elements * 2; // 2 bytes per fp16
411        }
412
413        let savings_ratio = 1.0 - (fp16_size as f64 / fp32_size as f64);
414        (fp32_size, fp16_size, savings_ratio)
415    }
416}
417
418/// Helper function to wrap any optimizer with mixed precision support
419pub fn with_mixed_precision<O: Optimizer>(
420    optimizer: O,
421    config: Option<MixedPrecisionConfig>,
422) -> MixedPrecisionOptimizer<O> {
423    match config {
424        Some(config) => MixedPrecisionOptimizer::new(optimizer, config),
425        None => MixedPrecisionOptimizer::with_defaults(optimizer),
426    }
427}
428
429#[cfg(test)]
430mod tests {
431    use super::*;
432    use crate::sgd::SGD;
433    use torsh_core::device::Device;
434    use torsh_tensor::creation;
435
436    #[test]
437    fn test_mixed_precision_config() {
438        let config = MixedPrecisionConfig::default();
439        assert!(!config.enabled);
440        assert_eq!(config.loss_scale, 65536.0);
441        assert!(config.dynamic_scale);
442    }
443
444    #[test]
445    fn test_static_loss_scaler() {
446        let scaler = StaticLossScaler::new(1024.0);
447        assert_eq!(scaler.get_scale(), 1024.0);
448    }
449
450    #[test]
451    fn test_dynamic_loss_scaler() {
452        let mut scaler = DynamicLossScaler::new(1024.0, 2.0, 2, 0.5);
453        assert_eq!(scaler.get_scale(), 1024.0);
454
455        // Simulate overflow
456        scaler.on_overflow_detected();
457        assert_eq!(scaler.get_scale(), 512.0);
458
459        // Simulate successful steps
460        scaler.on_successful_step();
461        scaler.on_successful_step();
462        assert_eq!(scaler.get_scale(), 1024.0);
463    }
464
465    #[test]
466    fn test_mixed_precision_optimizer_creation() {
467        let param = Arc::new(RwLock::new(creation::randn::<f32>(&[2, 3]).unwrap()));
468        let sgd = SGD::new(vec![param], 0.01, None, None, None, false);
469
470        let mp_optimizer = MixedPrecisionOptimizer::with_defaults(sgd);
471        assert!(mp_optimizer.is_enabled());
472        assert_eq!(mp_optimizer.get_loss_scale(), 65536.0);
473    }
474
475    #[test]
476    fn test_with_mixed_precision_helper() {
477        let param = Arc::new(RwLock::new(creation::randn::<f32>(&[2, 3]).unwrap()));
478        let sgd = SGD::new(vec![param], 0.01, None, None, None, false);
479
480        let mp_optimizer = with_mixed_precision(sgd, None);
481        assert!(mp_optimizer.is_enabled());
482    }
483
484    #[test]
485    fn test_memory_savings_estimation() {
486        let param1 = Arc::new(RwLock::new(creation::randn::<f32>(&[100, 100]).unwrap()));
487        let param2 = Arc::new(RwLock::new(creation::randn::<f32>(&[50, 50]).unwrap()));
488        let params = vec![param1, param2];
489
490        let (fp32_size, fp16_size, savings_ratio) = utils::estimate_memory_savings(&params);
491
492        assert_eq!(fp32_size, (10000 + 2500) * 4); // Total elements * 4 bytes
493        assert_eq!(fp16_size, (10000 + 2500) * 2); // Total elements * 2 bytes
494        assert!((savings_ratio - 0.5).abs() < 1e-6); // Should be ~50% savings
495    }
496
497    // Test disabled due to Device trait issues
498    // #[test]
499    // fn test_supports_mixed_precision() {
500    //     let cpu_device = Device::cpu();
501    //     assert!(!utils::supports_mixed_precision(&cpu_device));
502    // }
503}