MCPcopy Create free account
hub / github.com/Anoise/WTFlib / adjust_learning_rate

Function adjust_learning_rate

LDPS_Graph/utils/tools.py:9–39  ·  view source on GitHub ↗
(optimizer, scheduler, epoch, args, printout=True)

Source from the content-addressed store, hash-verified

7
8
9def 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
42class EarlyStopping:

Callers 1

trainMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected