Tensorflow维度不匹配报错:如何修改数据集代码适配4维输入模型
解决方法
报错的核心原因是你所用模型的首层为卷积层(block1_conv1为VGG类CV模型的首层结构),卷积层要求输入必须为4维张量,格式为(batch_size, 图片高度, 图片宽度, 通道数),你当前生成的玩具输入batch后维度为(2, 1),仅包含batch维度和单特征维度,缺少两个空间维度和通道维度。
快速修改代码
仅需要给生成的inputs补全缺失的维度即可,以下是最小修改版本,可直接替换原有数据集生成代码:
# 可根据你的模型要求调整高度、宽度、通道数的数值,示例为模拟32x32大小的3通道RGB图片 inputs = tf.range(10.)[:, None, None, None] * tf.ones((1, 32, 32, 3)) labels = inputs[:, 0, 0, 0] * 5. + tf.range(5.)[None, :] return tf.data.Dataset.from_tensor_slices( dict(x=inputs, y=labels)).repeat().batch(2)
说明
- 代码中的32(图片高度)、32(图片宽度)、3(通道数)可根据你自己模型要求的输入尺寸自定义调整,比如模型要求输入为224*224的灰度图,就将对应参数改为224、224、1即可
- 该方案为模拟输入的快速适配方法,仅用于验证模型可正常运行,若要训练有效模型需要替换为真实的图片类数据集
内容的提问来源于stack exchange,提问作者caasswa
相关产品推荐
相关产品推荐

