summaryrefslogtreecommitdiff
path: root/models/model.py
diff options
context:
space:
mode:
Diffstat (limited to 'models/model.py')
-rw-r--r--models/model.py4
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)