diff --git a/src/rbm/layer.py b/src/rbm/layer.py index 215826d..091c23b 100644 --- a/src/rbm/layer.py +++ b/src/rbm/layer.py @@ -20,25 +20,26 @@ class RbmLayer: self.state.to_file(self.state_filename) def load(self): - self.state = RbmState.from_file(self.state_filename) + state = RbmState.from_file(self.state_filename) + if state is not None: + self.state = state def train(self, batch: np.ndarray, cd_func: Callable, status: Status): - training_size = batch.shape[0] - num_cases = min(self.params.mini_batch_size, training_size) - d_progress = 100.0 / (training_size/num_cases * self.params.num_epochs) - progress = 0 - last_progress = progress + 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 - - keep_running = True - - training_remain = training_size training_seen = 0 + keep_running = True while training_remain > 0 and keep_running: - mini_batch_size = min(self.params.mini_batch_size, training_remain) - mini_batch = batch[batch_row_index:batch_row_index + mini_batch_size] - training_remain -= mini_batch_size - batch_row_index += mini_batch_size + 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) @@ -53,7 +54,7 @@ class RbmLayer: 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/num_cases + 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 @@ -62,8 +63,7 @@ class RbmLayer: self.state.b_h += inc_bh self.state.w_hv += inc_whv - progress = round(epochs*d_progress) - + progress = round(training_seen*d_progress) if progress != last_progress: # Calculate error status.progress = round(progress) @@ -73,7 +73,7 @@ class RbmLayer: break last_progress = progress - training_seen += 1 + training_seen += 1 status.on_change() @@ -113,10 +113,6 @@ def xor(): params = RbmParams() params.do_rao_blackwell = True params.num_gibbs_samples = 3 - params.mini_batch_size = 100 - params.learning_rate = 0.1 - params.momentum = 0.5 - params.num_epochs = 1000 # Create layer layer = RbmLayer("Layer_0", 3, 16, params) diff --git a/src/rbm/params.py b/src/rbm/params.py index b41d1b5..1c2a875 100644 --- a/src/rbm/params.py +++ b/src/rbm/params.py @@ -11,12 +11,12 @@ class Params: class RbmParams(Params): def __init__(self): - self.learning_rate = 0 - self.momentum = 0 + self.learning_rate = 0.1 + self.momentum = 0.5 self.weight_decay = 0 - self.num_epochs = 1 + self.num_epochs = 1000 self.num_gibbs_samples = 1 - self.mini_batch_size = 1 + self.mini_batch_size = 0 self.do_gaussian_visible = False self.do_gaussian_hidden = False diff --git a/src/rbm/state.py b/src/rbm/state.py index cc90f74..e9bedb0 100644 --- a/src/rbm/state.py +++ b/src/rbm/state.py @@ -19,16 +19,17 @@ class RbmState: @classmethod def from_file(cls, filename: str): + obj = None try: with np.load(filename) as X: w_hv, b_v, b_h = [X[i] for i in ('whv', 'bv', 'bh')] + obj = cls(w_hv, b_v, b_h) print(f"{filename} loaded successfully!") except FileNotFoundError: pass except KeyError: pass - obj = cls(w_hv, b_v, b_h) return obj def v_to_h(self, visible: np.ndarray) -> np.ndarray: