Decide Optimizer (Adam or SGD)
(args, net)
| 41 | |
| 42 | |
| 43 | def get_optimizer(args, net): |
| 44 | """ |
| 45 | Decide Optimizer (Adam or SGD) |
| 46 | """ |
| 47 | param_groups = net.parameters() |
| 48 | |
| 49 | if args.optimizer == 'sgd': |
| 50 | optimizer = optim.SGD(param_groups, |
| 51 | lr=args.lr, |
| 52 | weight_decay=args.weight_decay, |
| 53 | momentum=args.momentum, |
| 54 | nesterov=False) |
| 55 | elif args.optimizer == 'adam': |
| 56 | optimizer = optim.Adam(param_groups, |
| 57 | lr=args.lr, |
| 58 | weight_decay=args.weight_decay, |
| 59 | amsgrad=args.amsgrad) |
| 60 | elif args.optimizer == 'radam': |
| 61 | optimizer = RAdam(param_groups, |
| 62 | lr=args.lr, |
| 63 | weight_decay=args.weight_decay) |
| 64 | else: |
| 65 | raise ValueError('Not a valid optimizer') |
| 66 | |
| 67 | def poly_schd(epoch): |
| 68 | return math.pow(1 - epoch / args.max_epoch, args.poly_exp) |
| 69 | |
| 70 | def poly2_schd(epoch): |
| 71 | if epoch < args.poly_step: |
| 72 | poly_exp = args.poly_exp |
| 73 | else: |
| 74 | poly_exp = 2 * args.poly_exp |
| 75 | return math.pow(1 - epoch / args.max_epoch, poly_exp) |
| 76 | |
| 77 | if args.lr_schedule == 'scl-poly': |
| 78 | if cfg.REDUCE_BORDER_EPOCH == -1: |
| 79 | raise ValueError('ERROR Cannot Do Scale Poly') |
| 80 | |
| 81 | rescale_thresh = cfg.REDUCE_BORDER_EPOCH |
| 82 | scale_value = args.rescale |
| 83 | lambda1 = lambda epoch: \ |
| 84 | math.pow(1 - epoch / args.max_epoch, |
| 85 | args.poly_exp) if epoch < rescale_thresh else scale_value * math.pow( |
| 86 | 1 - (epoch - rescale_thresh) / (args.max_epoch - rescale_thresh), |
| 87 | args.repoly) |
| 88 | scheduler = optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=lambda1) |
| 89 | elif args.lr_schedule == 'poly2': |
| 90 | scheduler = optim.lr_scheduler.LambdaLR(optimizer, |
| 91 | lr_lambda=poly2_schd) |
| 92 | elif args.lr_schedule == 'poly': |
| 93 | scheduler = optim.lr_scheduler.LambdaLR(optimizer, |
| 94 | lr_lambda=poly_schd) |
| 95 | else: |
| 96 | raise ValueError('unknown lr schedule {}'.format(args.lr_schedule)) |
| 97 | |
| 98 | return optimizer, scheduler |
| 99 | |
| 100 |