TensorFlow Estimator设置n_classes=10却被识别为1的报错问题
解决TensorFlow Estimator中n_classes设置与标签不匹配的问题
嘿,我一眼就看出问题出在哪了——你用了one_hot=True加载MNIST数据,但DNNClassifier默认不识别这种格式的标签!
问题根源
你设置了n_classes=10没错,但加载数据时one_hot=True会把标签转换成10维的one-hot编码数组(比如数字5对应[0,0,0,0,0,1,0,0,0,0])。而DNNClassifier默认期望的是单维的整数标签(比如数字5直接用5表示)。这种格式不匹配,就导致Estimator错误地判断你要做的是二分类(n_classes=1),从而抛出那个报错。
两种快速解决方法
方法1:把one-hot标签转成整数标签
只需要在每次获取标签后,用np.argmax()把10维的one-hot数组转换成对应的整数索引就行:
修改训练循环里的代码:
for i in range(100000): xdata, ydata = mnist.train.next_batch(500) # 关键:将one-hot标签转换为整数标签 ydata = np.argmax(ydata, axis=1) train_input_fn = tf.estimator.inputs.numpy_input_fn( x={"x":xdata}, y=ydata, num_epochs=None, shuffle=True) classifier.train(input_fn=train_input_fn, steps=2000)
测试部分也要同步修改:
test_input_fn = tf.estimator.inputs.numpy_input_fn( x= {"x":mnist.test.images}, y= np.argmax(mnist.test.labels, axis=1), # 同样转换测试标签 num_epochs=1, shuffle=False)
方法2:让Estimator支持one-hot标签
如果你不想修改标签格式,可以在创建DNNClassifier时添加multi_class=True参数,明确告诉它要处理多分类的one-hot标签:
classifier = tf.estimator.DNNClassifier( feature_columns=feature_columns, hidden_units=[500, 500, 500], n_classes=10, model_dir="/tmp/MT", multi_class=True # 添加这个参数开启one-hot标签支持 )
额外小提醒
- 你的训练循环
for i in range(100000)每次都跑2000步,总步数会达到20亿,完全没必要,建议把循环次数调小(比如range(10))或者把steps改成1,避免训练时间过长。 - 如果之前运行过代码,
/tmp/MT目录里会残留旧的模型参数,可能导致奇怪的冲突,建议先删除这个目录再重新运行。
内容的提问来源于stack exchange,提问作者5Volts
相关产品推荐
相关产品推荐

