TensorFlow Estimator自定义input_fn无法正常工作问题求助
解决自定义input_fn替代tf.estimator.inputs.numpy_input_fn的问题
嘿,看你在把TensorFlow示例CNN的输入函数从官方的tf.estimator.inputs.numpy_input_fn改成自定义版本时踩坑了,我来帮你捋捋可能的问题和解决办法~
首先先把你给出的代码补全(看起来你没写完),应该是类似这样的:
def train_input_fn(x, y): dataset = tf.data.Dataset.from_tensor_slices(({"x": x}, y)) dataset = dataset.shuffle(100000).repeat().batch(100) return dataset # 调用训练时的代码 my_classifier.train( input_fn=lambda: train_input_fn(...) # 这里应该是传入训练数据x_train和y_train )
最可能的几个问题点及修复方案
参数传递错误
你在lambda里必须正确传入训练数据x_train和y_train,如果这些变量不在lambda的作用域里,或者传错了名字,会直接导致找不到数据的错误。正确的调用应该是:my_classifier.train( input_fn=lambda: train_input_fn(x_train, y_train), steps=2000 # 一定要指定训练步数,不然会无限训练下去 )数据类型/格式不匹配
官方的numpy_input_fn会自动处理numpy数组到TensorFlow张量的类型转换,但自定义input_fn需要你自己确保数据类型和模型的特征列匹配。比如如果你的特征列定义的是tf.feature_column.numeric_column("x", dtype=tf.float32),那要提前把训练数据转成对应类型:x_train = x_train.astype(np.float32) y_train = y_train.astype(np.int32) # 分类任务的标签一般是int类型数据集shuffle的buffer设置不合理
你用了shuffle(100000),如果你的训练集总样本数还没这么多,这个设置就没意义,建议直接用训练集的大小作为buffer,这样能保证充分打乱:dataset = dataset.shuffle(buffer_size=len(x_train))特征字典键名不匹配
要确保你创建数据集时的特征字典键名(比如{"x": x}里的"x")和模型特征列的名字完全一致,大小写、拼写都不能错。比如你的特征列如果是:feature_columns = [tf.feature_column.numeric_column("image_input", shape=(28,28))]那特征字典的键就得改成
"image_input",不然会报KeyError。
修正后的完整示例代码
假设你在做MNIST手写数字分类,完整的代码应该是这样的:
import tensorflow as tf import numpy as np from tensorflow.examples.tutorials.mnist import input_data mnist = input_data.read_data_sets("./mnist_data/", one_hot=False) x_train = mnist.train.images.reshape(-1, 28, 28) y_train = mnist.train.labels x_train = x_train.astype(np.float32) y_train = y_train.astype(np.int32) # 定义模型的特征列 feature_columns = [tf.feature_column.numeric_column("x", shape=(28,28))] # 创建分类器 my_classifier = tf.estimator.DNNClassifier( feature_columns=feature_columns, hidden_units=[256, 128], n_classes=10, model_dir="./mnist_cnn_model" ) # 自定义训练输入函数 def train_input_fn(x, y, batch_size=100): dataset = tf.data.Dataset.from_tensor_slices(({"x": x}, y)) dataset = dataset.shuffle(buffer_size=len(x)) dataset = dataset.repeat() dataset = dataset.batch(batch_size) return dataset # 启动训练 my_classifier.train( input_fn=lambda: train_input_fn(x_train, y_train), steps=2000 )
常见错误排查
- 如果遇到
KeyError: 'xxx':立刻检查特征字典的键名和特征列的名字是否完全一致。 - 如果遇到
TypeError: Expected dtype xxx but got dtype yyy:把训练数据转换成特征列指定的类型。 - 如果遇到
OutOfRangeError:说明数据集提前耗尽了,检查repeat()是否正确调用,或者如果是有限重复(比如repeat(5)),要确保训练steps不超过总batch数。
内容的提问来源于stack exchange,提问作者Arsenal591
相关产品推荐
相关产品推荐

