MCPcopy Create free account
hub / github.com/lazyprogrammer/machine_learning_examples / build

Method build

unsupervised_class2/rbm_tf.py:26–71  ·  view source on GitHub ↗
(self, D, M)

Source from the content-addressed store, hash-verified

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

Callers 1

__init__Method · 0.95

Calls 2

free_energyMethod · 0.95
forward_logitsMethod · 0.95

Tested by

no test coverage detected