髋关节植入物影像分类模型预测输入维度不兼容问题排查
髋关节植入物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())
已尝试的解决方法
- 在卷积层设置
padding='same',解决了初始的层内维度不匹配问题 - 显式指定模型输入形状为(244,244,1)(添加Input层)
- 修改滤波器大小、神经元数量、池化大小,解决层无法进一步降维的错误
解决方案
报错核心有两个原因:
- 预测时图片尺寸写错:单图预测代码里把
[224,224]写成了[244,244],和模型输入要求的224×224不匹配 - 单图缺少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
相关产品推荐
相关产品推荐

