| 7 | |
| 8 | |
| 9 | def adjust_learning_rate(optimizer, scheduler, epoch, args, printout=True): |
| 10 | # lr = args.learning_rate * (0.2 ** (epoch // 2)) |
| 11 | if args.lradj == 'type1': |
| 12 | lr_adjust = {epoch: args.learning_rate * (0.5 ** ((epoch - 1) // 1))} |
| 13 | elif args.lradj == 'type2': |
| 14 | lr_adjust = { |
| 15 | 2: 5e-5, 4: 1e-5, 6: 5e-6, 8: 1e-6, |
| 16 | 10: 5e-7, 15: 1e-7, 20: 5e-8 |
| 17 | } |
| 18 | elif args.lradj == 'type3': |
| 19 | lr_adjust = {epoch: args.learning_rate if epoch < 3 else args.learning_rate * (0.9 ** ((epoch - 3) // 1))} |
| 20 | elif args.lradj == 'constant': |
| 21 | lr_adjust = {epoch: args.learning_rate} |
| 22 | elif args.lradj == '3': |
| 23 | lr_adjust = {epoch: args.learning_rate if epoch < 10 else args.learning_rate*0.1} |
| 24 | elif args.lradj == '4': |
| 25 | lr_adjust = {epoch: args.learning_rate if epoch < 15 else args.learning_rate*0.1} |
| 26 | elif args.lradj == '5': |
| 27 | lr_adjust = {epoch: args.learning_rate if epoch < 25 else args.learning_rate*0.1} |
| 28 | elif args.lradj == '6': |
| 29 | lr_adjust = {epoch: args.learning_rate if epoch < 5 else args.learning_rate*0.1} |
| 30 | elif args.lradj in ['TST', 'Mvstgn']: |
| 31 | lr_adjust = {epoch: scheduler.get_last_lr()[0]} |
| 32 | else: |
| 33 | lr_adjust = {epoch: scheduler.get_last_lr()[0]} |
| 34 | |
| 35 | if epoch in lr_adjust.keys(): |
| 36 | lr = lr_adjust[epoch] |
| 37 | for param_group in optimizer.param_groups: |
| 38 | param_group['lr'] = lr |
| 39 | if printout: print(args.lradj,'=> Adjust updating learning rate to {}'.format(lr)) |
| 40 | |
| 41 | |
| 42 | class EarlyStopping: |