1use 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#[derive(Debug, Clone)]
18pub struct MixedPrecisionConfig {
19 pub enabled: bool,
21 pub loss_scale: f32,
23 pub dynamic_scale: bool,
25 pub init_scale: f32,
27 pub scale_growth_factor: f32,
29 pub scale_growth_interval: u32,
31 pub backoff_factor: f32,
33 pub max_scale: f32,
35 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
55pub 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 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 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 pub fn get_loss_scale(&self) -> f32 {
96 self.loss_scaler.get_scale()
97 }
98
99 pub fn is_enabled(&self) -> bool {
101 self.config.enabled
102 }
103
104 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 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 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 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 if self.has_inf_or_nan(grad)? {
143 overflow_detected = true;
144 break;
145 }
146
147 grad.mul_scalar_(inv_scale)?;
149
150 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 for param_arc in params {
163 let mut param = param_arc.write();
164 param.zero_grad();
165 }
166
167 self.loss_scaler.on_overflow_detected();
169 } else {
170 self.loss_scaler.on_successful_step();
171 }
172
173 Ok(overflow_detected)
174 }
175
176 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(¶m_id) {
187 let param_fp32 = param.to_dtype(DType::F32)?;
189 *master_weight = param_fp32;
190 }
191 }
192
193 Ok(())
194 }
195
196 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(¶m_id) {
207 let param_fp16 = master_weight.to_dtype(param.dtype())?;
209 crate::param_update::assign(&mut param, ¶m_fp16)?;
210 }
211 }
212
213 Ok(())
214 }
215
216 fn has_inf_or_nan(&self, tensor: &Tensor) -> Result<bool> {
218 let data = tensor.to_vec()?;
221 Ok(data.iter().any(|&x| x.is_infinite() || x.is_nan()))
222 }
223
224 pub fn inner(&self) -> &O {
226 &self.optimizer
227 }
228
229 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 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
275enum 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(_) => {} LossScaler::Dynamic(scaler) => scaler.on_overflow_detected(),
293 }
294 }
295
296 fn on_successful_step(&mut self) {
297 match self {
298 LossScaler::Static(_) => {} LossScaler::Dynamic(scaler) => scaler.on_successful_step(),
300 }
301 }
302}
303
304struct 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
319struct 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 self.scale *= self.backoff_factor;
346 self.scale = self.scale.max(1.0); self.growth_tracker = 0;
348 }
349
350 fn on_successful_step(&mut self) {
351 self.growth_tracker += 1;
353
354 if self.growth_tracker >= self.growth_interval {
356 self.scale *= self.growth_factor;
357 self.scale = self.scale.min(2.0_f32.powi(24)); self.growth_tracker = 0;
359 }
360 }
361}
362
363pub mod utils {
365 use super::*;
366
367 pub fn supports_mixed_precision(device: &DeviceType) -> bool {
369 match device {
370 DeviceType::Cuda(_) => true, DeviceType::Metal(_) => true, DeviceType::Cpu => false, DeviceType::Wgpu(_) => false, }
375 }
376
377 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 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 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; fp16_size += num_elements * 2; }
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
418pub 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 scaler.on_overflow_detected();
457 assert_eq!(scaler.get_scale(), 512.0);
458
459 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(¶ms);
491
492 assert_eq!(fp32_size, (10000 + 2500) * 4); assert_eq!(fp16_size, (10000 + 2500) * 2); assert!((savings_ratio - 0.5).abs() < 1e-6); }
496
497 }