From b4cd05d803f117885016b199a52c4767f6dee013 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Wed, 17 Dec 2025 19:29:46 +0100 Subject: [PATCH] - fixed max index - call status once inside epoch loop --- src/rbm/layer.py | 2 -- src/rbm/rbm.py | 12 ++++++------ 2 files changed, 6 insertions(+), 8 deletions(-) diff --git a/src/rbm/layer.py b/src/rbm/layer.py index 65223f6..55a19bd 100644 --- a/src/rbm/layer.py +++ b/src/rbm/layer.py @@ -76,8 +76,6 @@ class Layer: training_seen += 1 - status.on_change({"progress": {"value": round(training_seen*d_progress), "unit": "%"}, "err_rms": {"value": err_rms, "unit": ""}}) - def v_to_ph(self, v: np.ndarray) -> np.ndarray: state = self.state.v_to_h(v) if self.params.do_gaussian_visible: diff --git a/src/rbm/rbm.py b/src/rbm/rbm.py index 3f09235..78299f5 100644 --- a/src/rbm/rbm.py +++ b/src/rbm/rbm.py @@ -25,7 +25,7 @@ class MyStatus(Status): shape = self.stack.from_index(0).shape[0:2] + (1,) # User input - max_index = self.batch.shape[0] + max_index = self.batch.shape[0] - 1 key = cv.waitKeyEx(1) if key == ord('q'): do_continue = False @@ -74,20 +74,20 @@ def main(prj_name: str = "test"): stack.state_load() # Load train data - training_data_path = os.path.join(prj_root, f"{prj_name}.training.dat") - batch = read_armadillo(training_data_path) + training_data = read_armadillo(os.path.join(prj_root, f"{prj_name}.training.dat")) + test_data = read_armadillo(os.path.join(prj_root, f"{prj_name}.test.dat")) # Prepare status listener - my_status = MyStatus(stack, batch) + my_status = MyStatus(stack, test_data) # Train - stack.train(batch, status=my_status) + stack.train(training_data, status=my_status) # Save state stack.state_save() if __name__ == "__main__": - main("norb_small_16h") + main("norb_small_16h_v2") cv.destroyAllWindows() print("Test: [passed]")