- 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 status import Status
|
||||||
from cd_train import cd_jens
|
from cd_train import cd_jens
|
||||||
|
|
||||||
class RbmLayer:
|
class Layer:
|
||||||
def __init__(self, name: str, num_visible, num_hidden, params: RbmParams):
|
def __init__(self, name: str, dim: tuple[int, int], params: RbmParams):
|
||||||
self.name = name
|
self.name = name
|
||||||
self.state = RbmState.from_layer_params(num_visible, num_hidden)
|
self.state = RbmState.from_layer_params(dim)
|
||||||
self.params = params
|
self.params = params
|
||||||
self.state_filename = f"{self.name}_state.npz"
|
self.state_filename = f"{self.name}_state.npz"
|
||||||
|
|
||||||
def init(self, std: float):
|
def init(self, std: float):
|
||||||
self.state.init(mu=0, std=std)
|
self.state.init(mu=0, std=std)
|
||||||
|
|
||||||
def save(self):
|
def save(self, filename: str = None):
|
||||||
self.state.to_file(self.state_filename)
|
if filename is None:
|
||||||
|
filename = self.state_filename
|
||||||
|
self.state.to_file(filename)
|
||||||
|
|
||||||
def load(self):
|
def load(self, filename: str = None):
|
||||||
state = RbmState.from_file(self.state_filename)
|
if filename is None:
|
||||||
|
filename = self.state_filename
|
||||||
|
state = RbmState.from_file(filename)
|
||||||
if state is not None:
|
if state is not None:
|
||||||
self.state = state
|
self.state = state
|
||||||
|
|
||||||
@@ -115,7 +119,7 @@ def xor():
|
|||||||
params.num_gibbs_samples = 3
|
params.num_gibbs_samples = 3
|
||||||
|
|
||||||
# Create layer
|
# Create layer
|
||||||
layer = RbmLayer("Layer_0", 3, 16, params)
|
layer = Layer("Layer_0", (3, 16), params)
|
||||||
|
|
||||||
# Init weights
|
# Init weights
|
||||||
layer.init(0.01)
|
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
|
self.b_h = b_h
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_layer_params(cls, num_visible, num_hidden):
|
def from_layer_params(cls, dim: tuple[int, int]):
|
||||||
w_hv = np.zeros((num_visible, num_hidden))
|
w_hv = np.zeros(dim)
|
||||||
b_v = np.zeros((1, num_visible))
|
b_v = np.zeros((1, dim[0]))
|
||||||
b_h = np.zeros((1, num_hidden))
|
b_h = np.zeros((1, dim[1]))
|
||||||
obj = cls(w_hv, b_v, b_h)
|
obj = cls(w_hv, b_v, b_h)
|
||||||
return obj
|
return obj
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user