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
+24 -18
View File
@@ -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")