- added stack

- refactored
This commit is contained in:
2025-12-16 21:36:55 +01:00
parent f33df4b0e0
commit 71b723e4e9
4 changed files with 77 additions and 12 deletions
+12 -8
View File
@@ -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)
+20
View File
@@ -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!")
+41
View File
@@ -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")
+4 -4
View File
@@ -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