From 5cbef6b09fa6625e3e164613de9bf05e68387dd4 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Thu, 18 Dec 2025 12:24:17 +0100 Subject: [PATCH] refactored --- src/rbm/entity.py | 94 +++++++++++++++++++++++++++++++ src/rbm/layer.py | 98 +++------------------------------ src/rbm/{rbm.py => test_rbm.py} | 0 3 files changed, 102 insertions(+), 90 deletions(-) create mode 100644 src/rbm/entity.py rename src/rbm/{rbm.py => test_rbm.py} (100%) diff --git a/src/rbm/entity.py b/src/rbm/entity.py new file mode 100644 index 0000000..abf12de --- /dev/null +++ b/src/rbm/entity.py @@ -0,0 +1,94 @@ +import numpy as np +from collections.abc import Callable +from params import RbmParams +from state import RbmState +from matrix import sample, prob, rms_error_accu +from status import Status + +class Entity: + def __init__(self, shape: tuple[int, int], params: RbmParams): + self.state = RbmState.from_layer_params(shape) + self.params = params + + 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 + + status.on_change({}) + d_progress = 100.0 / (training_remain/batch_size * self.params.num_epochs) + 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 + + # Calculate error + if status.want_report(round(training_seen*d_progress)): + err_rms = rms_error_accu(mini_batch - self.h_to_pv(self.v_to_ph(v_states))) + if not status.on_change({"progress": {"value": round(training_seen * d_progress), "unit": "%"}, + "err_rms": {"value": err_rms, "unit": ""}}): + keep_running = False + break + + training_seen += 1 + + 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 + +if __name__ == "__main__": + print("Test: [passed]") diff --git a/src/rbm/layer.py b/src/rbm/layer.py index 947f055..080173f 100644 --- a/src/rbm/layer.py +++ b/src/rbm/layer.py @@ -1,113 +1,31 @@ import numpy as np -from collections.abc import Callable from params import RbmParams from state import RbmState -from matrix import sample, prob, rms_error_accu from status import Status from cd_train import cd_jens +from entity import Entity class Layer: def __init__(self, name: str, shape: tuple[int, int, int, int], params: RbmParams): self.name = name self.shape = shape - self.state = RbmState.from_layer_params((shape[0]*shape[1]+shape[2], shape[3])) - self.params = params + self.entity = Entity((shape[0]*shape[1]+shape[2], shape[3]), params) self.state_filename = f"{self.name}_state.npz" def init(self, std: float): - self.state.init(mu=0, std=std) + self.entity.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) + self.entity.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 - - status.on_change({}) - d_progress = 100.0 / (training_remain/batch_size * self.params.num_epochs) - 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 - - # Calculate error - if status.want_report(round(training_seen*d_progress)): - err_rms = rms_error_accu(mini_batch - self.h_to_pv(self.v_to_ph(v_states))) - if not status.on_change({"progress": {"value": round(training_seen * d_progress), "unit": "%"}, - "err_rms": {"value": err_rms, "unit": ""}}): - keep_running = False - break - - training_seen += 1 - - 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 + self.entity.state = state def xor(): # Create params @@ -128,7 +46,7 @@ def xor(): 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()) + layer.entity.train(training_batch, cd_jens, Status()) # Save weights layer.save() @@ -136,8 +54,8 @@ def xor(): # 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) + h = layer.entity.gibbs_v_to_h(pattern) + v = layer.entity.gibbs_h_to_v(h) print(f"P{pattern} : {v}") diff --git a/src/rbm/rbm.py b/src/rbm/test_rbm.py similarity index 100% rename from src/rbm/rbm.py rename to src/rbm/test_rbm.py