如何从TensorFlow MNIST教程导出Estimator并导入用于预测?
Got it, let's walk through exactly how to export the trained Estimator from that MNIST CNN tutorial and then load it back up for prediction. I’ll break this into straightforward steps that fit right into the existing tutorial code.
First, we need to add code to export the model after it finishes training. The key here is defining a serving input function that tells TensorFlow what input format the exported model should accept (this needs to match what your model expects during training).
1.1 Add the Serving Input Function
Insert this function somewhere in your cnn_mnist.py file (before the training loop works well):
def serving_input_receiver_fn(): # Define the input placeholder that matches our model's training input # Our model expects flattened 28x28 images (784 pixels) as input 'x' input_placeholder = tf.compat.v1.placeholder(tf.float32, shape=[None, 784]) # Wrap the placeholder in a dictionary matching the input key used in training inputs = {'x': input_placeholder} # Return a ServingInputReceiver that links the input to the model return tf.estimator.export.ServingInputReceiver(inputs, inputs)
1.2 Export the Model After Training
Find the part of the code where training finishes (right after classifier.train(...)), and add this code to export the trained model:
# Define a base directory to save the exported model export_dir_base = "./mnist_cnn_exported_model" # Export the trained Estimator exported_model_path = classifier.export_saved_model( export_dir_base, serving_input_receiver_fn=serving_input_receiver_fn ) print(f"Successfully exported model to: {exported_model_path}")
When you run the training script now, it will create a timestamped subfolder inside mnist_cnn_exported_model (this is TensorFlow’s way of versioning exported models). Note down this folder path—you’ll need it for prediction.
Now let’s write a separate script (or add to your existing one) to load the model and run predictions on new data.
2.1 Load the Exported Model
Use TensorFlow’s predictor utility to load the saved model. Replace the <timestamped_folder> placeholder with the actual folder name from your export step:
import tensorflow as tf from tensorflow.examples.tutorials.mnist import input_data # Path to your exported model's timestamped folder exported_model_dir = "./mnist_cnn_exported_model/<timestamped_folder>" # Load the model into a predictor object predictor = tf.contrib.predictor.from_saved_model(exported_model_dir)
2.2 Prepare Input Data and Run Predictions
We’ll use MNIST test data as an example, but you can use any flattened 28x28 grayscale images (scaled to 0-1) here:
# Load MNIST test data (or use your own images) mnist = input_data.read_data_sets("/tmp/mnist_data", one_hot=False) # Grab 5 test images to predict on test_images = mnist.test.images[:5] # Format input to match what the model expects (dictionary with key 'x') input_data = {"x": test_images} # Run prediction predictions = predictor(input_data) # Print the results! print("Predicted class labels:", predictions["classes"]) print("Prediction probabilities (confidence):", predictions["probabilities"])
Key Notes to Avoid Issues
- Input Format Match: Make sure your prediction input is flattened (784 elements per image) and scaled the same way as training data (the MNIST tutorial scales pixels to 0-1 automatically).
- Prediction Keys: The
predictionsdictionary uses keys defined in yourmodel_fn—in the original tutorial, these areclasses(the predicted digit) andprobabilities(confidence scores for each digit 0-9). - TensorFlow Version Compatibility: Since the tutorial uses TF 1.8, stick with TF 1.x syntax (avoid TF2.x APIs here unless you’re migrating the code).
内容的提问来源于stack exchange,提问作者Razvan

