如何将TFRecord文件输入模型训练并修复输入形状不匹配报错
问题原因
报错的核心是TFRecord数据集返回的单样本缺少batch维度,而Keras模型输入默认要求第一个维度为可变batch维度(即报错信息中的None维度)。
修复方案
方案1:给数据集添加batch操作(最推荐)
在获取数据集后调用.batch()方法指定批次大小即可,测试阶段可设为1,正式训练可根据显存容量调整为2、4等数值。
修改训练代码如下:
dataset = get_dataset('test.tfrecords') # 新增batch操作,这里测试用batch size=1 dataset = dataset.batch(1) my_model.fit(dataset, epochs = 1)
修改后可执行以下代码确认形状匹配:
print(dataset.element_spec) # 预期输出:(TensorSpec(shape=(None, 240, 240, 155, 4), dtype=tf.float32, name=None), TensorSpec(shape=(None, 240, 240, 155), dtype=tf.int32, name=None))
注意:3D卷积的显存占用较高,设置大batch size前请先确认显存容量足够
方案2:手动给单样本添加batch维度
如果仅需要单样本训练/推理,不需要批量处理,可以在数据预处理阶段手动给每个样本添加第0维的batch维度,修改get_image_and_label函数即可:
def get_image_and_label(features): image, label = features['image'], features['label'] # 手动添加batch维度 image = tf.expand_dims(image, axis=0) label = tf.expand_dims(label, axis=0) return image, label
内容的提问来源于stack exchange,提问作者user15515518
相关产品推荐
相关产品推荐

