You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何从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.

Step 1: Update the Training Code to Export the Estimator

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.

Step 2: Load the Exported Model 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 predictions dictionary uses keys defined in your model_fn—in the original tutorial, these are classes (the predicted digit) and probabilities (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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.29 08:23:09