-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathopt.py
More file actions
83 lines (66 loc) · 2.37 KB
/
Copy pathopt.py
File metadata and controls
83 lines (66 loc) · 2.37 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
import torch.optim as optim
import torch
import sys
class OptimizerParameters(object):
def __init__(self, learning_rate=0.001, momentum=0.9, momentum2=0.999,\
epsilon=1e-8, weight_decay=0.0005, damp=0):
super(OptimizerParameters, self).__init__()
self.learning_rate = learning_rate
self.momentum = momentum
self.momentum2 = momentum2
self.epsilon = epsilon
self.damp = damp
self.weight_decay = weight_decay
def get_learning_rate(self):
return self.learning_rate
def get_momentum(self):
return self.momentum
def get_momentum2(self):
return self.momentum2
def get_epsilon(self):
return self.epsilon
def get_weight_decay(self):
return self.weight_decay
def get_damp(self):
return self.damp
def get_optimizer(opt_type, model_params, opt_params):
if opt_type == "adam":
return optim.Adam(model_params, \
lr=opt_params.get_learning_rate(), \
betas=(opt_params.get_momentum(), opt_params.get_momentum2()), \
eps=opt_params.get_epsilon(),
weight_decay = opt_params.get_weight_decay() \
)
elif opt_type == "sgd":
return optim.SGD(model_params, \
lr=opt_params.get_learning_rate(), \
momentum=opt_params.get_momentum(), \
weight_decay=opt_params.get_weight_decay(), \
dampening=opt_params.get_damp() \
)
else:
print("Error when initializing optimizer, {} is not a valid optimizer type.".format(opt_type), \
file=sys.stderr)
return None
def save_checkpoint(state, curr_epoch):
torch.save(state, './models/model_e%d.pth.tar' % (curr_epoch))
def adjust_learning_rate(schedule, optimizer, epoch, lr, gamma):
if epoch in schedule:
lr *= gamma
for param_group in optimizer.param_groups:
param_group['lr'] = lr
return optimizer
# Computes and stores the average and current value
class AverageMeter(object):
def __init__(self):
self.reset()
def reset(self):
self.val = torch.tensor(0.0)
self.avg = torch.tensor(0.0)
self.sum = torch.tensor(0.0)
self.count = torch.tensor(0.0)
def update(self, val, n=1):
self.val = val
self.sum += val * n
self.count += n
self.avg = self.sum / self.count