From 362970574fe25b56c04536a06b25cb8937fb2d05 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Sun, 21 Dec 2025 19:34:23 +0100 Subject: [PATCH] refactored --- src/tests/test_norbs.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/src/tests/test_norbs.py b/src/tests/test_norbs.py index 93ea616..1edccb0 100644 --- a/src/tests/test_norbs.py +++ b/src/tests/test_norbs.py @@ -37,17 +37,19 @@ if __name__ == "__main__": model.load() # Load train data - batch = read_armadillo(os.path.join(prj_root, f"{prj_name}.training.dat")) + train_batch = read_armadillo(os.path.join(prj_root, f"{prj_name}.training.dat")) + + # Load test data + test_batch = read_armadillo(os.path.join(prj_root, f"{prj_name}.test.dat")) # Train - model.train(batch, TrainingParams(learning_rate=0.00001, momentum=0.9, do_rao_blackwell=True, num_epochs=100, num_gibbs_samples=3)) + model.train(train_batch, TrainingParams(learning_rate=0.00001, momentum=0.9, do_rao_blackwell=True, num_epochs=100, num_gibbs_samples=3)) # save state model.save() - num_patterns = len(batch) - fig, axes = plt.subplots(1, num_patterns, figsize=(12, 3)) - for index, inp in enumerate(batch): + fig, axes = plt.subplots(1, len(test_batch), figsize=(12, 3)) + for index, inp in enumerate(test_batch): out_normalized = model.backward(model.forward(inp)) img = 2*(out_normalized + 0.5) img = np.reshape(img, (96, 96))