Get optimizer.
(args, params)
| 79 | |
| 80 | |
| 81 | def get_optimizer(args, params): |
| 82 | """ |
| 83 | Get optimizer. |
| 84 | """ |
| 85 | if args.optimizer == 'adam': |
| 86 | return optim.Adam(params, lr=args.lr, weight_decay=args.weight_decay) |
| 87 | elif args.optimizer == 'adamw': |
| 88 | return optim.AdamW(params, lr=args.lr, weight_decay=args.weight_decay) |
| 89 | elif args.optimizer == 'sgd': |
| 90 | return optim.SGD(params, lr=args.lr, momentum=args.momentum, nesterov=True, weight_decay=args.weight_decay) |
| 91 | else: |
| 92 | raise ValueError(f'Optimizer {args.optimizer} not available.') |
| 93 | |
| 94 | |
| 95 | def get_scheduler(args, optimizer: torch.optim): |
no outgoing calls
no test coverage detected