mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
Fix saving native AMP scaler state (#1777)
Saving was introduced in #1561.
This commit is contained in:
@@ -46,6 +46,9 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/).
|
||||
|
||||
- Fixed lr key name in case of param groups in LearningRateLogger ([#1719](https://github.com/PyTorchLightning/pytorch-lightning/pull/1719))
|
||||
|
||||
- Fixed saving native AMP scaler state (introduced in [#1561](https://github.com/PyTorchLightning/pytorch-lightning/pull/1561))
|
||||
|
||||
|
||||
## [0.7.5] - 2020-04-27
|
||||
|
||||
### Changed
|
||||
|
||||
@@ -338,8 +338,8 @@ class TrainerIOMixin(ABC):
|
||||
|
||||
checkpoint['state_dict'] = model.state_dict()
|
||||
|
||||
# restore native amp scaling
|
||||
if self.use_amp and self.use_native_amp and 'native_amp_scaling_state' in checkpoint:
|
||||
# save native amp scaling
|
||||
if self.use_amp and self.use_native_amp:
|
||||
checkpoint['native_amp_scaling_state'] = self.scaler.state_dict()
|
||||
|
||||
if hasattr(model, "hparams"):
|
||||
|
||||
Reference in New Issue
Block a user