MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / run_training

Function run_training

tensorflow/contrib/factorization/examples/mnist.py:183–276  ·  view source on GitHub ↗

Train MNIST for a number of steps.

()

Source from the content-addressed store, hash-verified

181
182
183def run_training():
184 """Train MNIST for a number of steps."""
185 # Get the sets of images and labels for training, validation, and
186 # test on MNIST.
187 train_dir = tempfile.mkdtemp()
188 data_sets = input_data.read_data_sets(train_dir, FLAGS.fake_data)
189
190 # Tell TensorFlow that the model will be built into the default Graph.
191 with tf.Graph().as_default():
192 # Generate placeholders for the images and labels.
193 images_placeholder, labels_placeholder = placeholder_inputs()
194
195 # Build a Graph that computes predictions from the inference model.
196 logits, clustering_loss, kmeans_init, kmeans_training_op = inference(
197 images_placeholder,
198 FLAGS.num_clusters,
199 FLAGS.hidden1,
200 FLAGS.hidden2)
201
202 # Add to the Graph the Ops for loss calculation.
203 loss = mnist.loss(logits, labels_placeholder)
204
205 # Add to the Graph the Ops that calculate and apply gradients.
206 train_op = tf.group(mnist.training(loss, FLAGS.learning_rate),
207 kmeans_training_op)
208
209 # Add the Op to compare the logits to the labels during evaluation.
210 eval_correct = mnist.evaluation(logits, labels_placeholder)
211
212 # Add the variable initializer Op.
213 init = tf.global_variables_initializer()
214
215 # Create a session for running Ops on the Graph.
216 sess = tf.Session()
217
218 # Run the Op to initialize the variables.
219 sess.run(init)
220
221 feed_dict = fill_feed_dict(data_sets.train,
222 images_placeholder,
223 labels_placeholder,
224 batch_size=max(FLAGS.batch_size, 5000))
225 # Run the Op to initialize the clusters.
226 sess.run(kmeans_init, feed_dict=feed_dict)
227
228 # Start the training loop.
229 max_test_prec = 0
230 for step in xrange(FLAGS.max_steps):
231 start_time = time.time()
232
233 # Fill a feed dictionary with the actual set of images and labels
234 # for this particular training step.
235 feed_dict = fill_feed_dict(data_sets.train,
236 images_placeholder,
237 labels_placeholder,
238 FLAGS.batch_size)
239
240 # Run one step of the model.

Callers 1

test_trainMethod · 0.70

Calls 13

runMethod · 0.95
mkdtempMethod · 0.80
timeMethod · 0.80
placeholder_inputsFunction · 0.70
inferenceFunction · 0.70
fill_feed_dictFunction · 0.70
do_evalFunction · 0.70
maxFunction · 0.50
as_defaultMethod · 0.45
GraphMethod · 0.45
lossMethod · 0.45
groupMethod · 0.45

Tested by 1

test_trainMethod · 0.56