- state load/save revised
- num gibbs samples is also part of entity - working 3-layer deep test model
This commit is contained in:
+8
-3
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user