如何打印Estimator.predict生成器对象的内容?
问题:TensorFlow Estimator.predict生成器无法迭代,报错TypeError: 'NoneType' object cannot be interpreted as an integer
问题背景
执行my_predictions = estimator.predict(input_fn=functools.partial(ds.eval_input_fn, params))后得到生成器对象<generator object Estimator.predict at 0x7fc2de02ef20>。模型训练后评估正常,输出为:Evaluations: {'loss': 0.031095365, 'global_step': 25666}。根据TensorFlow文档,该生成器应返回预测张量的评估值,但尝试查看内容时持续报错。
已尝试的方法
print(*my_predictions, sep='\n')list(my_predictions)next(my_predictions)for x in my_predictions: print(x)
报错信息
File "/usr/local/lib/python3.8/dist-packages/tensorflow_estimator/python/estimator/estimator.py", line 647, in predict for i in range(self._extract_batch_length(preds_evaluated)): TypeError: 'NoneType' object cannot be interpreted as an integer
生成器对象属性
__class__ <class 'generator'> __del__ <method-wrapper '__del__' of generator object at 0x7fc2de02ef20> __delattr__ <method-wrapper '__delattr__' of generator object at 0x7fc2de02ef20> __dir__ <built-in method __dir__ of generator object at 0x7fc2de02ef20> __doc__ None __eq__ <method-wrapper '__eq__' of generator object at 0x7fc2de02ef20> __format__ <built-in method __format__ of generator object at 0x7fc2de02ef20> __ge__ <method-wrapper '__ge__' of generator object at 0x7fc2de02ef20> __getattribute__ <method-wrapper '__getattribute__' of generator object at 0x7fc2de02ef20> __gt__ <method-wrapper '__gt__' of generator object at 0x7fc2de02ef20> __hash__ <method-wrapper '__hash__' of generator object at 0x7fc2de02ef20> __init__ <method-wrapper '__init__' of generator object at 0x7fc2de02ef20> __init_subclass__ <built-in method __init_subclass__ of type object at 0x8fc1c0> __iter__ <method-wrapper '__iter__' of generator object at 0x7fc2de02ef20> __le__ <method-wrapper '__le__' of generator object at 0x7fc2de02ef20> __lt__ <method-wrapper '__lt__' of generator object at 0x7fc2de02ef20> __name__ predict __ne__ <method-wrapper '__ne__' of generator object at 0x7fc2de02ef20> __new__ <built-in method __new__ of type object at 0x9075a0> __next__ <method-wrapper '__next__' of generator object at 0x7fc2de02ef20> __qualname__ Estimator.predict __reduce__ <built-in method __reduce__ of generator object at 0x7fc2de02ef20> __reduce_ex__ <built-in method __reduce_ex__ of generator object at 0x7fc2de02ef20> __repr__ <method-wrapper '__repr__' of generator object at 0x7fc2de02ef20> __setattr__ <method-wrapper '__setattr__' of generator object at 0x7fc2de02ef20> __sizeof__ <built-in method __sizeof__ of generator object at 0x7fc2de02ef20> __str__ <method-wrapper '__str__' of generator object at 0x7fc2de02ef20> __subclasshook__ <built-in method __subclasshook__ of type object at 0x8fc1c0>
排查方向与解决方案建议
核心问题分析
报错源于_extract_batch_length(preds_evaluated)返回None,导致range()无法处理。常见原因:
- input_fn输出不符合predict要求:评估用input_fn可能返回
(features, labels)元组,但predict的input_fn应仅返回features(无需labels)。 - 模型predict分支输出异常:
model_fn中ModeKeys.PREDICT分支返回的EstimatorSpec未正确设置predictions,或predictions结构无效(如为None、无可用batch长度的张量)。
具体排查步骤
- 调整input_fn:若原
eval_input_fn返回(features, labels),需包装为仅返回features的版本:
再用def predict_input_fn(params): features, _ = ds.eval_input_fn(params) return featuresfunctools.partial(predict_input_fn, params)作为input_fn传入predict。 - 检查model_fn的PREDICT分支:确保该分支返回的
EstimatorSpec包含有效predictions:if mode == tf.estimator.ModeKeys.PREDICT: predictions = {"predicted_value": logits} return tf.estimator.EstimatorSpec(mode, predictions=predictions) - 验证input_fn输出:单独调用input_fn,打印返回值的结构与内容,确认是模型可处理的features格式,无None值。
- 简化测试:用手动构造的单条测试数据传入input_fn,验证能否正常生成预测,逐步排查数据量或格式问题。
内容的提问来源于stack exchange,提问作者Alessandro
相关产品推荐
相关产品推荐

