- state load/save revised

- num gibbs samples is also part of entity
- working 3-layer deep test model
This commit is contained in:
2025-12-21 17:34:11 +01:00
parent 22d189cf2f
commit 4bd5cffd6b
5 changed files with 57 additions and 30 deletions
+8 -3
View File
@@ -2,10 +2,11 @@ from .state import RbmState
from .matrix import prob, Mat
class EntityParams:
def __init__(self, do_gaussian_visible: bool = False, do_gaussian_hidden: bool = False):
def __init__(self, do_gaussian_visible: bool = False, do_gaussian_hidden: bool = False, num_gibbs_samples: int = 1):
# Entity parameters
self.do_gaussian_visible = do_gaussian_visible
self.do_gaussian_hidden = do_gaussian_hidden
self.num_gibbs_samples = num_gibbs_samples
@classmethod
def from_dict(cls, params: dict):
@@ -14,6 +15,8 @@ class EntityParams:
obj.do_gaussian_visible = params["doGaussianVisible"]
if "doGaussianHidden" in params:
obj.do_gaussian_hidden = params["doGaussianHidden"]
if "num_gibbs_samples" in params:
obj.num_gibbs_samples = params["num_gibbs_samples"]
return obj
@@ -44,7 +47,8 @@ class Entity:
self.state.b_h += grad.b_h
self.state.w_hv += grad.w_hv
def forward(self, v: Mat, num_gibbs: int = 1) -> Mat:
def forward(self, v: Mat, num_gibbs: int = 0) -> Mat:
num_gibbs = self.params.num_gibbs_samples if num_gibbs == 0 else num_gibbs
h = self._v_to_ph(v)
for i in range(num_gibbs-1):
h = self._h_to_pv(h)
@@ -52,7 +56,8 @@ class Entity:
return h
def reconstruct(self, h: Mat, num_gibbs: int = 1) -> Mat:
def reconstruct(self, h: Mat, num_gibbs: int = 0) -> Mat:
num_gibbs = self.params.num_gibbs_samples if num_gibbs == 0 else num_gibbs
v = self._h_to_pv(h)
for i in range(num_gibbs-1):
v = self._v_to_ph(v)