[bugfix] - remove dead cd_jens; fix cd_gaussian_binary dbv; fix cd_binary_gaussian Gibbs sampling
- Remove cd_jens: dead code never called, contained multiple bugs. - cd_gaussian_binary: dbv formula simplified — b_v terms cancelled between positive and negative phases, making the subtraction redundant. - cd_binary_gaussian: Gibbs loop (CD-k > 1) now properly samples both visible (sample(prob(...))) and hidden (+ sample_gaussian) states when do_rao_blackwell=False. Previously used means throughout, giving biased gradient estimates for k > 1. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
+7
-44
@@ -2,46 +2,6 @@ from .matrix import sample, sample_gaussian, prob, rms_error_accu, Mat, np
|
|||||||
from .entity import Entity
|
from .entity import Entity
|
||||||
from .status import Status
|
from .status import Status
|
||||||
|
|
||||||
def cd_jens(entity: Entity, v_states: Mat):
|
|
||||||
params = entity.training_params
|
|
||||||
v_probs = prob(v_states)
|
|
||||||
h_states = entity.forward(v_states)
|
|
||||||
h_probs = h_states
|
|
||||||
|
|
||||||
if entity.params.do_gaussian_hidden:
|
|
||||||
h_states += sample_gaussian(h_states)
|
|
||||||
else:
|
|
||||||
h_probs = entity.forward(v_states)
|
|
||||||
if params.do_rao_blackwell:
|
|
||||||
h_states = h_probs
|
|
||||||
else:
|
|
||||||
h_states = sample(h_probs)
|
|
||||||
|
|
||||||
# Update weights (positive phase)
|
|
||||||
dw = np.dot(np.transpose(v_states), h_states)
|
|
||||||
dbv = np.sum(v_states, 0)
|
|
||||||
dbh = np.sum(h_states, 0)
|
|
||||||
|
|
||||||
# Gibbs sampling with training params
|
|
||||||
for i in range(params.num_gibbs_samples):
|
|
||||||
if params.do_gibbs_sample_hidden:
|
|
||||||
v_probs = entity.reconstruct(sample(h_probs))
|
|
||||||
else:
|
|
||||||
v_probs = entity.reconstruct(h_probs)
|
|
||||||
|
|
||||||
# Create hidden representation given v
|
|
||||||
if params.do_gibbs_sample_visible:
|
|
||||||
h_probs = entity.forward(sample(v_probs))
|
|
||||||
else:
|
|
||||||
h_probs = entity.forward(v_probs)
|
|
||||||
|
|
||||||
# Update weights (negative phase)
|
|
||||||
dw -= np.dot(np.transpose(v_probs), h_probs)
|
|
||||||
dbv -= np.sum(v_probs, 0)
|
|
||||||
dbh -= np.sum(h_probs, 0)
|
|
||||||
|
|
||||||
return dw, dbv, dbh
|
|
||||||
|
|
||||||
# %%%%%%%%% START POSITIVE PHASE %%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%
|
# %%%%%%%%% START POSITIVE PHASE %%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%
|
||||||
# data = batchdata(:,:,batch);
|
# data = batchdata(:,:,batch);
|
||||||
# poshidprobs = 1./(1 + exp(-data*vishid - repmat(hidbiases,numcases,1)));
|
# poshidprobs = 1./(1 + exp(-data*vishid - repmat(hidbiases,numcases,1)));
|
||||||
@@ -99,7 +59,7 @@ def cd_gaussian_binary(entity: Entity, data_pos: Mat):
|
|||||||
# Update weights (positive phase)
|
# Update weights (positive phase)
|
||||||
dw = np.dot(np.transpose(data_pos), h_probs_pos)
|
dw = np.dot(np.transpose(data_pos), h_probs_pos)
|
||||||
dbh = np.sum(h_probs_pos, 0)
|
dbh = np.sum(h_probs_pos, 0)
|
||||||
dbv = np.sum(data_pos-entity.state.b_v, 0)
|
dbv = np.sum(data_pos, 0)
|
||||||
|
|
||||||
# Negative phase — sample binary hidden states to get a sharper fantasy particle
|
# Negative phase — sample binary hidden states to get a sharper fantasy particle
|
||||||
data_neg = entity.v_given_h(sample(h_probs_pos))
|
data_neg = entity.v_given_h(sample(h_probs_pos))
|
||||||
@@ -108,7 +68,7 @@ def cd_gaussian_binary(entity: Entity, data_pos: Mat):
|
|||||||
# Update weights (negative phase)
|
# Update weights (negative phase)
|
||||||
dw -= np.dot(np.transpose(data_neg), h_probs_neg)
|
dw -= np.dot(np.transpose(data_neg), h_probs_neg)
|
||||||
dbh -= np.sum(h_probs_neg, 0)
|
dbh -= np.sum(h_probs_neg, 0)
|
||||||
dbv -= np.sum(data_neg-entity.state.b_v, 0)
|
dbv -= np.sum(data_neg, 0)
|
||||||
|
|
||||||
return dw, dbv, dbh
|
return dw, dbv, dbh
|
||||||
|
|
||||||
@@ -169,10 +129,13 @@ def cd_binary_gaussian(entity: Entity, data_pos: Mat):
|
|||||||
data_neg = prob(entity.v_given_h(h_probs_pos))
|
data_neg = prob(entity.v_given_h(h_probs_pos))
|
||||||
h_probs_neg = entity.h_given_v(data_neg)
|
h_probs_neg = entity.h_given_v(data_neg)
|
||||||
|
|
||||||
# Gibbs sampling
|
# Gibbs sampling — sample both visible and hidden to continue the Markov chain
|
||||||
for _ in range(params.num_gibbs_samples-1):
|
for _ in range(params.num_gibbs_samples-1):
|
||||||
data_neg = prob(entity.v_given_h(h_probs_neg))
|
v_probs = prob(entity.v_given_h(h_probs_neg))
|
||||||
|
data_neg = v_probs if params.do_rao_blackwell else sample(v_probs)
|
||||||
h_probs_neg = entity.h_given_v(data_neg)
|
h_probs_neg = entity.h_given_v(data_neg)
|
||||||
|
if not params.do_rao_blackwell:
|
||||||
|
h_probs_neg = h_probs_neg + sample_gaussian(h_probs_neg)
|
||||||
|
|
||||||
# Update weights (negative phase)
|
# Update weights (negative phase)
|
||||||
dw -= np.dot(np.transpose(data_neg), h_probs_neg)
|
dw -= np.dot(np.transpose(data_neg), h_probs_neg)
|
||||||
|
|||||||
Reference in New Issue
Block a user