refactored

This commit is contained in:
2025-12-19 15:20:04 +01:00
parent a293fc31a0
commit a47922cb1c
8 changed files with 92 additions and 124 deletions
+21 -3
View File
@@ -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)