- added gibbs sampling for forward pass
This commit is contained in:
+31
-3
@@ -13,6 +13,9 @@ class RbmLayer:
|
|||||||
self.params = params
|
self.params = params
|
||||||
self.state_filename = f"{self.name}_state.npz"
|
self.state_filename = f"{self.name}_state.npz"
|
||||||
|
|
||||||
|
def init(self):
|
||||||
|
self.state.init()
|
||||||
|
|
||||||
def save(self):
|
def save(self):
|
||||||
self.state.to_file(self.state_filename)
|
self.state.to_file(self.state_filename)
|
||||||
|
|
||||||
@@ -89,6 +92,22 @@ class RbmLayer:
|
|||||||
return prob(state)
|
return prob(state)
|
||||||
|
|
||||||
|
|
||||||
|
def gibbs_v_to_h(self, v: np.ndarray) -> np.ndarray:
|
||||||
|
h = None
|
||||||
|
for i in range(self.params.num_gibbs_samples):
|
||||||
|
h = self.v_to_ph(v)
|
||||||
|
v = self.h_to_pv(h)
|
||||||
|
|
||||||
|
return h
|
||||||
|
|
||||||
|
def gibbs_h_to_v(self, h: np.ndarray) -> np.ndarray:
|
||||||
|
v = None
|
||||||
|
for i in range(self.params.num_gibbs_samples):
|
||||||
|
v = self.h_to_pv(h)
|
||||||
|
h = self.v_to_ph(v)
|
||||||
|
|
||||||
|
return v
|
||||||
|
|
||||||
def xor():
|
def xor():
|
||||||
params = RbmParams()
|
params = RbmParams()
|
||||||
params.do_rao_blackwell = True
|
params.do_rao_blackwell = True
|
||||||
@@ -96,15 +115,24 @@ def xor():
|
|||||||
params.mini_batch_size = 100
|
params.mini_batch_size = 100
|
||||||
params.learning_rate = 0.1
|
params.learning_rate = 0.1
|
||||||
params.momentum = 0.5
|
params.momentum = 0.5
|
||||||
params.num_epochs = 10000
|
params.num_epochs = 100
|
||||||
|
|
||||||
status = Status()
|
status = Status()
|
||||||
layer = RbmLayer("Layer_0", 3, 16, params)
|
layer = RbmLayer("Layer_0", 3, 16, params)
|
||||||
training_batch = np.array([[0,0,0], [0,1,1], [1,0,1], [1,1,0]], dtype=np.float64)
|
layer.init()
|
||||||
test_batch = np.array([[0,0,0], [0,1,0], [1,0,0], [1,1,0]], dtype=np.float64)
|
|
||||||
|
|
||||||
|
# Train
|
||||||
|
training_batch = np.array([[0,0,0], [0,1,1], [1,0,1], [1,1,0]], dtype=np.float64)
|
||||||
layer.train(training_batch, cd_jens, status)
|
layer.train(training_batch, cd_jens, status)
|
||||||
|
|
||||||
|
# Test
|
||||||
|
test_batch = np.array([[0,0,0], [0,1,0], [1,0,0], [1,1,0]], dtype=np.float64)
|
||||||
|
for pattern in test_batch:
|
||||||
|
h = layer.gibbs_v_to_h(pattern)
|
||||||
|
v = layer.gibbs_h_to_v(h)
|
||||||
|
print(f"P{pattern} : {v}")
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
xor()
|
xor()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user