- added stack

- refactored
This commit is contained in:
2025-12-16 21:36:55 +01:00
parent f33df4b0e0
commit 71b723e4e9
4 changed files with 77 additions and 12 deletions
+12 -8
View File
@@ -6,21 +6,25 @@ from helper import sample, prob, uniform, rms_error_accu
from status import Status
from cd_train import cd_jens
class RbmLayer:
def __init__(self, name: str, num_visible, num_hidden, params: RbmParams):
class Layer:
def __init__(self, name: str, dim: tuple[int, int], params: RbmParams):
self.name = name
self.state = RbmState.from_layer_params(num_visible, num_hidden)
self.state = RbmState.from_layer_params(dim)
self.params = params
self.state_filename = f"{self.name}_state.npz"
def init(self, std: float):
self.state.init(mu=0, std=std)
def save(self):
self.state.to_file(self.state_filename)
def save(self, filename: str = None):
if filename is None:
filename = self.state_filename
self.state.to_file(filename)
def load(self):
state = RbmState.from_file(self.state_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.state = state
@@ -115,7 +119,7 @@ def xor():
params.num_gibbs_samples = 3
# Create layer
layer = RbmLayer("Layer_0", 3, 16, params)
layer = Layer("Layer_0", (3, 16), params)
# Init weights
layer.init(0.01)