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