MCPcopy Create free account
hub / github.com/M-Nauta/ProtoTree / get_optimizer

Function get_optimizer

util/args.py:190–255  ·  view source on GitHub ↗

Construct the optimizer as dictated by the parsed arguments :param tree: The tree that should be optimized :param args: Parsed arguments containing hyperparameters. The '--optimizer' argument specifies which type of optimizer will be used. Optimizer specific argumen

(tree, args: argparse.Namespace)

Source from the content-addressed store, hash-verified

188 return args
189
190def get_optimizer(tree, args: argparse.Namespace) -> torch.optim.Optimizer:
191 """
192 Construct the optimizer as dictated by the parsed arguments
193 :param tree: The tree that should be optimized
194 :param args: Parsed arguments containing hyperparameters. The '--optimizer' argument specifies which type of
195 optimizer will be used. Optimizer specific arguments (such as learning rate and momentum) can be passed
196 this way as well
197 :return: the optimizer corresponding to the parsed arguments, parameter set that can be frozen, and parameter set of the net that will be trained
198 """
199
200 optim_type = args.optimizer
201 #create parameter groups
202 params_to_freeze = []
203 params_to_train = []
204
205 dist_params = []
206 for name,param in tree.named_parameters():
207 if 'dist_params' in name:
208 dist_params.append(param)
209 # set up optimizer
210 if 'resnet50_inat' in args.net or ('resnet50' in args.net and args.dataset=='CARS'): #to reproduce experimental results
211 # freeze resnet50 except last convolutional layer
212 for name,param in tree._net.named_parameters():
213 if 'layer4.2' not in name:
214 params_to_freeze.append(param)
215 else:
216 params_to_train.append(param)
217
218 if optim_type == 'SGD':
219 paramlist = [
220 {"params": params_to_freeze, "lr": args.lr_net, "weight_decay_rate": args.weight_decay, "momentum": args.momentum},
221 {"params": params_to_train, "lr": args.lr_block, "weight_decay_rate": args.weight_decay,"momentum": args.momentum},
222 {"params": tree._add_on.parameters(), "lr": args.lr_block, "weight_decay_rate": args.weight_decay,"momentum": args.momentum},
223 {"params": tree.prototype_layer.parameters(), "lr": args.lr,"weight_decay_rate": 0,"momentum": 0}]
224 if args.disable_derivative_free_leaf_optim:
225 paramlist.append({"params": dist_params, "lr": args.lr_pi, "weight_decay_rate": 0})
226 else:
227 paramlist = [
228 {"params": params_to_freeze, "lr": args.lr_net, "weight_decay_rate": args.weight_decay},
229 {"params": params_to_train, "lr": args.lr_block, "weight_decay_rate": args.weight_decay},
230 {"params": tree._add_on.parameters(), "lr": args.lr_block, "weight_decay_rate": args.weight_decay},
231 {"params": tree.prototype_layer.parameters(), "lr": args.lr,"weight_decay_rate": 0}]
232
233 if args.disable_derivative_free_leaf_optim:
234 paramlist.append({"params": dist_params, "lr": args.lr_pi, "weight_decay_rate": 0})
235
236 else: #other network architectures
237 for name,param in tree._net.named_parameters():
238 params_to_freeze.append(param)
239 paramlist = [
240 {"params": params_to_freeze, "lr": args.lr_net, "weight_decay_rate": args.weight_decay},
241 {"params": tree._add_on.parameters(), "lr": args.lr_block, "weight_decay_rate": args.weight_decay},
242 {"params": tree.prototype_layer.parameters(), "lr": args.lr,"weight_decay_rate": 0}]
243 if args.disable_derivative_free_leaf_optim:
244 paramlist.append({"params": dist_params, "lr": args.lr_pi, "weight_decay_rate": 0})
245
246 if optim_type == 'SGD':
247 return torch.optim.SGD(paramlist,

Callers 1

run_treeFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected