From 73d4486fb76512596ea2597694953afe671d5183 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Sun, 4 Jan 2026 16:06:41 +0100 Subject: [PATCH] test_norbs: added mode hidden gaussian - hidden binary --- src/tests/test_norbs.py | 18 ++++++++++++------ 1 file changed, 12 insertions(+), 6 deletions(-) diff --git a/src/tests/test_norbs.py b/src/tests/test_norbs.py index 410c6b9..9dbf6ef 100644 --- a/src/tests/test_norbs.py +++ b/src/tests/test_norbs.py @@ -6,9 +6,16 @@ from rbm.entity import Entity, EntityParams, TrainingParams from rbm.matrix import Mat, np, read_armadillo class TestModel(Model): - def __init__(self, name: str, work_dir: str = '.'): + def __init__(self, name: str, work_dir: str = '.', do_gaussian_hidden=False): 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, num_epochs=1000)) + + if do_gaussian_hidden: + # Hidden gaussian + self.unit1 = Entity((96*96, 16), EntityParams(do_gaussian_visible=True, do_gaussian_hidden=True), TrainingParams(learning_rate=0.00001, momentum=0.9, num_epochs=1000, do_rao_blackwell=True)) + else: + # Hidden binary + self.unit1 = Entity((96 * 96, 16), EntityParams(do_gaussian_visible=True, do_gaussian_hidden=False), + TrainingParams(learning_rate=0.000002, momentum=0.9, num_epochs=1000)) def forward(self, x: Mat): x = self.unit1.forward(x) @@ -24,20 +31,19 @@ if __name__ == "__main__": prj_root = "/home/jens/work/repos/Rbm" # Create model - model = TestModel("norb_small_16h_v2", "results") + model = TestModel(prj_name, "results") # Init state model.init(0.01) # load state - model.load() +# model.load() # 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)