From 8b973505643513942d837143048682e5462495e9 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Thu, 1 Jan 2026 20:02:22 +0100 Subject: [PATCH] cd_binary_binary: added gibbs sampling --- src/rbm/entity.py | 1 - src/rbm/train.py | 6 ++++-- src/tests/test_learn_encoded_labels.py | 6 +++--- 3 files changed, 7 insertions(+), 6 deletions(-) diff --git a/src/rbm/entity.py b/src/rbm/entity.py index 8031121..8196763 100644 --- a/src/rbm/entity.py +++ b/src/rbm/entity.py @@ -102,7 +102,6 @@ class Entity: state = self.state.v_to_h(v) return state - def v_given_h(self, h: Mat) -> Mat: state = self.state.h_to_v(h) return state diff --git a/src/rbm/train.py b/src/rbm/train.py index 5c9f048..3d88523 100644 --- a/src/rbm/train.py +++ b/src/rbm/train.py @@ -100,8 +100,10 @@ def cd_binary_binary(entity: Entity, data_pos: Mat, params: TrainingParams): data_neg = prob(entity.v_given_h(h_probs_pos)) h_probs_neg = prob(entity.h_given_v(data_neg)) - # Gibbs sampling with training params - # ToDo + # Gibbs sampling + for _ in range(params.num_gibbs_samples-1): + data_neg = prob(entity.v_given_h(h_probs_neg)) + h_probs_neg = prob(entity.h_given_v(data_neg)) # Update weights (negative phase) dw -= np.dot(np.transpose(data_neg), h_probs_neg) diff --git a/src/tests/test_learn_encoded_labels.py b/src/tests/test_learn_encoded_labels.py index 8bd08e6..9db049e 100644 --- a/src/tests/test_learn_encoded_labels.py +++ b/src/tests/test_learn_encoded_labels.py @@ -10,7 +10,7 @@ from rbm.train import TrainingParams class LabelLearner(Model): def __init__(self, name: str, work_dir: str = '.'): super().__init__(name, work_dir) - self.unit1 = Entity((16, 16), EntityParams()) + self.unit1 = Entity((16, 8), EntityParams()) def forward(self, x: Mat): x = self.unit1.forward(x) @@ -45,7 +45,7 @@ if __name__ == "__main__": model.load() # Train - model.train(encoded, TrainingParams(learning_rate=0.01, momentum=0.9, do_rao_blackwell=True, num_epochs=10000)) + model.train(encoded, TrainingParams(learning_rate=0.01, momentum=0.9, do_rao_blackwell=True, num_epochs=1000)) # save state model.save() @@ -57,7 +57,7 @@ if __name__ == "__main__": axes[index].imshow(img) axes[index].axis('off') axes[index].set_title(f'{test_labels[index]}') - + print(inp) plt.show() print("Test: [passed]")