diff --git a/src/rbm/train.py b/src/rbm/train.py index 8cda7a4..1a785a3 100644 --- a/src/rbm/train.py +++ b/src/rbm/train.py @@ -63,6 +63,7 @@ def cd_jens(entity: Entity, v_states: Mat): # %%%%%%%%% END OF NEGATIVE PHASE %%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%% def cd_binary_binary(entity: Entity, data_pos: Mat): params = entity.training_params + # Positive phase h_probs_pos = prob(entity.h_given_v(data_pos)) @@ -92,19 +93,25 @@ def cd_binary_binary(entity: Entity, data_pos: Mat): return dw, dbv, dbh def cd_gaussian_binary(entity: Entity, data_pos: Mat): + params = entity.training_params + # Positive phase - h_probs_pos = entity.h_given_v(data_pos) + if params.do_batch_sample: + h_probs_pos = entity.h_given_v(sample_gaussian(data_pos)) + else: + h_probs_pos = entity.h_given_v(data_pos) # Update weights (positive phase) dw = np.dot(np.transpose(data_pos), h_probs_pos) dbh = np.sum(h_probs_pos, 0) dbv = np.sum(data_pos, 0) - # Sample hidden states - h_states_pos = sample(h_probs_pos) + if not params.do_rao_blackwell: + # Sample hidden states + h_probs_pos = sample(h_probs_pos) # Negative phase - data_neg = entity.v_given_h(h_states_pos) + data_neg = entity.v_given_h(h_probs_pos) h_probs_neg = entity.h_given_v(data_neg) # Update weights (negative phase) @@ -115,19 +122,24 @@ def cd_gaussian_binary(entity: Entity, data_pos: Mat): return dw, dbv, dbh def cd_gaussian_gaussian(entity: Entity, data_pos: Mat): + params = entity.training_params + # Positive phase - h_probs_pos = entity.h_given_v(data_pos) + if params.do_batch_sample: + h_probs_pos = entity.h_given_v(sample_gaussian(data_pos)) + else: + h_probs_pos = entity.h_given_v(data_pos) # Sample hidden states - h_states_pos = h_probs_pos + sample_gaussian(h_probs_pos) + h_probs_pos += sample_gaussian(h_probs_pos) # Update weights (positive phase) - dw = np.dot(np.transpose(data_pos), h_states_pos) - dbh = np.sum(h_states_pos, 0) + dw = np.dot(np.transpose(data_pos), h_probs_pos) + dbh = np.sum(h_probs_pos, 0) dbv = np.sum(data_pos, 0) # Negative phase - data_neg = entity.v_given_h(h_states_pos) + data_neg = entity.v_given_h(h_probs_pos) h_probs_neg = entity.h_given_v(data_neg) # Update weights (negative phase)