基于MNIST训练后,如何用Keras的CNN预测本地目录自定义图像?
Great question! Since your CNN is trained on MNIST's 28x28 grayscale digits with specific formatting, you just need to preprocess your local images to match that format before feeding them into model.predict(). Let's break this down step by step, using your existing training code as a base.
First: Save Your Trained Model (Optional but Recommended)
Your current code trains the model but doesn't save it—so you'd have to retrain every time you want to make a prediction. Add this line right after model.fit() to save the model for later use:
# Save the trained model to a file model.save('mnist_cnn_model.h5')
Later, when you want to load the model without retraining, use:
from keras.models import load_model model = load_model('mnist_cnn_model.h5')
Step-by-Step: Preprocess & Predict Local Images
You'll need libraries like PIL (for image handling) and numpy (for array manipulation). Here's the full workflow:
1. Import Required Libraries
Add these imports to your code:
from PIL import Image import numpy as np
2. Load & Preprocess Your Local Image
Follow these steps to get your image ready for the model:
# 1. Load the local image (replace 'your_digit.png' with your image path) img = Image.open('your_digit.png') # 2. Convert to grayscale (matches MNIST's single-channel input) img_gray = img.convert('L') # 3. Resize to 28x28 (exact input shape your model expects) img_resized = img_gray.resize((28, 28)) # 4. Convert to numpy array img_array = np.array(img_resized) # 5. Reverse colors (critical if your image has a white background! MNIST uses black background with white digits) # If your image is white text on black background, skip this step img_array = 255 - img_array # 6. Reshape to match the model's input shape: (batch_size, height, width, channels) # We use batch_size=1 since we're predicting one image at a time img_input = img_array.reshape(1, 28, 28, 1)
3. Run Prediction & Get Results
Now feed the preprocessed image into your model:
# Get prediction probabilities for each digit (0-9) predictions = model.predict(img_input) # Find the digit with the highest probability predicted_digit = np.argmax(predictions) print(f"Predicted Digit: {predicted_digit}") print(f"Prediction Confidence: {np.max(predictions):.2f}")
Key Notes for Better Accuracy
- Image Alignment: Try to center the digit in your image, just like MNIST's samples—extra background noise can throw off the model.
- Color Check: Double-check if you need the color reversal step. You can visualize the preprocessed image with
plt.imshow(img_array, cmap='gray')to make sure it looks like a MNIST digit. - Pixel Range: Your training data uses raw 0-255 pixel values (no normalization), so keep your preprocessed image in the same range.
内容的提问来源于stack exchange,提问作者armin ariana

