Initial commit

This commit is contained in:
2025-12-16 14:52:36 +01:00
commit b6e94e75ac
11 changed files with 269 additions and 0 deletions
+41
View File
@@ -0,0 +1,41 @@
import numpy as np
from helper import uniform
class RbmState:
def __init__(self, w_hv: np.ndarray, b_v: np.ndarray, b_h: np.ndarray):
self.num_visible, self.num_hidden = w_hv.shape
self.w_hv = w_hv
self.b_v = b_v
self.b_h = b_h
@classmethod
def from_layer_params(cls, num_visible, num_hidden):
w_hv = np.zeros((num_visible, num_hidden))
b_v = np.zeros((1, num_visible))
b_h = np.zeros((1, num_hidden))
obj = cls(w_hv, b_v, b_h)
return obj
@classmethod
def from_file(cls, filename: str):
try:
with np.load(filename) as X:
w_hv, b_v, b_h = [X[i] for i in ('whv', 'bv', 'bh')]
print(f"{filename} loaded successfully!")
except FileNotFoundError:
pass
except KeyError:
pass
obj = cls(w_hv, b_v, b_h)
return obj
def to_file(self, filename: str):
np.savez(filename, whv=self.w_hv, bv=self.b_v, bh=self.b_h)
print(f"{filename} saved successfully!")
def init(self, mu: float = 0.5, std: float = 1.0):
self.w_hv = uniform(self.w_hv.shape, mu, std)
self.b_v = uniform(self.b_v.shape, 0, 0)
self.b_h = uniform(self.b_h.shape, 0, 0)