| 28 | |
| 29 | |
| 30 | def adjust_lr(config, optimizer, step_count): |
| 31 | # adjust learning rate, step lr every lr_decay_steps |
| 32 | if step_count < config.lr_warm_step: |
| 33 | lr = config.lr_init * step_count / config.lr_warm_step |
| 34 | for param_group in optimizer.param_groups: |
| 35 | param_group['lr'] = lr |
| 36 | else: |
| 37 | lr = config.lr_init * config.lr_decay_rate ** ((step_count - config.lr_warm_step) // config.lr_decay_steps) |
| 38 | for param_group in optimizer.param_groups: |
| 39 | param_group['lr'] = lr |
| 40 | |
| 41 | return lr |
| 42 | |
| 43 | |
| 44 | def update_weights(model, batch, optimizer, replay_buffer, config, scaler, vis_result=False): |