如何在独立程序中恢复已训练TensorFlow模型以进行预测?
Hey there! I see you've got a trained model saved and just need to use it for predictions without retraining—let's get this sorted out. Here's a step-by-step guide tailored to your code structure:
First, a Quick Context Check
Make sure during training you used tf.train.Saver() to save your model. If you haven't added that yet (though you said training is done, this is just for reference), here's what the saving snippet would look like in your training code:
# Add this inside your training session saver = tf.train.Saver() # After training finishes, save the model saver.save(sess, "./your_trained_model.ckpt")
Step 1: Replicate the Exact Model Structure
Your prediction script needs identical variable definitions and the model function as your training code. Any tiny difference (like changing a filter size or variable name) will break model loading. Copy all those variable declarations (conv1_weights, conv1_biases, etc.) and the full model function into your prediction program.
Step 2: Set Up Prediction-Specific Inputs & Outputs
Replace training-only placeholders and operations with ones optimized for prediction:
# Adjust EVAL_BATCH_SIZE to fit your use case (e.g., 1 for single image predictions) EVAL_BATCH_SIZE = 1 # Define input placeholder for prediction data predict_data = tf.placeholder( data_type(), shape=(EVAL_BATCH_SIZE, IMAGE_SIZE, IMAGE_SIZE, NUM_CHANNELS)) # Get prediction outputs (set train=False to disable dropout!) predict_logits = model(predict_data, train=False) predict_probs = tf.nn.softmax(predict_logits)
Step 3: Load the Saved Model and Run Predictions
Create a Saver object, start a session, and restore your trained weights:
saver = tf.train.Saver() with tf.Session() as sess: # Replace with your actual model save path saver.restore(sess, "./your_trained_model.ckpt") print("Model loaded successfully!") # Prepare your input data (critical: match training preprocessing exactly!) # Example: load and preprocess an image test_image = your_custom_preprocessing_function("path/to/your/image.jpg") # Reshape to match the placeholder shape test_data = test_image.reshape((EVAL_BATCH_SIZE, IMAGE_SIZE, IMAGE_SIZE, NUM_CHANNELS)) # Run the prediction probabilities = sess.run(predict_probs, feed_dict={predict_data: test_data}) predicted_label = np.argmax(probabilities, axis=1)[0] print(f"Predicted Label: {predicted_label}, Confidence: {probabilities[0][predicted_label]:.2f}")
Key Tips to Avoid Headaches
- Match Preprocessing: Your input data must go through the exact same normalization, resizing, or scaling as your training data. If you scaled pixels to [0,1] during training, do the same for prediction!
- Skip Training Ops: Don't include training-specific code like optimizers or loss functions in your prediction script—they're unnecessary and can cause conflicts.
- Variable Consistency: Double-check that all variable names, shapes, and data types are identical to the training code. Even a small change (like using
tf.float64instead oftf.float32) will trigger shape mismatch errors.
Troubleshooting Common Issues
NotFoundError: Verify your model path is correct. TensorFlow looks for files likeyour_trained_model.ckpt.index,your_trained_model.ckpt.data-00000-of-00001, and.meta—ensure all are present in the directory.Shape Mismatch: ConfirmIMAGE_SIZE,NUM_CHANNELS, andNUM_LABELSmatch your training code. Also check that your input data shape aligns with the placeholder.- Unreliable Predictions: If outputs don't make sense, confirm your preprocessing is identical to training. You can also print a small slice of a variable (e.g.,
sess.run(conv1_weights[:1,:1,:1,:1])) to ensure the model loaded correctly.
内容的提问来源于stack exchange,提问作者MarkAlanFrank

