From 6546c376571e08e31b592464e860e92c8706f8a2 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Sat, 3 Jan 2026 19:30:47 +0100 Subject: [PATCH] train: fixed gaussian sampling --- src/rbm/train.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/src/rbm/train.py b/src/rbm/train.py index 1a785a3..773f38a 100644 --- a/src/rbm/train.py +++ b/src/rbm/train.py @@ -97,7 +97,7 @@ def cd_gaussian_binary(entity: Entity, data_pos: Mat): # Positive phase if params.do_batch_sample: - h_probs_pos = entity.h_given_v(sample_gaussian(data_pos)) + h_probs_pos = entity.h_given_v(data_pos + sample_gaussian(data_pos)) else: h_probs_pos = entity.h_given_v(data_pos) @@ -126,12 +126,13 @@ def cd_gaussian_gaussian(entity: Entity, data_pos: Mat): # Positive phase if params.do_batch_sample: - h_probs_pos = entity.h_given_v(sample_gaussian(data_pos)) + h_probs_pos = entity.h_given_v(data_pos + sample_gaussian(data_pos)) else: h_probs_pos = entity.h_given_v(data_pos) - # Sample hidden states - h_probs_pos += sample_gaussian(h_probs_pos) + if not params.do_rao_blackwell: + # Sample hidden states + h_probs_pos += sample_gaussian(h_probs_pos) # Update weights (positive phase) dw = np.dot(np.transpose(data_pos), h_probs_pos)