revised stack

This commit is contained in:
2025-12-17 11:31:13 +01:00
parent 80ba543323
commit 3ec170234b
3 changed files with 33 additions and 25 deletions
+7 -5
View File
@@ -9,14 +9,16 @@ if __name__ == "__main__":
for n, dim in enumerate(dims): for n, dim in enumerate(dims):
layer = Layer(f"Layer-{n}", dim, params) layer = Layer(f"Layer-{n}", dim, params)
stack.layer_add(layer) stack.append(layer)
stack.init(std=0.1) stack.init(std=0.1)
stack.state_save() stack.state_save()
try: lay0 = stack.layers[0]
stack.layer_add(Layer(f"Layer-{0}", (9,9), params)) lay1 = stack.from_name("Layer-1")
except StackException as e: lay2 = stack.from_index(2)
print(f"Exception raised: [{e}] -> success!") stack.remove(lay2)
stack.remove(lay1)
stack.remove(lay0)
print("Test: [passed]") print("Test: [passed]")
+24 -18
View File
@@ -12,30 +12,36 @@ class Stack:
def __init__(self, stack_type: StackType, name: str): def __init__(self, stack_type: StackType, name: str):
self.stack_type = stack_type self.stack_type = stack_type
self.name = name self.name = name
self.layers: dict[str, Layer] = {} self.layers: list[Layer] = []
def layer_add(self, layer: Layer): def num_layers(self):
if layer.name in self.layers.keys(): return len(self.layers)
raise StackException(f"Layer \"{layer.name}\" already exists")
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): def remove(self, layer: Layer):
if layer.name not in self.layers.keys(): for index, lay in enumerate(self.layers):
raise StackException(f"Layer \"{layer.name}\" does not exist") if lay == layer:
del(self.layers[index])
return
raise StackException(f"Layer \"{layer.name}\" does not exist")
del(self.layers[layer.name]) def from_index(self, index : int) -> Layer:
result = self.layers[index]
def layer_from_name(self, name: str) -> Layer:
result = None
if name in self.layers.keys():
result = self.layers[name]
return result 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): def init(self, std: float):
for k in self.layers.keys(): for layer in self.layers:
self.layers[k].init(std) layer.init(std)
def state_save(self): def state_save(self):
for k in self.layers.keys(): for index, layer in enumerate(self.layers):
self.layers[k].save(f"{self.name}-{k}-state.npz") layer.save(f"{self.name}-{index}-state.npz")
+2 -2
View File
@@ -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) layer_obj = Layer(f"{layer_name}-{layer_id}", (num_visible_x*num_visible_y+num_context, num_hidden), params)
# Add layer to stack # Add layer to stack
obj.layer_add(layer_obj) obj.append(layer_obj)
return obj return obj
@@ -66,6 +66,6 @@ class StackFactory:
if __name__ == "__main__": 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]") print("Test: [passed]")