(model, configs)
| 5 | |
| 6 | |
| 7 | def get_optimizer(model, configs): |
| 8 | optim_configs = {k: v for k, v in configs.items() if k != '_name_'} |
| 9 | if configs['_name_'] == 'adamw': |
| 10 | return torch.optim.AdamW(model.parameters(), **optim_configs) |
| 11 | elif configs['_name_'] == 'sgd': |
| 12 | return torch.optim.SGD(model.parameters(), **optim_configs) |
| 13 | elif configs['_name_'] == 'adam': |
| 14 | return torch.optim.Adam(model.parameters(), **optim_configs) |
| 15 | |
| 16 | |
| 17 | def get_scheduler(model, optimizer, configs): |