revised stack
This commit is contained in:
+7
-5
@@ -9,14 +9,16 @@ if __name__ == "__main__":
|
||||
|
||||
for n, dim in enumerate(dims):
|
||||
layer = Layer(f"Layer-{n}", dim, params)
|
||||
stack.layer_add(layer)
|
||||
stack.append(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!")
|
||||
lay0 = stack.layers[0]
|
||||
lay1 = stack.from_name("Layer-1")
|
||||
lay2 = stack.from_index(2)
|
||||
stack.remove(lay2)
|
||||
stack.remove(lay1)
|
||||
stack.remove(lay0)
|
||||
|
||||
print("Test: [passed]")
|
||||
|
||||
+24
-18
@@ -12,30 +12,36 @@ class Stack:
|
||||
def __init__(self, stack_type: StackType, name: str):
|
||||
self.stack_type = stack_type
|
||||
self.name = name
|
||||
self.layers: dict[str, Layer] = {}
|
||||
self.layers: list[Layer] = []
|
||||
|
||||
def layer_add(self, layer: Layer):
|
||||
if layer.name in self.layers.keys():
|
||||
raise StackException(f"Layer \"{layer.name}\" already exists")
|
||||
def num_layers(self):
|
||||
return len(self.layers)
|
||||
|
||||
self.layers[layer.name] = layer
|
||||
def append(self, layer: Layer):
|
||||
self.layers.append(layer)
|
||||
return self.num_layers()-1
|
||||
|
||||
def layer_remove(self, layer: Layer):
|
||||
if layer.name not in self.layers.keys():
|
||||
raise StackException(f"Layer \"{layer.name}\" does not exist")
|
||||
def remove(self, layer: Layer):
|
||||
for index, lay in enumerate(self.layers):
|
||||
if lay == layer:
|
||||
del(self.layers[index])
|
||||
return
|
||||
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]
|
||||
def from_index(self, index : int) -> Layer:
|
||||
result = self.layers[index]
|
||||
return result
|
||||
|
||||
def from_name(self, name: str):
|
||||
for layer in self.layers:
|
||||
if layer.name in name:
|
||||
return layer
|
||||
return None
|
||||
|
||||
def init(self, std: float):
|
||||
for k in self.layers.keys():
|
||||
self.layers[k].init(std)
|
||||
for layer in self.layers:
|
||||
layer.init(std)
|
||||
|
||||
def state_save(self):
|
||||
for k in self.layers.keys():
|
||||
self.layers[k].save(f"{self.name}-{k}-state.npz")
|
||||
for index, layer in enumerate(self.layers):
|
||||
layer.save(f"{self.name}-{index}-state.npz")
|
||||
|
||||
@@ -51,7 +51,7 @@ class StackFactory:
|
||||
layer_obj = Layer(f"{layer_name}-{layer_id}", (num_visible_x*num_visible_y+num_context, num_hidden), params)
|
||||
|
||||
# Add layer to stack
|
||||
obj.layer_add(layer_obj)
|
||||
obj.append(layer_obj)
|
||||
|
||||
return obj
|
||||
|
||||
@@ -66,6 +66,6 @@ class StackFactory:
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
stack = StackFactory.from_file("/home/jens/work/repos/Rbm/many.prj")
|
||||
stack = StackFactory.from_file("/home/jens/work/repos/Rbm/test.prj")
|
||||
|
||||
print("Test: [passed]")
|
||||
|
||||
Reference in New Issue
Block a user