设置num_epochs=None后TensorFlow Estimator训练陷入死循环的问题
你遇到的这个死循环问题,根源其实在num_epochs=None的设置和循环调用est.train的搭配上。我之前用TensorFlow 1.x的estimator时也踩过类似的坑,给你拆解一下原因和解决办法:
问题根源
当你设置num_epochs=None时,numpy_input_fn会生成一个永远不会停止的数据集——它会无限重复你的训练数据,不会抛出终止信号。你原本想循环15次、每次跑20步,但实际情况是:第一次调用est.train(steps=20)时,由于数据集无限,部分旧版本的TensorFlow里,estimator可能无法正确识别steps参数的终止条件,或者因为迭代器的状态在循环中被保留,导致训练一旦开始就停不下来,看起来就像死循环了。
解决办法
有两种简单的方式,选哪种看你的需求:
一次性跑完总步数(推荐)
既然你预期总步数是20×15=300,直接在一次train调用里指定总步数就行,完全不用循环,这也更符合estimator的设计逻辑:train_input = tf.estimator.inputs.numpy_input_fn( x={'x': sst_train}, y=precip_train, shuffle=True, batch_size=100, num_epochs=None ) est.train(input_fn=train_input, steps=300)这样estimator会自动处理无限数据集,跑完300步后就会停止。
每次循环重新创建输入函数
如果一定要用循环控制训练节奏,那得在每次循环里重新生成输入函数,确保每次训练都用一个全新的数据集迭代器:for i in range(15): # 每次循环都重新定义输入函数 train_input = tf.estimator.inputs.numpy_input_fn( x={'x': sst_train}, y=precip_train, shuffle=True, batch_size=100, num_epochs=None ) est.train(input_fn=train_input, steps=20)这样每次调用
est.train时,都是用新的迭代器,能准确跑完20步后进入下一次循环,15次后完成训练。
另外提一句,如果你不需要无限重复数据,也可以把num_epochs设为具体值(比如每次循环设为1),不过这种方式需要根据数据集大小和batch_size计算步数,不如前两种省心。
内容的提问来源于stack exchange,提问作者Duan Shiheng

