-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathparameter_search.py
More file actions
23 lines (19 loc) · 924 Bytes
/
Copy pathparameter_search.py
File metadata and controls
23 lines (19 loc) · 924 Bytes
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
import torch.optim as optim
lr = [0.1, 0.01, 0.001]
optimizers = ["adadelta", "adam", "sgd"]
momentum = [0, 0.5, 0.9]
def find_optimizer(lr=lr, optimizers=optimizers, model=None):
for i in range(len(optimizers)):
if optimizers[i] == "adadelta":
for j in range(len(lr)):
print(lr[j], "adad")
yield optim.Adadelta(model.parameters(), lr=lr[j]), ("adadelta", str(lr[j]))
elif optimizers[i] == "adam":
for j in range(len(lr)):
print(lr[j], "adam")
yield optim.Adam(model.parameters(), lr=lr[j]), ("adam", str(lr[j]))
elif optimizers[i]=="sgd":
for j in range(len(lr)):
for z in range(len(momentum)):
print(lr[j], momentum[z], "sgd")
yield optim.SGD(model.parameters(), lr=lr[j], momentum=momentum[z]), ("sgd", str(lr[j]), str(momentum[z]))