如何将冻结的MNIST模型加载到TensorFlow Estimator用于预测?
Hey there! Great question—yes, you absolutely can reload your frozen frozen_mnist.pb model into a TensorFlow Estimator object to use the estimator.predict() method. It just requires a bit of custom setup since Estimators are built around model functions rather than pre-frozen graphs out of the box. Let me walk you through how to do this step by step:
Step 1: Define a Custom Model Function for the Frozen Graph
You'll need to create a model function that loads your frozen graph instead of constructing the model from scratch. This function handles mapping your input features to the graph's input tensor and extracting prediction outputs.
import tensorflow as tf def frozen_mnist_model_fn(features, labels, mode): # Load the frozen graph from the .pb file with tf.io.gfile.GFile('frozen_mnist.pb', 'rb') as f: graph_def = tf.GraphDef() graph_def.ParseFromString(f.read()) # Import the graph into the current TensorFlow graph # Replace 'input:0' and 'output/predictions:0' with your actual tensor names # (You can find these using TensorBoard or by inspecting the frozen graph) input_tensor, output_logits = tf.import_graph_def( graph_def, input_map={'input:0': features['image']}, # Map your input feature to the graph's input return_elements=['input:0', 'output/predictions:0'] ) # Handle prediction mode (the only mode we need for this use case) if mode == tf.estimator.ModeKeys.PREDICT: # Format predictions into the structure Estimator expects predictions = { 'class_ids': tf.argmax(output_logits, axis=1, output_type=tf.int32), 'probabilities': output_logits } return tf.estimator.EstimatorSpec(mode=mode, predictions=predictions) # We don't need TRAIN or EVAL modes since the model is frozen (can't be retrained) raise NotImplementedError("This model function only supports PREDICT mode")
Step 2: Create the Estimator Instance
Use your custom model function to initialize an Estimator:
mnist_estimator = tf.estimator.Estimator(model_fn=frozen_mnist_model_fn)
Step 3: Prepare a Prediction Input Function
Define an input function that feeds your test data into the Estimator. This should match the input format your frozen model expects:
def predict_input_fn(): # Replace this with your actual preprocessed test data # For MNIST, this would be normalized 28x28 images test_images = ... # Shape: [num_samples, 28, 28, 1] or [num_samples, 784] # Create a dataset from your test images dataset = tf.data.Dataset.from_tensor_slices({'image': test_images}) dataset = dataset.batch(32) # Use a batch size that works for your data return dataset
Step 4: Run Predictions
Now you can use the predict() method just like you would with a regular Estimator:
predictions = mnist_estimator.predict(input_fn=predict_input_fn) # Iterate through the predictions for idx, pred in enumerate(predictions): print(f"Sample {idx}: Predicted class = {pred['class_ids']}, Probabilities = {pred['probabilities']}")
Key Notes to Keep in Mind
- Tensor Names: Make sure you use the exact names of the input and output tensors from your frozen graph. You can inspect these by loading the graph in TensorBoard or using
tf.get_default_graph().get_all_tensor_names()after importing the graph. - Input Format: Your input features must match the shape and data type expected by the frozen model (e.g., normalized pixel values between 0 and 1 for MNIST).
- Training Limitation: Since this is a frozen model, you can't use the Estimator for training or evaluation—only prediction. If you need to retrain, you'd have to work with the original model checkpoint instead of the frozen graph.
内容的提问来源于stack exchange,提问作者Andrea Rossi

