Train MNIST for a number of steps.
()
| 181 | |
| 182 | |
| 183 | def 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. |