import numpy as np from collections.abc import Callable from params import RbmParams from state import RbmState from helper import sample, prob, uniform, rms_error_accu from status import Status from cd_train import cd_jens class Layer: def __init__(self, name: str, dim: tuple[int, int], params: RbmParams): self.name = name self.state = RbmState.from_layer_params(dim) self.params = params self.state_filename = f"{self.name}_state.npz" def init(self, std: float): self.state.init(mu=0, std=std) def save(self, filename: str = None): if filename is None: filename = self.state_filename self.state.to_file(filename) def load(self, filename: str = None): if filename is None: filename = self.state_filename state = RbmState.from_file(filename) if state is not None: self.state = state def train(self, batch: np.ndarray, cd_func: Callable, status: Status): training_remain = batch.shape[0] batch_size = min(self.params.mini_batch_size, training_remain) if batch_size == 0: batch_size = training_remain d_progress = 100.0 / (training_remain/batch_size * self.params.num_epochs) last_progress = 0 batch_row_index = 0 training_seen = 0 keep_running = True while training_remain > 0 and keep_running: batch_size_remain = min(batch_size, training_remain) mini_batch = batch[batch_row_index:batch_row_index + batch_size_remain] training_remain -= batch_size_remain batch_row_index += batch_size_remain inc_bv = np.zeros(self.state.b_v.shape) inc_bh = np.zeros(self.state.b_h.shape) inc_whv = np.zeros(self.state.w_hv.shape) v_states = mini_batch if self.params.do_batch_sample: v_states = sample(mini_batch) for epochs in range(self.params.num_epochs): # Contrastive divergence learning: calculate gradients dwhv, dbv, dbh = cd_func(v_states, self.params, self.v_to_ph, self.h_to_pv) # Adjust weight and biases kl = self.params.learning_rate/batch_size inc_bv = self.params.momentum*inc_bv + kl*dbv inc_bh = self.params.momentum*inc_bh + kl*dbh inc_whv = self.params.momentum*inc_whv + kl*dwhv - self.params.weight_decay*self.state.w_hv self.state.b_v += inc_bv self.state.b_h += inc_bh self.state.w_hv += inc_whv progress = round(training_seen*d_progress) if progress != last_progress: # Calculate error status.progress = round(progress) status.err = rms_error_accu(mini_batch - self.h_to_pv(self.v_to_ph(v_states))) if not status.on_change(): keep_running = False break last_progress = progress training_seen += 1 status.on_change() def v_to_ph(self, v: np.ndarray) -> np.ndarray: state = self.state.v_to_h(v) if self.params.do_gaussian_visible: return state return prob(state) def h_to_pv(self, h: np.ndarray) -> np.ndarray: state = self.state.h_to_v(h) if self.params.do_gaussian_visible: return state return prob(state) def gibbs_v_to_h(self, v: np.ndarray) -> np.ndarray: h = None for i in range(self.params.num_gibbs_samples): h = self.v_to_ph(v) v = self.h_to_pv(h) return h def gibbs_h_to_v(self, h: np.ndarray) -> np.ndarray: v = None for i in range(self.params.num_gibbs_samples): v = self.h_to_pv(h) h = self.v_to_ph(v) return v def xor(): # Create params params = RbmParams() params.do_rao_blackwell = True params.num_gibbs_samples = 3 # Create layer layer = Layer("Layer_0", (3, 16), params) # Init weights layer.init(0.01) # Load weights (if exists) layer.load() # Prepare training data training_batch = np.array([[0,1,1], [0,0,0], [1,1,0], [1,0,1]], dtype=np.float64) # Train layer layer.train(training_batch, cd_jens, Status()) # Save weights layer.save() # Test with test data test_batch = np.array([[0,0,0], [0,1,0], [1,0,0], [1,1,0]], dtype=np.float64) for pattern in test_batch: h = layer.gibbs_v_to_h(pattern) v = layer.gibbs_h_to_v(h) print(f"P{pattern} : {v}") if __name__ == "__main__": xor()