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

如何将冻结的MNIST模型加载到TensorFlow Estimator用于预测?

Reloading a Frozen MNIST Model into TensorFlow Estimator for Predictions

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 09:23:52