refactored

This commit is contained in:
2025-12-21 19:34:23 +01:00
parent ff05c729e8
commit 362970574f
+7 -5
View File
@@ -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))