From 71b723e4e9996f1af9b66eb214396ea3bb020cc9 Mon Sep 17 00:00:00 2001 From: Jens Ahrensfeld Date: Tue, 16 Dec 2025 21:36:55 +0100 Subject: [PATCH] - added stack - refactored --- src/rbm/layer.py | 20 ++++++++++++-------- src/rbm/rbm.py | 20 ++++++++++++++++++++ src/rbm/stack.py | 41 +++++++++++++++++++++++++++++++++++++++++ src/rbm/state.py | 8 ++++---- 4 files changed, 77 insertions(+), 12 deletions(-) create mode 100644 src/rbm/stack.py diff --git a/src/rbm/layer.py b/src/rbm/layer.py index 091c23b..311dcc6 100644 --- a/src/rbm/layer.py +++ b/src/rbm/layer.py @@ -6,21 +6,25 @@ from helper import sample, prob, uniform, rms_error_accu from status import Status from cd_train import cd_jens -class RbmLayer: - def __init__(self, name: str, num_visible, num_hidden, params: RbmParams): +class Layer: + def __init__(self, name: str, dim: tuple[int, int], params: RbmParams): self.name = name - self.state = RbmState.from_layer_params(num_visible, num_hidden) + 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): - self.state.to_file(self.state_filename) + def save(self, filename: str = None): + if filename is None: + filename = self.state_filename + self.state.to_file(filename) - def load(self): - state = RbmState.from_file(self.state_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 @@ -115,7 +119,7 @@ def xor(): params.num_gibbs_samples = 3 # Create layer - layer = RbmLayer("Layer_0", 3, 16, params) + layer = Layer("Layer_0", (3, 16), params) # Init weights layer.init(0.01) diff --git a/src/rbm/rbm.py b/src/rbm/rbm.py index e69de29..06ac3ae 100644 --- a/src/rbm/rbm.py +++ b/src/rbm/rbm.py @@ -0,0 +1,20 @@ +from layer import Layer +from stack import Stack, StackType, StackException +from params import RbmParams + +if __name__ == "__main__": + params = RbmParams() + stack = Stack(StackType.Deep, "Stack") + dims = [(2, 3), (3, 4), (4, 5)] + + for n, dim in enumerate(dims): + layer = Layer(f"Layer-{n}", dim, params) + stack.layer_add(layer) + + stack.init(std=0.1) + stack.state_save() + + try: + stack.layer_add(Layer(f"Layer-{0}", (9,9), params)) + except StackException as e: + print(f"Exception raised: [{e}] -> success!") \ No newline at end of file diff --git a/src/rbm/stack.py b/src/rbm/stack.py new file mode 100644 index 0000000..318958e --- /dev/null +++ b/src/rbm/stack.py @@ -0,0 +1,41 @@ +from enum import Enum +from layer import Layer + +class StackType(Enum): + Deep = "Deep", + Rnn = "Rnn" + +class StackException(Exception): + pass + +class Stack: + def __init__(self, stack_type: StackType, name: str): + self.stack_type = stack_type + self.name = name + self.layers: dict[str, Layer] = {} + + def layer_add(self, layer: Layer): + if layer.name in self.layers.keys(): + raise StackException(f"Layer \"{layer.name}\" already exists") + + self.layers[layer.name] = layer + + def layer_remove(self, layer: Layer): + if layer.name not in self.layers.keys(): + raise StackException(f"Layer \"{layer.name}\" does not exist") + + del(self.layers[layer.name]) + + def layer_from_name(self, name: str) -> Layer: + result = None + if name in self.layers.keys(): + result = self.layers[name] + return result + + def init(self, std: float): + for k in self.layers.keys(): + self.layers[k].init(std) + + def state_save(self): + for k in self.layers.keys(): + self.layers[k].save(f"{self.name}-{k}-state.npz") diff --git a/src/rbm/state.py b/src/rbm/state.py index e9bedb0..9e757cb 100644 --- a/src/rbm/state.py +++ b/src/rbm/state.py @@ -10,10 +10,10 @@ class RbmState: self.b_h = b_h @classmethod - def from_layer_params(cls, num_visible, num_hidden): - w_hv = np.zeros((num_visible, num_hidden)) - b_v = np.zeros((1, num_visible)) - b_h = np.zeros((1, num_hidden)) + def from_layer_params(cls, dim: tuple[int, int]): + w_hv = np.zeros(dim) + b_v = np.zeros((1, dim[0])) + b_h = np.zeros((1, dim[1])) obj = cls(w_hv, b_v, b_h) return obj