Files
pyRBM/src/rbm/entity.py
T
jens 4bd5cffd6b - state load/save revised
- num gibbs samples is also part of entity
- working 3-layer deep test model
2025-12-21 17:34:11 +01:00

84 lines
2.4 KiB
Python

from .state import RbmState
from .matrix import prob, Mat
class EntityParams:
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):
obj = EntityParams()
if "doGaussianVisible" in params:
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
class Entity:
def __init__(self, shape: tuple[int, int], params: EntityParams):
self.shape = shape
self.params = params
self.state = RbmState.from_layer_params(shape)
self.grad = RbmState.from_layer_params(shape)
def __call__(self, x: Mat):
return self.forward(x)
def grad_zero(self):
self.grad = RbmState.from_layer_params(self.shape)
def grad_compute(self, d_bv: Mat, d_bh: Mat, d_whv: Mat, learning_rate: float, momentum: float, weight_decay: float):
# Compute gradient
self.grad.b_v = (momentum * self.grad.b_v + learning_rate * d_bv)
self.grad.b_h = (momentum * self.grad.b_h + learning_rate * d_bh)
self.grad.w_hv = (momentum * self.grad.w_hv + learning_rate * d_whv - weight_decay * self.state.w_hv)
return self.grad
def state_adjust(self, grad: RbmState):
# Adjust state
self.state.b_v += grad.b_v
self.state.b_h += grad.b_h
self.state.w_hv += grad.w_hv
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)
h = self._v_to_ph(h)
return h
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)
v = self._h_to_pv(v)
return v
def _v_to_ph(self, v: Mat) -> Mat:
state = self.state.v_to_h(v)
if self.params.do_gaussian_hidden:
return state
return prob(state)
def _h_to_pv(self, h: Mat) -> Mat:
state = self.state.h_to_v(h)
if self.params.do_gaussian_visible:
return state
return prob(state)
if __name__ == "__main__":
print("Test: [passed]")