TensorFlow中model.predict调用重复数据集死循环及输入形状不兼容问题求助
TensorFlow中model.predict调用重复数据集死循环及输入形状不兼容问题求助
咱们先把你遇到的几个问题的根源拆解清楚:
- 为什么带
.repeat()的测试数据集会让predict无限循环?因为.repeat()会让数据集无限重复迭代,训练时没问题是因为你指定了epochs和steps_per_epoch,模型知道何时停止;但predict没有这些限制,会一直从数据集中取数据,自然停不下来。 - 直接用未做batch处理的
test_dataset_sim报形状错误?你的模型输入期望的是带batch维度的张量(形状为(None, 24),None代表任意batch大小),但未batch的数据集每次返回单个样本,形状是(24,),缺少了batch维度,所以模型无法识别。 - 只做batch不加
repeat()出现OUT_OF_RANGE警告?这是数据集遍历完毕后的正常提示,但可能导致预测结果不完整,咱们可以用正确的方式规避。
下面给你几个实用的解决方案,任选其一即可:
方案一:给测试数据集仅做batch处理,明确指定predict的步数
这是最推荐的方式,既解决形状问题,又避免无限循环:
- 处理测试数据集时,只做batch操作,不要加
repeat() - 计算测试集需要的预测步数(总样本数除以batch大小,向上取整)
- 调用
predict时指定steps参数
代码示例:
# 处理测试数据集:仅batch,不repeat test_dataset = test_dataset.batch(test_batch_sz) # 计算测试集总样本数 test_samples_num = len(list(test_dataset.unbatch())) # 计算需要的预测步数(向上取整) import math steps = math.ceil(test_samples_num / test_batch_sz) # 执行预测,指定steps后模型会在完成对应步数后停止 y_pred = model.predict(test_dataset, steps=steps)
这样操作后,既不会出现无限循环,输入形状也符合模型要求,同时OUT_OF_RANGE警告也会消失。
方案二:给单个样本手动添加batch维度(适合小样本场景)
如果不想对数据集做batch处理,可以给每个样本手动补上batch维度,再逐个预测:
predictions = [] # 遍历未batch的测试数据集 for sample in test_dataset_sim: # 给样本添加batch维度,从(24,)变为(1, 24) sample_with_batch = tf.expand_dims(sample, axis=0) # 执行单样本预测,关闭日志输出避免刷屏 pred = model.predict(sample_with_batch, verbose=0) predictions.append(pred) # 把所有预测结果合并成一个数组 import numpy as np y_pred = np.concatenate(predictions, axis=0)
这个方法适合测试样本量较小的场景,大样本下效率不如batch处理。
方案三:用take限制repeat的迭代次数(不推荐,仅作补充)
如果一定要用repeat(),可以配合take方法限制取数据的batch数量,本质和方案一逻辑类似:
test_samples_num = len(list(test_dataset.unbatch())) test_batch_num = math.ceil(test_samples_num / test_batch_sz) # 仅重复取指定数量的batch后停止 test_dataset = test_dataset.batch(test_batch_sz).repeat().take(test_batch_num) # 执行预测,此时数据集不会无限迭代 y_pred = model.predict(test_dataset)
最后再划个重点:训练数据集用repeat()是合理的,因为需要多轮迭代;但测试/预测数据集绝对不要加repeat(),通过指定steps参数控制预测停止时机,同时必须保证输入带batch维度(即数据集经过.batch()处理)。
备注:内容来源于stack exchange,提问作者Jonathan Roy
相关产品推荐
相关产品推荐

