-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathrun.py
More file actions
44 lines (40 loc) · 1.91 KB
/
Copy pathrun.py
File metadata and controls
44 lines (40 loc) · 1.91 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
import optimizers
import networks
from keras.datasets import mnist
import torch
if __name__ == "__main__":
N_data = 1000
N_iters = 1000
(X_train, y_train), (X_test, y_test) = mnist.load_data()
X_train, y_train = torch.tensor(X_train[:N_data,:,:], dtype=torch.float), torch.tensor(y_train[:N_data])
def loss_func(data):
loss_func = torch.nn.CrossEntropyLoss(reduction='none')
relu = torch.nn.ReLU()
return relu(loss_func(data[0], data[1])+torch.log(torch.tensor([0.5])))
print("Adam with summation\n--------")
sum_net = networks.SimpleNN()
adam_s = optimizers.adam_sum(sum_net, loss_func, 0.01)
comp_set_card = adam_s.train((X_train, y_train),N_iters)
X_test, y_test = torch.tensor(X_test, dtype=torch.float), torch.tensor(y_test)
acc = networks.eval(sum_net, (X_test, y_test))
print("Accuracy: {:.2f}%".format(acc*100))
print("Compression set size: {}/{}".format(comp_set_card, len(y_train)))
print("-------")
print("Adam with subgradient\n--------")
max_net = networks.SimpleNN()
adam_m = optimizers.adam_max(max_net, loss_func, 0.01)
comp_set_card = adam_m.train((X_train, y_train),N_iters)
X_test, y_test = torch.tensor(X_test, dtype=torch.float), torch.tensor(y_test)
acc = networks.eval(max_net, (X_test, y_test))
print("Accuracy: {:.2f}%".format(acc*100))
print("Compression set size: {}/{}".format(comp_set_card, len(y_train)))
print("-------")
print("Subsurface\n--------")
sub_net = networks.SimpleNN()
sub = optimizers.subsurface(sub_net, loss_func, 0.01)
comp_set_card = sub.train((X_train, y_train),N_iters)
X_test, y_test = torch.tensor(X_test, dtype=torch.float), torch.tensor(y_test)
acc = networks.eval(sub_net, (X_test, y_test))
print("Accuracy: {:.2f}%".format(acc*100))
print("Compression set size: {}/{}".format(comp_set_card, len(y_train)))
print("-------")