from .matrix import sample, sample_gaussian, prob, rms_error_accu, Mat, np from .entity import Entity 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 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)) # Sample hidden states if not params.do_rao_blackwell: h_probs_pos = sample(h_probs_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) # Negative phase data_neg = prob(entity.v_given_h(h_probs_pos)) h_probs_neg = prob(entity.h_given_v(data_neg)) # 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) dbh -= np.sum(h_probs_neg, 0) dbv -= np.sum(data_neg, 0) return dw, dbv, dbh def cd_gaussian_binary(entity: Entity, data_pos: Mat): # Positive phase 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) # Negative phase data_neg = entity.v_given_h(h_states_pos) h_probs_neg = entity.h_given_v(data_neg) # Update weights (negative phase) dw -= np.dot(np.transpose(data_neg), h_probs_neg) dbh -= np.sum(h_probs_neg, 0) dbv -= np.sum(data_neg, 0) return dw, dbv, dbh def cd_gaussian_gaussian(entity: Entity, data_pos: Mat): # Positive phase h_probs_pos = entity.h_given_v(data_pos) # Sample hidden states h_states_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) dbv = np.sum(data_pos, 0) # Negative phase data_neg = entity.v_given_h(h_states_pos) h_probs_neg = entity.h_given_v(data_neg) # Update weights (negative phase) dw -= np.dot(np.transpose(data_neg), h_probs_neg) dbh -= np.sum(h_probs_neg, 0) dbv -= np.sum(data_neg, 0) return dw, dbv, dbh def to_mini_batch(batch: Mat, mini_batch_size: int): mini_batches = [] remain = batch.shape[0] index = 0 while remain > 0: batch_size = min(mini_batch_size, remain) mini_batches.append(batch[index:index + batch_size]) remain -= batch_size index += batch_size return mini_batches def train(entity: Entity, batch: Mat, status: Status): params = entity.training_params if params is None: return False mini_batch_size = min(params.mini_batch_size, batch.shape[0]) if params.mini_batch_size > 0 else batch.shape[0] d_progress = 100.0 / (batch.shape[0]*params.num_epochs) progress = 0 keep_running = True status.on_change() cd_func = None if entity.type == Entity.Type.GB_RBM: cd_func = cd_gaussian_binary if entity.type == Entity.Type.BB_RBM: cd_func = cd_binary_binary if entity.type == Entity.Type.GG_RBM: cd_func = cd_gaussian_gaussian for mini_batch in to_mini_batch(batch, mini_batch_size): if not keep_running: break entity.grad_zero() for epochs in range(params.num_epochs): # Contrastive divergence learning: calculate gradients dwhv, dbv, dbh = cd_func(entity, mini_batch) # Adjust weight and biases grad = entity.grad_compute(dbv, dbh, dwhv, learning_rate=params.learning_rate/batch.shape[0], momentum=params.momentum, weight_decay=params.weight_decay) entity.state_adjust(grad) # check if status update is needed if status.want_report(round(progress)): # Calculate error err_rms = rms_error_accu(batch - entity.reconstruct(entity.forward(batch))) if not status.on_change({"progress": {"value": round(progress), "unit": "%"}, "err_rms": {"value": err_rms, "unit": ""}}): keep_running = False break progress += d_progress*mini_batch.shape[0] # Update final status err_rms = rms_error_accu(batch - entity.reconstruct(entity.forward(batch))) status.on_change({"progress": {"value": round(progress), "unit": "%"}, "err_rms_total": {"value": err_rms, "unit": ""}}) return True