使用tf.contrib.learn.LinearRegressor报错,需切换至tf.train.get_global_step
解决TensorFlow线性回归相关的两个问题
一、处理tf.contrib.learn.LinearRegressor的更新报错
你遇到的Please switch to tf.train.get_global_step提示,本质是因为tf.contrib.learn属于TensorFlow的旧版API,已经被标记为废弃(deprecated),官方早就停止了维护,转而推荐使用更稳定的tf.estimator系列API(也就是你后来尝试的tf.estimator.LinearRegressor)。
关于你提到的SKCompat类,它是旧版用来兼容Scikit-learn接口的过渡类,现在完全没有必要再用了——直接迁移到tf.estimator才是长久之计,这也是官方的明确建议。所以不建议再纠结旧版的报错,直接转向新版API解决问题更高效。
二、解决tf.estimator.LinearRegressor无输出的问题
你执行新版代码后看不到输出,主要有两个原因:
1. evaluate()的结果需要主动打印
estimator.evaluate()会返回一个包含评估指标(比如损失值等)的字典,但它不会自动打印结果,你需要手动输出:
eval_results = estimator.evaluate(input_fn=eval_input_fn) print("评估结果:", eval_results)
2. predict()返回的是迭代器,需要遍历获取结果
estimator.predict()返回的是一个生成器对象,并非直接的预测结果列表,你需要通过遍历或者转换为列表来查看输出:
predictions = estimator.predict(input_fn=eval_input_fn) # 遍历打印每个预测结果 for pred in predictions: print("预测值:", pred['predictions'][0]) # 或者转换为列表一次性查看 pred_list = list(predictions) print("所有预测结果:", pred_list)
另外补充一点:estimator.train()如果不指定steps参数,会一直训练下去(直到输入数据耗尽,如果是无限生成的input_fn就会无限训练),建议加上steps参数控制训练步数,比如:
estimator.train(input_fn=training_input_fn, steps=10000)
这样修改后,你就能看到训练、评估和预测的具体结果了。
内容的提问来源于stack exchange,提问作者Trufa
相关产品推荐
相关产品推荐

