From 48f3642e44635577edf976c3953c80b998839f0f Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Fri, 2 Jan 2026 18:36:17 +0100 Subject: [PATCH] fixed TrainingParams --- src/tests/test_norbs.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/src/tests/test_norbs.py b/src/tests/test_norbs.py index 855f647..7e06aa9 100644 --- a/src/tests/test_norbs.py +++ b/src/tests/test_norbs.py @@ -2,14 +2,13 @@ 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.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)) def forward(self, x: Mat): x = self.unit1.forward(x) @@ -40,7 +39,7 @@ if __name__ == "__main__": test_batch = read_armadillo(os.path.join(prj_root, f"{prj_name}.test.dat")) # Train - model.train(train_batch, TrainingParams(learning_rate=0.000001, momentum=0.9, do_rao_blackwell=False, num_epochs=1000, num_gibbs_samples=3)) + model.train(train_batch) # save state model.save()