refactored
This commit is contained in:
+21
-3
@@ -1,7 +1,25 @@
|
||||
from .params import EntityParams
|
||||
from .state import RbmState
|
||||
from .matrix import prob, Mat
|
||||
|
||||
class EntityParams:
|
||||
def __init__(self):
|
||||
# Entity parameters
|
||||
self.num_gibbs_samples = 1
|
||||
self.do_gaussian_visible = False
|
||||
self.do_gaussian_hidden = False
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, params: dict):
|
||||
obj = EntityParams()
|
||||
if "numGibbs" in params:
|
||||
obj.num_gibbs_samples = params["numGibbs"]
|
||||
if "doGaussianVisible" in params:
|
||||
obj.do_gaussian_visible = params["doGaussianVisible"]
|
||||
if "doGaussianHidden" in params:
|
||||
obj.do_gaussian_hidden = params["doGaussianHidden"]
|
||||
|
||||
return obj
|
||||
|
||||
class Entity:
|
||||
def __init__(self, shape: tuple[int, int], params: EntityParams):
|
||||
self.shape = shape
|
||||
@@ -22,7 +40,7 @@ class Entity:
|
||||
|
||||
return prob(state)
|
||||
|
||||
def gibbs_v_to_h(self, v: Mat) -> Mat:
|
||||
def forward(self, v: Mat) -> Mat:
|
||||
h = self.v_to_ph(v)
|
||||
for i in range(self.params.num_gibbs_samples-1):
|
||||
h = self.h_to_pv(h)
|
||||
@@ -30,7 +48,7 @@ class Entity:
|
||||
|
||||
return h
|
||||
|
||||
def gibbs_h_to_v(self, h: Mat) -> Mat:
|
||||
def reconstruct(self, h: Mat) -> Mat:
|
||||
v = self.h_to_pv(h)
|
||||
for i in range(self.params.num_gibbs_samples-1):
|
||||
v = self.v_to_ph(v)
|
||||
|
||||
Reference in New Issue
Block a user