ResNet图像分类预测报错:无法将数组重塑为(1,224,224,3)
Hey there! Let's break down why you're hitting this ValueError and get your prediction code working smoothly.
What's Causing the Error?
First, let's unpack the error message:
ValueError: cannot reshape array of size 50176 into shape (1,224,224,3)
Do the quick math:
- Your target shape requires
224 × 224 × 3 = 150528elements (for a 3-channel RGB image with a batch dimension) - The array you're trying to reshape only has
50176 = 224 × 224elements—this is exactly the size of a single-channel grayscale image.
This means your img_data is a grayscale (1-channel) array, but your ResNet model was trained to expect 3-channel RGB input (since you set the input size to 224×224×3 during training).
How to Fix It
Follow these steps to align your image data with the model's requirements:
Convert single-channel grayscale to 3-channel RGB
Before reshaping, replicate the single channel three times to mimic an RGB image. Use this numpy operation:# Convert (224,224) grayscale array to (224,224,3) RGB img_data = np.stack((img_data,) * 3, axis=-1)If you're loading images with PIL, you can skip the numpy step by converting to RGB directly when opening:
from PIL import Image img = Image.open("your_test_image.jpg").convert("RGB").resize((224,224)) img_data = np.array(img)Reshape to the model's expected batch input shape
Now that you have a 3-channel image, add the batch dimension (the1at the start, since we're predicting one image at a time):data = img_data.reshape(1, 224, 224, 3)Fix the prediction line
In your original code, you passedorig(the original grayscale array) topredict_classes—this would still cause a mismatch. Instead, use the reshapeddata:# For older Keras versions with predict_classes model_out = model.predict_classes(data)[0] # For newer Keras versions (predict_classes is deprecated) predictions = model.predict(data) model_out = np.argmax(predictions, axis=1)[0]
Full Corrected Code Snippet
Here's your loop with all fixes applied:
fig = plt.figure() for num, data in enumerate(test_data): img_num = data[1] img_data = data[0] # Convert grayscale to 3-channel RGB if needed if len(img_data.shape) == 2: img_data = np.stack((img_data,) * 3, axis=-1) y = fig.add_subplot(9, 3, num+1) orig = img_data # Reshape for model input data = img_data.reshape(1, 224, 224, 3) # Run prediction predictions = model.predict(data) model_out = np.argmax(predictions, axis=1)[0] str_label = 'Dog' if model_out == 1 else 'Cat' y.imshow(orig) plt.title(str_label) y.axes.get_xaxis().set_visible(False) y.axes.get_yaxis().set_visible(False) plt.show() plt.savefig('test_labeled.jpg')
内容的提问来源于stack exchange,提问作者Madhi

