- added gibbs sampling for forward pass
This commit is contained in:
+31
-3
@@ -13,6 +13,9 @@ class RbmLayer:
|
||||
self.params = params
|
||||
self.state_filename = f"{self.name}_state.npz"
|
||||
|
||||
def init(self):
|
||||
self.state.init()
|
||||
|
||||
def save(self):
|
||||
self.state.to_file(self.state_filename)
|
||||
|
||||
@@ -89,6 +92,22 @@ class RbmLayer:
|
||||
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():
|
||||
params = RbmParams()
|
||||
params.do_rao_blackwell = True
|
||||
@@ -96,15 +115,24 @@ def xor():
|
||||
params.mini_batch_size = 100
|
||||
params.learning_rate = 0.1
|
||||
params.momentum = 0.5
|
||||
params.num_epochs = 10000
|
||||
params.num_epochs = 100
|
||||
|
||||
status = Status()
|
||||
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)
|
||||
test_batch = np.array([[0,0,0], [0,1,0], [1,0,0], [1,1,0]], dtype=np.float64)
|
||||
layer.init()
|
||||
|
||||
# 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)
|
||||
|
||||
# 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__":
|
||||
xor()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user