From 42520e5761e989e9a43e017b3af6ae66fc75d15f Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Tue, 6 Jan 2026 09:44:45 +0100 Subject: [PATCH] improved tests --- src/tests/test_deep_norbs.py | 16 ++++++---------- src/tests/test_model.py | 18 ++++++++++-------- src/tests/test_norbs.py | 2 +- 3 files changed, 17 insertions(+), 19 deletions(-) diff --git a/src/tests/test_deep_norbs.py b/src/tests/test_deep_norbs.py index 2cd681f..ea37de6 100644 --- a/src/tests/test_deep_norbs.py +++ b/src/tests/test_deep_norbs.py @@ -2,15 +2,14 @@ import os import matplotlib.pyplot as plt from rbm.model import Model -from rbm.entity import Entity, EntityParams +from rbm.entity import Entity, EntityParams, TrainingParams from rbm.matrix import Mat, np, read_armadillo -from rbm.train import TrainingParams 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=True)) - self.unit2 = Entity((16, 16), EntityParams()) + self.unit1 = Entity((96*96, 333), EntityParams(do_gaussian_visible=True, do_gaussian_hidden=False), TrainingParams(learning_rate=0.001, momentum=0.9, num_epochs=1000), enable_training=False) + self.unit2 = Entity((333, 256), EntityParams(), TrainingParams(learning_rate=0.1, momentum=0.9, num_epochs=10000, do_rao_blackwell=True)) def forward(self, x: Mat): x = self.unit1.forward(x) @@ -31,7 +30,7 @@ if __name__ == "__main__": model = TestModel(prj_name, "results") # Init state - model.init(0.01) + model.init(0.1) # load state model.load() @@ -43,10 +42,7 @@ if __name__ == "__main__": test_batch = read_armadillo(os.path.join(prj_root, f"norb_small_16h_v2.test.dat")) # Train - model.train(train_batch,[ - TrainingParams(learning_rate=0.00001, momentum=0.9, do_rao_blackwell=True, num_epochs=1000, num_gibbs_samples=3), - TrainingParams(learning_rate=0.01, momentum=0.9, do_rao_blackwell=True, num_epochs=1000, num_gibbs_samples=1) - ]) + model.train(train_batch) # save state model.save() @@ -56,7 +52,7 @@ if __name__ == "__main__": out_normalized = model.backward(model.forward(inp)) img = 2*(out_normalized + 0.5) img = np.reshape(img, (96, 96)) - axes[index].imshow(img) + axes[index].imshow(np.asnumpy(img)) axes[index].axis('off') plt.show() diff --git a/src/tests/test_model.py b/src/tests/test_model.py index 042e07d..188a0e4 100644 --- a/src/tests/test_model.py +++ b/src/tests/test_model.py @@ -1,22 +1,24 @@ from rbm.model import Model -from rbm.entity import Entity, EntityParams +from rbm.entity import Entity, EntityParams, TrainingParams from rbm.matrix import Mat, np -from rbm.train import TrainingParams class TestModel(Model): def __init__(self, name: str, work_dir: str = '.'): super().__init__(name, work_dir) - self.unit1 = Entity((16, 64), EntityParams()) - self.unit2 = Entity((64, 16), EntityParams()) - self.unit3 = Entity((16, 64), EntityParams()) + self.unit1 = Entity((1024, 333), EntityParams(), TrainingParams(learning_rate=0.1, momentum=0.9, do_rao_blackwell=True, num_epochs=1000)) + self.unit2 = Entity((333, 64), EntityParams(), TrainingParams(learning_rate=0.1, momentum=0.9, do_rao_blackwell=True, num_epochs=1000)) + self.unit3 = Entity((64, 128), EntityParams(), TrainingParams(learning_rate=0.1, momentum=0.9, do_rao_blackwell=True, num_epochs=1000)) + self.unit4 = Entity((128, 128), EntityParams(), TrainingParams(learning_rate=0.1, momentum=0.9, do_rao_blackwell=True, num_epochs=1000)) def forward(self, x: Mat): x = self.unit1.forward(x) x = self.unit2.forward(x) x = self.unit3.forward(x) + x = self.unit4.forward(x) return x def backward(self, x: Mat): + x = self.unit4.reconstruct(x) x = self.unit3.reconstruct(x) x = self.unit2.reconstruct(x) x = self.unit1.reconstruct(x) @@ -27,16 +29,16 @@ if __name__ == "__main__": model = TestModel("TestModel", "results") # Init state - model.init(0.01) + model.init(0.1) # load state model.load() # create batch - batch = (np.random.rand(64, 16) > 0.5).astype(np.float64) + batch = (np.random.rand(64, 1024) > 0.5).astype(np.float64) # Train - model.train(batch, TrainingParams(learning_rate=0.01, momentum=0.9, do_rao_blackwell=True, num_epochs=1000, num_gibbs_samples=3)) + model.train(batch) # save state model.save() diff --git a/src/tests/test_norbs.py b/src/tests/test_norbs.py index a8ad663..e4cec0e 100644 --- a/src/tests/test_norbs.py +++ b/src/tests/test_norbs.py @@ -16,7 +16,7 @@ class TestModel(Model): else: # Hidden binary self.unit1 = Entity((96 * 96, 333), EntityParams(do_gaussian_visible=True, do_gaussian_hidden=False), - TrainingParams(learning_rate=0.0001, momentum=0.9, num_epochs=1000)) + TrainingParams(learning_rate=0.001, momentum=0.9, num_epochs=1000)) def forward(self, x: Mat): x = self.unit1.forward(x)