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

基于MNIST训练后,如何用Keras的CNN预测本地目录自定义图像?

How to Predict Custom Local Images with a Trained Keras CNN Model (MNIST)

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.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 08:44:35