refactored

This commit is contained in:
2025-12-16 14:59:07 +01:00
parent b6e94e75ac
commit 944b96448c
3 changed files with 11 additions and 11 deletions
+2 -8
View File
@@ -10,21 +10,15 @@ class CdTrain:
self.state = state
self.params = params
def v_to_h(self, visible: np.ndarray) -> np.ndarray:
return np.vecmat(visible, self.state.w_hv) + self.state.b_h
def h_to_v(self, hidden: np.ndarray) -> np.ndarray:
return np.vecmat(hidden, np.transpose(self.state.w_hv)) + self.state.b_v
def v_to_h_prob(self, v: np.ndarray) -> np.ndarray:
state = self.v_to_h(v)
state = self.state.v_to_h(v)
if self.params.do_gaussian_visible:
return state
return prob(state)
def h_to_v_prob(self, h: np.ndarray) -> np.ndarray:
state = self.h_to_v(h)
state = self.state.h_to_v(h)
if self.params.do_gaussian_visible:
return state
+3 -3
View File
@@ -20,7 +20,7 @@ class RbmLayer:
def train_batch(self, v_states: np.ndarray, cd_func: Callable):
for epochs in range(self.params.num_epochs):
# Contrastive divergence learning: calculate gradients
cd_func(v_states, dwhv, dbh, dbv);
dwhv, dbh, dbv = cd_func(v_states)
@@ -36,9 +36,9 @@ if __name__ == "__main__":
v0 = uniform((1,2))
h0 = uniform((1,3))
h1 = l1.v_to_h(v0)
h1 = l1.state.v_to_h(v0)
print(h1.shape)
v1 = l1.h_to_v(h1)
v1 = l1.state.h_to_v(h1)
s_v1 = sample(v1)
p_h = prob(s_v1)
+6
View File
@@ -31,6 +31,12 @@ class RbmState:
obj = cls(w_hv, b_v, b_h)
return obj
def v_to_h(self, visible: np.ndarray) -> np.ndarray:
return np.vecmat(visible, self.w_hv) + self.b_h
def h_to_v(self, hidden: np.ndarray) -> np.ndarray:
return np.vecmat(hidden, np.transpose(self.w_hv)) + self.b_v
def to_file(self, filename: str):
np.savez(filename, whv=self.w_hv, bv=self.b_v, bh=self.b_h)
print(f"{filename} saved successfully!")