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

ResNet图像分类预测报错:无法将数组重塑为(1,224,224,3)

Fixing the Reshape Error for ResNet Image Classification

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 = 150528 elements (for a 3-channel RGB image with a batch dimension)
  • The array you're trying to reshape only has 50176 = 224 × 224 elements—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:

  1. 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)
    
  2. Reshape to the model's expected batch input shape
    Now that you have a 3-channel image, add the batch dimension (the 1 at the start, since we're predicting one image at a time):

    data = img_data.reshape(1, 224, 224, 3)
    
  3. Fix the prediction line
    In your original code, you passed orig (the original grayscale array) to predict_classes—this would still cause a mismatch. Instead, use the reshaped data:

    # 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.27 07:33:23