Set up the optimizer.
(param_groups, args)
| 173 | |
| 174 | |
| 175 | def get_optimizer(param_groups, args): |
| 176 | """Set up the optimizer.""" |
| 177 | if args.cpu_optimizer: |
| 178 | # Apex FusedAdam uses decoupled weight decay so use the same here |
| 179 | if args.cpu_torch_adam: |
| 180 | cpu_adam_optimizer = torch.optim.AdamW |
| 181 | else: |
| 182 | from deepspeed.ops.adam import DeepSpeedCPUAdam |
| 183 | cpu_adam_optimizer = DeepSpeedCPUAdam |
| 184 | optimizer = cpu_adam_optimizer(param_groups, |
| 185 | lr=args.lr, weight_decay=args.weight_decay) |
| 186 | else: |
| 187 | # Use FusedAdam. |
| 188 | if args.optimizer == 'adam': |
| 189 | optimizer = Adam(param_groups, |
| 190 | lr=args.lr, |
| 191 | weight_decay=args.weight_decay, |
| 192 | betas=(args.adam_beta1, args.adam_beta2), |
| 193 | eps=args.adam_eps) |
| 194 | elif args.optimizer == 'adafactor': |
| 195 | from transformers import Adafactor |
| 196 | optimizer = Adafactor(param_groups, lr=args.lr, relative_step=False, warmup_init=False) |
| 197 | else: |
| 198 | raise NotImplementedError |
| 199 | |
| 200 | print(f'Optimizer = {optimizer.__class__.__name__}') |
| 201 | if hasattr(args, "deepspeed") and args.deepspeed: |
| 202 | raise NotImplementedError |
| 203 | # fp16 wrapper is not required for DeepSpeed. |
| 204 | # return optimizer |
| 205 | |
| 206 | # Wrap into fp16 optimizer. |
| 207 | if args.fp16: |
| 208 | optimizer = FP16_Optimizer(optimizer, |
| 209 | static_loss_scale=args.loss_scale, |
| 210 | dynamic_loss_scale=args.dynamic_loss_scale, |
| 211 | dynamic_loss_args={ |
| 212 | 'scale_window': args.loss_scale_window, |
| 213 | 'min_scale': args.min_scale, |
| 214 | 'delayed_shift': args.hysteresis}) |
| 215 | |
| 216 | return optimizer |
| 217 | |
| 218 | |
| 219 | def get_learning_rate_scheduler(optimizer, args): |
no test coverage detected