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)
| 188 | return args |
| 189 | |
| 190 | def 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, |