tf.estimator.train_and_evaluate二次调用仅训练1步问题解决
问题:使用tf.estimator.train_and_evaluate训练第二个模型时提前终止
我用tf.estimator.train_and_evaluate依次训练两个TensorFlow模型,代码如下:
# 第一个模型 train_spec = tf.estimator.TrainSpec(input_fn=my_input_fn("train"), max_steps=max_steps) eval_spec = tf.estimator.EvalSpec(input_fn=my_input_fn("valid"), steps=None) tf.estimator.train_and_evaluate(model1, train_spec, eval_spec) # 第二个模型 hook = [tf.estimator.StopAtStepHook(num_steps=max_steps)] train_spec = tf.estimator.TrainSpec(input_fn=my_input_fn("train"), max_steps=None, hooks=hook) eval_spec = tf.estimator.EvalSpec(input_fn=my_input_fn("valid"), steps=None) tf.estimator.train_and_evaluate(model2, train_spec, eval_spec)
第一个模型训练正常,但第二个模型仅完成1步训练就终止,日志显示全局步骤1时保存指标,最终步骤损失为0.32117385。预期两个模型训练过程一致,实际第二个模型提前结束。
原因与修复方案
问题核心出在第二个模型的input_fn配置上:如果my_input_fn内部使用tf.estimator.inputs.numpy_input_fn,num_epochs默认值为1或显式设置为1会导致训练数据只遍历一轮,直接触发训练终止。
修复步骤:
- 将训练集对应的
numpy_input_fn中num_epochs设置为None,让数据可以循环迭代 - 保持
TrainSpec的max_steps=None,通过StopAtStepHook(num_steps=max_steps)控制训练总步数 - 确保未启用任何早停相关Hook
调整后的my_input_fn示例(针对训练场景):
def my_input_fn(mode): if mode == "train": return tf.estimator.inputs.numpy_input_fn( x={"x": train_data}, y=train_labels, batch_size=batch_size, num_epochs=None, # 关键修改:设为None实现数据循环 shuffle=True ) # 验证集无需循环,保持num_epochs默认值1即可 elif mode == "valid": return tf.estimator.inputs.numpy_input_fn( x={"x": valid_data}, y=valid_labels, batch_size=batch_size, shuffle=False )
内容的提问来源于stack exchange,提问作者Ergun Bicici
相关产品推荐
相关产品推荐

