You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何打印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()无法处理。常见原因:

  1. input_fn输出不符合predict要求:评估用input_fn可能返回(features, labels)元组,但predict的input_fn应仅返回features(无需labels)。
  2. 模型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 features
    
    再用functools.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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.26 03:45:49