(self, D, M)
| 24 | self.session = session |
| 25 | |
| 26 | def build(self, D, M): |
| 27 | # params |
| 28 | self.W = tf.Variable(tf.random.normal(shape=(D, M)) * np.sqrt(2.0 / M)) |
| 29 | # note: without limiting variance, you get numerical stability issues |
| 30 | self.c = tf.Variable(np.zeros(M).astype(np.float32)) |
| 31 | self.b = tf.Variable(np.zeros(D).astype(np.float32)) |
| 32 | |
| 33 | # data |
| 34 | self.X_in = tf.compat.v1.placeholder(tf.float32, shape=(None, D)) |
| 35 | |
| 36 | # conditional probabilities |
| 37 | # NOTE: tf.contrib.distributions.Bernoulli API has changed in Tensorflow v1.2 |
| 38 | V = self.X_in |
| 39 | p_h_given_v = tf.nn.sigmoid(tf.matmul(V, self.W) + self.c) |
| 40 | self.p_h_given_v = p_h_given_v # save for later |
| 41 | # self.rng_h_given_v = tf.contrib.distributions.Bernoulli( |
| 42 | # probs=p_h_given_v, |
| 43 | # dtype=tf.float32 |
| 44 | # ) |
| 45 | r = tf.random.uniform(shape=tf.shape(input=p_h_given_v)) |
| 46 | H = tf.cast(r < p_h_given_v, dtype=tf.float32) |
| 47 | |
| 48 | p_v_given_h = tf.nn.sigmoid(tf.matmul(H, tf.transpose(a=self.W)) + self.b) |
| 49 | # self.rng_v_given_h = tf.contrib.distributions.Bernoulli( |
| 50 | # probs=p_v_given_h, |
| 51 | # dtype=tf.float32 |
| 52 | # ) |
| 53 | r = tf.random.uniform(shape=tf.shape(input=p_v_given_h)) |
| 54 | X_sample = tf.cast(r < p_v_given_h, dtype=tf.float32) |
| 55 | |
| 56 | |
| 57 | # build the objective |
| 58 | objective = tf.reduce_mean(input_tensor=self.free_energy(self.X_in)) - tf.reduce_mean(input_tensor=self.free_energy(X_sample)) |
| 59 | self.train_op = tf.compat.v1.train.AdamOptimizer(1e-2).minimize(objective) |
| 60 | # self.train_op = tf.train.GradientDescentOptimizer(1e-3).minimize(objective) |
| 61 | |
| 62 | # build the cost |
| 63 | # we won't use this to optimize the model parameters |
| 64 | # just to observe what happens during training |
| 65 | logits = self.forward_logits(self.X_in) |
| 66 | self.cost = tf.reduce_mean( |
| 67 | input_tensor=tf.nn.sigmoid_cross_entropy_with_logits( |
| 68 | labels=self.X_in, |
| 69 | logits=logits, |
| 70 | ) |
| 71 | ) |
| 72 | |
| 73 | def fit(self, X, epochs=1, batch_sz=100, show_fig=False): |
| 74 | N, D = X.shape |
no test coverage detected