TensorFlow 1.5中Estimator.predict_scores结合Dataset API预测数量异常问题
嘿,这个问题我之前在TensorFlow 1.x里踩过坑!你遇到的predict只输出batch_size(16条)结果的情况,核心原因是自定义TFRecord输入函数在预测模式下没有正确配置数据集的遍历逻辑,导致estimator只拿到了一个batch的数据就停止了。
为什么会这样?
tf.estimator的predict方法会一直从输入函数提供的数据集里拉取数据,直到数据集抛出OutOfRangeError(也就是所有数据都遍历完)。而tf.estimator.inputs.numpy_input_fn默认帮你处理好了预测模式下的完整遍历;但自定义TFRecordDataset输入函数时,很容易忽略这些细节:
- 你可能在预测模式下的数据集配置有误,比如没有让它完整遍历一次所有样本;
- 或者解析TFRecord的函数不小心过滤掉了大部分数据(比如把训练时的过滤逻辑带到了预测里)。
直接给你修复方案
我给你一个标准的TFRecord输入函数模板,严格区分训练和预测模式,保证预测时能拿到所有样本:
def parse_tfrecord_example(serialized_example): # 替换成你的TFRecord解析逻辑,这里是示例 feature_spec = { 'features': tf.FixedLenFeature([你的特征维度], tf.float32), # 其他需要的特征... } parsed_features = tf.parse_single_example(serialized_example, feature_spec) return parsed_features['features'] # 返回用于预测的特征 def custom_input_fn(mode, batch_size=16): # 1. 加载TFRecord文件 dataset = tf.data.TFRecordDataset("你的TFRecord文件路径") # 2. 并行解析样本 dataset = dataset.map( parse_tfrecord_example, num_parallel_calls=tf.data.experimental.AUTOTUNE ) # 3. 按模式配置数据集 if mode == tf.estimator.ModeKeys.TRAIN: # 训练模式:打乱+无限重复+分批 dataset = dataset.shuffle(buffer_size=256) # buffer_size建议设为样本总量 dataset = dataset.repeat() # 无限重复,训练时由steps参数控制停止 dataset = dataset.batch(batch_size) elif mode == tf.estimator.ModeKeys.PREDICT: # 预测模式:不打乱+只遍历一次+分批 dataset = dataset.batch(batch_size) # 注意:TF1.5中Dataset默认是遍历一次就结束,所以不用额外加repeat(1),除非你需要重复预测 return dataset
调用predict的时候要指定模式:
predict_results = estimator.predict( input_fn=lambda: custom_input_fn(tf.estimator.ModeKeys.PREDICT, batch_size=16) ) # 把生成器转成列表,就能看到所有256条结果了 predict_scores = list(predict_results) print(len(predict_scores)) # 这里应该输出256
额外要检查的点
- 先确认你的TFRecord文件里确实有256条有效样本,用这个小脚本验证:
count = 0 for _ in tf.python_io.tf_record_iterator("你的TFRecord文件路径"): count += 1 print(f"TFRecord里的样本数:{count}")
- 检查解析函数
parse_tfrecord_example,别把训练时的过滤逻辑(比如过滤某些标签)带到预测里; - 确保predict时传入的batch_size是16,别不小心设成了256(那只会输出1条,但你的情况应该不是这个)。
内容的提问来源于stack exchange,提问作者Natello
相关产品推荐
相关产品推荐

