TensorFlow高阶Estimator更新问题:DNN Regressor.predict()无法转列表
我明白你碰到的这个麻烦——用TensorFlow 1.7.0里的DNNRegressor做预测时,classifier.predict()返回的生成器突然没法用list()转成结果了,之前明明好使,降级版本也没解决,确实头疼。下面给你几个针对性的解决方案:
解决DNNRegressor.predict()生成器无法用list()转换的问题
可能的核心原因
在TensorFlow 1.x的部分版本中,predict()返回的生成器依赖**活跃的TensorFlow会话(Session)**才能迭代取值。如果你的代码里会话管理逻辑有变化(比如更新后隐式会话的处理方式变了),直接调用list()就会失效——因为生成器内部需要会话来拉取计算结果,没有会话的话根本没法生成数据。
具体解决方案
方案1:显式在会话上下文中处理生成器
这是最稳妥的方式,确保生成器的迭代过程处在活跃的会话里,示例代码如下:
import tensorflow as tf # 假设你已经定义好特征列和回归器 feature_columns = [tf.feature_column.numeric_column("x", shape=[你的特征维度])] classifier = tf.estimator.DNNRegressor( feature_columns=feature_columns, hidden_units=[10, 20, 10], model_dir="./your_model_dir" ) # 定义预测用的输入函数 def predict_input_fn(test_data): # 替换成你的测试数据,确保特征key和训练时一致 return tf.data.Dataset.from_tensor_slices({"x": test_data}).batch(32) # 显式创建会话并处理预测 with tf.Session() as sess: # 初始化变量(如果是加载已训练模型,可能需要恢复模型参数) sess.run(tf.global_variables_initializer()) # 获取预测生成器 predictions_gen = classifier.predict(input_fn=lambda: predict_input_fn(你的测试数据)) # 在会话内转换为列表 pred_results = list(predictions_gen) # 现在可以正常使用pred_results了 print(pred_results)
方案2:手动迭代生成器提取结果
如果list()还是不生效,可以手动循环生成器,逐个取出预测值——DNNRegressor的预测结果是字典,对应的预测值存在"predictions"键下:
predictions_gen = classifier.predict(input_fn=lambda: predict_input_fn(你的测试数据)) pred_list = [] for pred_dict in predictions_gen: # 取出单个预测值(如果是多维度预测,根据需求调整索引) pred_value = pred_dict["predictions"][0] pred_list.append(pred_value)
方案3:检查输入函数的正确性
有时候输入函数的格式错误也会导致生成器无法正常迭代。确保你的input_fn返回的是符合Estimator要求的tf.data.Dataset,特征字典的key和训练时完全一致,数据类型、维度也和训练时匹配。比如:
# 正确的预测输入函数示例(适配单特征或多特征) def predict_input_fn(data): feature_dict = {"x": data} # 如果有多个特征,添加对应的key和数据 return tf.data.Dataset.from_tensor_slices(feature_dict).batch(32)
额外提示
TensorFlow 1.x的Estimator API在小版本间确实有一些细节变动,但降级版本没用的话,大概率不是版本本身的问题,而是代码中会话管理、模型加载或者输入数据的细节和之前不一样了。可以检查下模型目录是否正确加载了训练好的参数,或者输入数据的预处理逻辑有没有变化。
内容的提问来源于stack exchange,提问作者5Volts
相关产品推荐
相关产品推荐

