- added stack
- refactored
This commit is contained in:
+12
-8
@@ -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)
|
||||
|
||||
@@ -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!")
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user