2.6 MiB
2.6 MiB
In [1]:
from PIL import Image
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
import torchvision
import torchvision.transforms as transforms
import matplotlib.pyplot as plt
from rbm.model import Model
from rbm.entity import Entity, EntityParams, TrainingParams
from rbm.matrix import Mat, np, rms_error_accu
from rbm.torch import Optimizer
import mathIn [2]:
transform = transforms.Compose([
transforms.ToTensor()
]) In [3]:
train_data = torchvision.datasets.MNIST(root='./data', train=True, transform=transform, download=True)
test_data = torchvision.datasets.MNIST(root='./data', train=False, transform=transform, download=True)
train_loader = torch.utils.data.DataLoader(train_data, batch_size=32, shuffle=True, num_workers=2)
test_loader = torch.utils.data.DataLoader(test_data, batch_size=32, shuffle=True, num_workers=2)In [4]:
print(len(train_data))60000
In [5]:
print(list(train_data[1][0].size()[1:3]))[28, 28]
In [6]:
N = 400
size = 784
train_images = np.zeros(shape=(N,size))
for i in range(N):
image = train_data[i][0]
flat_array = torch.flatten(image)
train_images[i, :] = np.asarray(flat_array.numpy())
In [7]:
print(len(train_images))400
In [8]:
n_side = int(math.sqrt(len(train_images)))
fig, axes = plt.subplots(n_side, n_side, figsize=(30,30))
index = 0
for x in range(n_side):
for y in range(n_side):
inp = train_images[index]
img = np.reshape(inp, (28, 28))
axes[x,y].imshow(np.asnumpy(img))
axes[x,y].axis('off')
index += 1
plt.show()In [9]:
class TestModel(Model):
def __init__(self, name: str, work_dir: str = '.'):
super().__init__(name, work_dir)
self.unit1 = Entity((28*28, 512), EntityParams(do_gaussian_visible=False, do_gaussian_hidden=False), TrainingParams(learning_rate=0.01, momentum=0.9, num_epochs=1000, mini_batch_size=100, weight_decay=0.0, l2_lambda=0.8), True)
self.unit2 = Entity((512, 64), EntityParams(do_gaussian_visible=False, do_gaussian_hidden=False), TrainingParams(learning_rate=0.04, momentum=0.9, num_epochs=1000, mini_batch_size=100, weight_decay=0.0, do_rao_blackwell=True), False)
def forward(self, x: Mat):
x = self.unit1.forward(x)
# x = self.unit2.forward(x)
return x
def backward(self, x: Mat):
x = self.unit1.reconstruct(x)
# x = self.unit1.reconstruct(x)
return x
In [ ]:
work_dir = "results"
prj_name = "mnist_test"
prj_root = "/home/jens/work/repos/Rbm"
# Create model
model = TestModel(prj_name, "results")
# Init state
model.init(0.01)
# load state
model.load()
# Train
model.train(train_images)
# save state
model.save()results/mnist_test-0-state.npz loaded successfully! results/mnist_test-1-state.npz loaded successfully! ------------------------------------------- Entity-784x512: progress : 0% Entity-784x512: err_rms : 0.009141128622640915 Entity-784x512: l2_norm : 2.9563627676909805 ------------------------------------------- Entity-784x512: progress : 10% Entity-784x512: err_rms : 0.009093806900410771 Entity-784x512: l2_norm : 2.9629020988619446
In [25]:
# Plot reconstructions
n_side = int(math.sqrt(len(train_images)))
fig, axes = plt.subplots(n_side, n_side, figsize=(30,30))
index = 0
for x in range(n_side):
for y in range(n_side):
inp = train_images[index]
recon = model.backward(model.forward(inp))
recon -= np.min(recon)
recon = recon / np.max(recon)
img = np.reshape(recon, (28, 28))
axes[x,y].imshow(np.asnumpy(img))
axes[x,y].axis('off')
index += 1
plt.show()In [26]:
# Plot weights
weights = model.unit1.state.w_hv
n, n_hid = weights.shape
w = np.reshape(weights, (28, 28, n_hid))
print(w.shape)
n_side = int(math.sqrt(n_hid))
fig, axes = plt.subplots(n_side, n_side, figsize=(30,30))
index = 0
for x in range(n_side):
for y in range(n_side):
inp = np.asnumpy(w[:,:, index])
axes[x,y].imshow(inp)
axes[x,y].axis('off')
index += 1
plt.show()(28, 28, 512)
In [23]:
# torch styleIn [14]:
model.load()
optimizer = Optimizer(model.unit1)
report_interval = 1000
next_ep = report_interval
optimizer.zero_grad()
for epoch in range(10000):
running_loss = 0.0
optimizer.step(train_images)
if epoch >= next_ep:
print(f'Training epoch {epoch}...')
next_ep += report_interval
# Update final status
err_rms = rms_error_accu(train_images - model.unit1.reconstruct(model.unit1.forward(train_images)))
print(err_rms)
model.save()
results/mnist_test-0-state.npz loaded successfully! results/mnist_test-1-state.npz loaded successfully!
[31m---------------------------------------------------------------------------[39m [31mTypeError[39m Traceback (most recent call last) [36mCell[39m[36m [39m[32mIn[14][39m[32m, line 9[39m [32m 6[39m [38;5;28;01mfor[39;00m epoch [38;5;129;01min[39;00m [38;5;28mrange[39m([32m10000[39m): [32m 7[39m running_loss = [32m0.0[39m [32m----> [39m[32m9[39m [43moptimizer[49m[43m.[49m[43mstep[49m[43m([49m[43mtrain_images[49m[43m)[49m [32m 11[39m [38;5;28;01mif[39;00m epoch >= next_ep: [32m 12[39m [38;5;28mprint[39m([33mf[39m[33m'[39m[33mTraining epoch [39m[38;5;132;01m{[39;00mepoch[38;5;132;01m}[39;00m[33m...[39m[33m'[39m) [36mFile [39m[32m~/work/pyRBM/src/rbm/torch.py:38[39m, in [36mOptimizer.step[39m[34m(self, data)[39m [32m 35[39m dwhv, dbv, dbh = [38;5;28mself[39m.loss([38;5;28mself[39m.entity, data) [32m 37[39m [38;5;66;03m# Adjust weight and biases[39;00m [32m---> [39m[32m38[39m grad = [38;5;28;43mself[39;49m[43m.[49m[43mentity[49m[43m.[49m[43mgrad_compute[49m[43m([49m[43mdbv[49m[43m,[49m[43m [49m[43mdbh[49m[43m,[49m[43m [49m[43mdwhv[49m[43m,[49m[43m [49m[43mlearning_rate[49m[43m=[49m[43mparams[49m[43m.[49m[43mlearning_rate[49m[43m [49m[43m/[49m[43m [49m[43mdata[49m[43m.[49m[43mshape[49m[43m[[49m[32;43m0[39;49m[43m][49m[43m,[49m [32m 39[39m [43m [49m[43mmomentum[49m[43m=[49m[43mparams[49m[43m.[49m[43mmomentum[49m[43m,[49m[43m [49m[43mweight_decay[49m[43m=[49m[43mparams[49m[43m.[49m[43mweight_decay[49m[43m)[49m [32m 40[39m [38;5;66;03m# Adjust weights[39;00m [32m 41[39m [38;5;28mself[39m.entity.state_adjust(grad) [31mTypeError[39m: Entity.grad_compute() missing 1 required positional argument: 'l2_lambda'
In [ ]:
In [ ]:
In [ ]:
In [ ]:
In [ ]:
In [ ]:
In [ ]:
In [ ]:
In [ ]:
In [ ]:
In [ ]:
In [ ]:
In [ ]:
In [ ]:
In [ ]:
In [ ]:
In [ ]: