MCPcopy Create free account
hub / github.com/NVIDIA/semantic-segmentation / get_optimizer

Function get_optimizer

loss/optimizer.py:43–98  ·  view source on GitHub ↗

Decide Optimizer (Adam or SGD)

(args, net)

Source from the content-addressed store, hash-verified

41
42
43def 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

Callers 1

mainFunction · 0.90

Calls 1

RAdamClass · 0.90

Tested by

no test coverage detected