-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdemo.py
More file actions
117 lines (97 loc) · 3.66 KB
/
Copy pathdemo.py
File metadata and controls
117 lines (97 loc) · 3.66 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
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
import matplotlib.pyplot as plt
import numpy as np
import torch
from torchvision import datasets, transforms
state_space_model_type = "complex" # or "real"
optimizer = "adam" # or "psgd"
print(f"Domain of state vectors: {state_space_model_type}")
print(f"Optimizer: {optimizer}\n")
if state_space_model_type == "complex":
from state_space_models import ComplexStateSpaceModel as SSM
increase_state_size = 1
else:
from state_space_models import RealStateSpaceModel as SSM
increase_state_size = 2
if optimizer == "psgd":
print("Need to download the psgd optimizer (https://github.com/lixilinx/psgd_torch or from other places)")
import preconditioned_stochastic_gradient_descent as psgd
train_loader = torch.utils.data.DataLoader(
datasets.MNIST(
"../data",
train=True,
download=True,
transform=transforms.Compose([transforms.ToTensor()]),
),
batch_size=60,
shuffle=True,
)
test_loader = torch.utils.data.DataLoader(
datasets.MNIST(
"../data", train=False, transform=transforms.Compose([transforms.ToTensor()])
),
batch_size=60,
shuffle=False,
)
class SSMNet(torch.nn.Module):
def __init__(self):
super(SSMNet, self).__init__()
self.ssm1 = SSM(1, increase_state_size * 16, 16, resample_down=4)
self.ssm2 = SSM(16, increase_state_size * 128, 128, resample_down=4)
self.linear = torch.nn.Linear(128, 10)
def forward(self, u):
x, _ = self.ssm1(u)
x = x * torch.rsqrt(1 + x*x)
x, _ = self.ssm2(x)
x = x[:, -1]
x = x * torch.rsqrt(1 + x*x)
x = self.linear(x)
return x
device = torch.device("cuda:0")
ssmnet = SSMNet().to(device)
lr0 = 1e-3
if optimizer == "psgd":
opt = psgd.Kron(ssmnet.parameters(), lr_params=lr0, lr_preconditioner=0.1,
momentum=0.9, preconditioner_type="whitening", grad_clip_max_norm=100)
else:
opt = torch.optim.Adam(ssmnet.parameters(), lr=lr0)
num_epochs = 20
TrainLosses, TestErrs = [], []
for epoch in range(num_epochs):
for batch, (data, target) in enumerate(train_loader):
data, target = data.to(device), target.to(device)
def closure():
y = ssmnet(torch.reshape(data, [-1, 28*28, 1]))
y = torch.nn.functional.log_softmax(y, dim=-1)
xentropy = torch.nn.functional.nll_loss(y, target)
return xentropy
if optimizer == "psgd":
loss = opt.step(closure)
else: # Adam
opt.zero_grad()
loss = closure()
loss.backward()
opt.step()
TrainLosses.append(loss.item())
if (batch+1) % 100 == 0:
print(f"Epoch: {epoch + 1}; train loss: {np.mean(TrainLosses[-1000:])}")
# test loss
with torch.no_grad():
num_errs = 0
for data, target in test_loader:
data, target = data.to(device), target.to(device)
y = ssmnet(torch.reshape(data, [-1, 28*28, 1]))
_, pred = torch.max(y, dim=1)
num_errs += torch.sum(pred != target)
test_err_rate = num_errs.item() / len(test_loader.dataset)
TestErrs.append(test_err_rate)
print(f"Epoch: {epoch + 1}; test classification error rate: {TestErrs[-1]}")
# linear lr schedule
if optimizer == "psgd":
opt.preconditioner_update_probability = 0.1
opt.lr_params -= lr0 / num_epochs
else:
opt.param_groups[0]["lr"] -= lr0 / num_epochs
plt.plot(TrainLosses)
plt.show()
plt.plot(TestErrs)
plt.show()