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

髋关节植入物影像分类模型预测输入维度不兼容问题排查

髋关节植入物X光分类模型单图预测报错:输入维度不匹配

背景介绍

  • 基于X光影像构建髋关节植入物松动/正常分类模型
  • 数据存储在GCS桶的CSV文件中,包含图片路径和类别两列
  • 训练时将图片统一调整为224×224尺寸

约束条件

  • 可用影像数据有限,优先保证模型可运行,暂不追求预测精度

问题现象

调用model.predict做单图预测时,抛出以下错误:

ValueError: Input 0 of layer "sequential_7" is incompatible with the layer: expected shape=(None, 224, 224, 1), found shape=(None, 244, 1)

相关代码片段

1. 数据管道

# 数据管道
CLASS_NAMES = ['loose', 'control']

def decode_csv(csv_row): # csv_row包含文件路径和图片类别
    record_defaults = ["path", "image class"] # 数据集默认值
    filename, label_string = tf.io.decode_csv(csv_row, record_defaults) # 读取CSV每行数据
    
    image_bytes = tf.io.read_file(filename=filename) # 输出:base64图片字符串
    image_bytes = tf.image.decode_jpeg(image_bytes) # 输出:整数数组
    image_bytes = tf.image.convert_image_dtype(image_bytes, tf.float32) # 输出:0-1范围的浮点数
    image_bytes = tf.image.resize(image_bytes, [224, 224]) # 输出:统一尺寸后的图片
    
    label = tf.math.equal(CLASS_NAMES, label_string) # 将标签格式化为布尔数组,对应类别为True
    
    return image_bytes, label # 返回处理后的图片和标签

def load_dataset(csv_file, batch_size, training=True):
    ds = tf.data.TextLineDataset(filenames=csv_file).skip(1) # 跳过表头行
    ds = ds.map(decode_csv).cache()
    ds = ds.batch(batch_size=batch_size)
    
    if training:
        ds = ds.shuffle(10).repeat()
    return ds
train_ds = load_dataset("gs://qwiklabs-asl-04-06351f77b64f-hip-implant/hip-implant-data.csv", batch_size = 10)
validation_data = load_dataset("gs://qwiklabs-asl-04-06351f77b64f-hip-implant/hip-implant-data.csv", batch_size = 10, training=False)

2. 模型创建

# 创建模型
IMG_HEIGHT = 224
IMG_WIDTH = 224
IMG_CHANNELS = 64

model = Sequential([
    Conv2D(name="first-Conv2D-layer",filters=64, kernel_size=3, input_shape=(IMG_WIDTH, IMG_HEIGHT, 1), padding='same', activation='relu'),
    MaxPooling2D(name="first-pooling-layer",strides=2, padding='same'),
    Conv2D(name="second-Conv2D-layer", filters=32, kernel_size=3, activation='relu'),
    MaxPooling2D(name="second-pooling-layer", strides=2, padding='same'),
    Flatten(),
    Dense(units=400, activation='relu'),
    Dense(units=100, activation='relu'),
    Dropout(0.25),
    Dense(2),
    Softmax()
])

model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])

3. 单图预测代码

# 单图预测
image_path = tf.io.read_file("gs://qwiklabs-asl-04-06351f77b64f-hip-implant/Control/control (25).png")

new_image = decode_img(image_path, [244, 244])

print(new_image.shape)
plt.imshow(new_image.numpy())

prediction = model.predict(new_image)
print(prediction)

补充:decode_img函数

img = tf.io.read_file("gs://qwiklabs-asl-04-06351f77b64f-hip-implant/Control/control (25).png")

def decode_img(img, reshape_dims):
    img = tf.image.decode_jpeg(img) # 将base64图片字符串解码为整数数组
    img = tf.image.convert_image_dtype(img, tf.float32) # 将整数数组转换为0-1范围的浮点数
    img = tf.image.resize(img, reshape_dims) # 统一图片尺寸
    return img


img = decode_img(img, [224, 224])

plt.imshow(img.numpy())

已尝试的解决方法

  1. 在卷积层设置padding='same',解决了初始的层内维度不匹配问题
  2. 显式指定模型输入形状为(244,244,1)(添加Input层)
  3. 修改滤波器大小、神经元数量、池化大小,解决层无法进一步降维的错误

解决方案

报错核心有两个原因:

  1. 预测时图片尺寸写错:单图预测代码里把[224,224]写成了[244,244],和模型输入要求的224×224不匹配
  2. 单图缺少batch维度和通道维度:模型输入要求是(None, 224,224,1),其中None代表batch维度,而单图处理后是(224,224)或(224,224,3)(X光图应为单通道灰度图),需要补充维度

修正后的单图预测代码

# 修正后的单图预测
image_path = tf.io.read_file("gs://qwiklabs-asl-04-06351f77b64f-hip-implant/Control/control (25).png")

# 1. 修正尺寸为224×224
new_image = decode_img(image_path, [224, 224])

# 2. 转换为单通道并补充batch维度
new_image = tf.image.rgb_to_grayscale(new_image) # 转灰度图(适配模型单通道输入)
new_image = tf.expand_dims(new_image, axis=0) # 增加batch维度,形状变为(1,224,224,1)

print(new_image.shape)
plt.imshow(new_image.numpy()[0, ..., 0], cmap='gray') # 正确显示灰度图

prediction = model.predict(new_image)
print(prediction)
# 输出预测类别
pred_class = CLASS_NAMES[tf.argmax(prediction, axis=1)[0]]
print(f"预测类别:{pred_class}")

额外优化建议

  • 训练数据管道同步修改:在decode_csv函数中添加image_bytes = tf.image.rgb_to_grayscale(image_bytes),保证训练和预测的通道数一致
  • 统一使用IMG_HEIGHT/IMG_WIDTH变量,避免硬编码尺寸导致的错误

内容的提问来源于stack exchange,提问作者fastscholar

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 00:23:11