diff options
-rw-r--r-- | models/model.py | 4 |
1 files changed, 2 insertions, 2 deletions
diff --git a/models/model.py b/models/model.py index aad0a99..b86a050 100644 --- a/models/model.py +++ b/models/model.py @@ -118,8 +118,8 @@ class Model: # Prepare for optimizer and scheduler optim_hp = self.hp.get('optimizer', {}) # Scale learning rate to world size - if optim_hp['lr']: - optim_hp['lr'] *= xm.xrt_world_size() + lr = optim_hp.get('lr', '1-e3') + optim_hp['lr'] = lr * xm.xrt_world_size() sched_hp = self.hp.get('scheduler', {}) device = xm.xla_device() rgb_pn = wrapped_rgb_pn.to(device) |