refactored
This commit is contained in:
@@ -10,19 +10,14 @@ class Layer:
|
||||
self.name = name
|
||||
self.shape = shape
|
||||
self.entity = Entity((shape[0]*shape[1]+shape[2], shape[3]), params)
|
||||
self.state_filename = f"{self.name}_state.npz"
|
||||
|
||||
def init(self, std: float):
|
||||
self.entity.state.init(mu=0, std=std)
|
||||
|
||||
def save(self, filename: str = None):
|
||||
if filename is None:
|
||||
filename = self.state_filename
|
||||
self.entity.state.to_file(filename)
|
||||
|
||||
def load(self, filename: str = None):
|
||||
if filename is None:
|
||||
filename = self.state_filename
|
||||
state = RbmState.from_file(filename)
|
||||
if state is not None:
|
||||
self.entity.state = state
|
||||
|
||||
Reference in New Issue
Block a user