373 KiB
373 KiB
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 model.model import Model
from rbm.entity import Entity, EntityParams, TrainingParams
from rbm.matrix import Mat, np, rms_error_accu, sample_gaussian
from compat.torch import Optimizer
from image.sub_image import SubImageExtract, SubImageCombine, normalize
import math
import randomIn [2]:
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
]) In [3]:
train_data = torchvision.datasets.CIFAR10(root='./data', train=True, transform=transform, download=True)
test_data = torchvision.datasets.CIFAR10(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(train_data)Dataset CIFAR10
Number of datapoints: 50000
Root location: ./data
Split: Train
StandardTransform
Transform: Compose(
ToTensor()
Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))
)
In [5]:
N = 500
n_side = 10
PATCH = 8
STRIDE = 1
N_VIS = 3 * PATCH * PATCH
image, label = train_data[0]
size = len(torch.flatten(image))
train_images = np.zeros(shape=(N, size))
train_images_tensor = torch.Tensor(N, size)
for i in range(N):
image, label = train_data[i]
flat_array = torch.flatten(image)
train_images[i, :] = np.asarray(flat_array.numpy())
train_images_tensor[i, :] = flat_array
print(image.shape)torch.Size([3, 32, 32])
In [6]:
image_indices = np.array(random.sample(range(min(N, n_side*n_side)), n_side*n_side))
print(f"train_images.shape:{train_images.shape}")
mean_training_batch = np.reshape(np.repeat(np.mean(train_images, axis=1), 3072, axis=0), train_images.shape)
var_training_batch = np.reshape(np.repeat(np.std(train_images, axis=1), 3072, axis=0), train_images.shape)
train_images_norm = (train_images - mean_training_batch) / var_training_batch
sub_generator = SubImageExtract(PATCH, PATCH, STRIDE, STRIDE)
train_subimages = sub_generator(np.reshape(train_images, (N, 3, 32, 32)), 3)
print(train_subimages.shape)
train_subimages = np.reshape(train_subimages, (train_subimages.shape[0], N_VIS))
print(train_subimages.shape)train_images.shape:(500, 3072) (312500, 3, 8, 8) (312500, 192)
In [7]:
class TestModel(Model):
def __init__(self, name: str, work_dir: str = '.'):
super().__init__(name, work_dir)
self.unit1 = Entity((N_VIS, 128), EntityParams(do_gaussian_visible=True, do_gaussian_hidden=False), TrainingParams(learning_rate=0.005, momentum=0.9, num_epochs=1000, mini_batch_size=1000, weight_decay=0.0, l1_lambda=0.02))
def forward(self, x: Mat):
x = self.unit1.forward(x)
return x
def backward(self, x: Mat):
x = self.unit1.reconstruct(x)
return xIn [8]:
work_dir = "results"
prj_name = "cifar_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(normalize(train_subimages))
model.train(train_subimages)
# save state
model.save()------------------------------------------- Entity-192x128: progress : 0% Entity-192x128: err_rms : 0.02982608561419449 Entity-192x128: l2_norm : 0.5154587953841656 ------------------------------------------- Entity-192x128: progress : 10% Entity-192x128: err_rms : 0.028228882569056157 Entity-192x128: l2_norm : 0.5495957617469555 ------------------------------------------- Entity-192x128: progress : 20% Entity-192x128: err_rms : 0.027550804062706425 Entity-192x128: l2_norm : 0.5878319340382664 ------------------------------------------- Entity-192x128: progress : 30% Entity-192x128: err_rms : 0.02740561917531193 Entity-192x128: l2_norm : 0.5976701067086865 ------------------------------------------- Entity-192x128: progress : 40% Entity-192x128: err_rms : 0.027248082195068083 Entity-192x128: l2_norm : 0.5991757429205762 ------------------------------------------- Entity-192x128: progress : 50% Entity-192x128: err_rms : 0.026824203916998864 Entity-192x128: l2_norm : 0.6027710231695954 ------------------------------------------- Entity-192x128: progress : 60% Entity-192x128: err_rms : 0.02655947728844083 Entity-192x128: l2_norm : 0.6041359064549188 ------------------------------------------- Entity-192x128: progress : 70% Entity-192x128: err_rms : 0.026933387882988033 Entity-192x128: l2_norm : 0.603607961995845 ------------------------------------------- Entity-192x128: progress : 80% Entity-192x128: err_rms : 0.02728050450716867 Entity-192x128: l2_norm : 0.6036734185897112 ------------------------------------------- Entity-192x128: progress : 90% Entity-192x128: err_rms : 0.027154601608006385 Entity-192x128: l2_norm : 0.6046113788661706 ------------------------------------------- Entity-192x128: progress : 100% Entity-192x128: err_rms : 0.027059586883727227 Entity-192x128: l2_norm : 0.6035971257113926 ------------------------------------------- Entity-192x128: progress : 100% Entity-192x128: err_rms_total : 0.02683939309309671 Entity-192x128: l2_norm : 0.604365570485915
In [9]:
combiner = SubImageCombine(PATCH, PATCH, STRIDE, STRIDE, 32, 32, n_ch=3)
n_images = n_side
nx_steps = (32 - PATCH) // STRIDE + 1
ny_steps = (32 - PATCH) // STRIDE + 1
fig, axes = plt.subplots(n_images, 2, figsize=(6, n_images * 3))
for img_idx in range(n_images):
orig = np.reshape(train_images[img_idx], (3, 32, 32))
orig_display = np.asnumpy((orig - np.min(orig)) / (np.max(orig) - np.min(orig)))
start = img_idx * nx_steps * ny_steps
patches = train_subimages[start : start + nx_steps * ny_steps]
recon_patches = np.array([model.backward(model.forward(p)) for p in patches])
recon_full = combiner(recon_patches)
axes[img_idx, 0].imshow(orig_display.transpose(1, 2, 0))
axes[img_idx, 0].axis('off')
if img_idx == 0:
axes[img_idx, 0].set_title('Original')
axes[img_idx, 1].imshow(recon_full)
axes[img_idx, 1].axis('off')
if img_idx == 0:
axes[img_idx, 1].set_title('Reconstructed')
plt.tight_layout()
plt.show()In [10]:
# Plot weights
weights = model.unit1.state.w_hv
n, n_hid = weights.shape
weight_indices = np.array(random.sample(range(n_hid), n_side*n_side))
w = np.reshape(weights, (3, PATCH, PATCH, n_hid))
print(w.shape)
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[:,:,:, weight_indices[index]])
inp -= np.min(inp)
inp = inp / np.max(inp)
img = inp.transpose((1,2,0))
axes[x,y].imshow(np.asnumpy(img))
axes[x,y].axis('off')
index += 1
plt.show()(3, 8, 8, 128)
In [11]:
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_subimages[image_indices[index]]
img = (1+np.reshape(inp, (3, PATCH, PATCH)))/2
img = img.transpose((1,2,0))
axes[x,y].imshow(np.asnumpy(img))
axes[x,y].axis('off')
index += 1
plt.show()In [11]:
In [11]:
In [12]:
# Plot reconstructions
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_subimages[image_indices[index]]
recon = model.backward(model.forward(inp))
recon -= np.min(recon)
recon = recon / np.max(recon)
img = np.reshape(recon, (3, PATCH, PATCH))
img = img.transpose((1,2,0))
axes[x,y].imshow(np.asnumpy(img))
axes[x,y].axis('off')
index += 1
plt.show()In [12]: