diff --git a/src/tests/test_learn_encoded_labels.py b/src/tests/test_learn_encoded_labels.py index d259909..aa57bda 100644 --- a/src/tests/test_learn_encoded_labels.py +++ b/src/tests/test_learn_encoded_labels.py @@ -59,12 +59,22 @@ if __name__ == "__main__": model.save() # Plot - fig, axes = plt.subplots(1, len(encoded), figsize=(12, 3)) + cmap = 'Grays' + fig, axes = plt.subplots(3, len(encoded), figsize=(12, 3)) for index, inp in enumerate(encoded): - img = np.reshape(inp, (label_w, label_h)) - axes[index].imshow(img) - axes[index].axis('off') - axes[index].set_title(f'{test_labels[index]}') + hidden = model.forward(inp) + recon = model.backward(hidden) + img_hidden = np.reshape(hidden, (4, 4)) + img_recon = np.reshape(recon, (label_w, label_h)) + axes[0][index].imshow(img_hidden, cmap=cmap) + axes[0][index].axis('off') + axes[0][index].set_title(f'{test_labels[index]}') + axes[1][index].imshow(img_recon, cmap=cmap) + axes[1][index].axis('off') + axes[1][index].set_title(f'{test_labels[index]}') + axes[2][index].imshow(np.reshape(inp, (label_w, label_h)), cmap=cmap) + axes[2][index].axis('off') + axes[2][index].set_title(f'{test_labels[index]}') plt.show() diff --git a/src/tests/test_norbs.py b/src/tests/test_norbs.py index 7e06aa9..410c6b9 100644 --- a/src/tests/test_norbs.py +++ b/src/tests/test_norbs.py @@ -8,7 +8,7 @@ from rbm.matrix import Mat, np, read_armadillo class TestModel(Model): def __init__(self, name: str, work_dir: str = '.'): super().__init__(name, work_dir) - self.unit1 = Entity((96*96, 16), EntityParams(do_gaussian_visible=True, do_gaussian_hidden=False), TrainingParams(learning_rate=0.000001, momentum=0.9, do_rao_blackwell=True, num_epochs=1000)) + self.unit1 = Entity((96*96, 16), EntityParams(do_gaussian_visible=True, do_gaussian_hidden=False), TrainingParams(learning_rate=0.000001, momentum=0.9, num_epochs=1000)) def forward(self, x: Mat): x = self.unit1.forward(x) @@ -34,9 +34,10 @@ if __name__ == "__main__": # Load train data train_batch = read_armadillo(os.path.join(prj_root, f"{prj_name}.training.dat")) + test_batch = train_batch # Load test data - test_batch = read_armadillo(os.path.join(prj_root, f"{prj_name}.test.dat")) + #test_batch = read_armadillo(os.path.join(prj_root, f"{prj_name}.test.dat")) # Train model.train(train_batch)