使用LinearRegressor.train()报错‘...is not a callable object’的问题
问题分析与解决方案
你遇到的问题核心是对TensorFlow train()方法要求的input_fn参数理解有误,两种写法的本质区别在于传递给train()的是“可调用函数”还是“函数执行后的结果”:
错误原因
第一种写法里,lambda: my_input_fn(my_feature, targets)传递的是一个无参的可调用函数——当train()需要数据时,它会主动调用这个lambda,进而执行my_input_fn获取张量。
而第二种写法中,temp_my_input_fn(my_feature, targets)是直接执行了函数,返回的是(features, labels)这个张量元组,而不是一个可调用对象。train()方法尝试把这个元组当成函数去调用,自然就抛出了“不是可调用对象”的错误。
修正方案
有两种简单的修正方式:
方式1:用lambda包装(和第一种写法逻辑对齐)
保持嵌套函数结构不变,调用train()时用lambda把参数绑定成无参可调用对象:
def get_my_input_fn() : def my_input_func(features, targets, batch_size=1, shuffle=True, num_epochs=None) : ... # 你的原有数据处理逻辑 features, labels = ds.make_one_shot_iterator().get_next() return features, labels return my_input_func temp_my_input_fn = get_my_input_fn() # 用lambda包装,让train()能调用这个无参函数 _ = linear_regressor.train(input_fn=lambda: temp_my_input_fn(my_feature, targets), steps=100)
方式2:修改嵌套函数,返回绑定参数的闭包
如果想避免显式写lambda,可以让外层函数接收features和targets,内层函数变成直接使用外层变量的无参闭包:
def get_my_input_fn(features, targets) : def my_input_func(batch_size=1, shuffle=True, num_epochs=None) : ... # 直接使用外层传入的features和targets features, labels = ds.make_one_shot_iterator().get_next() return features, labels return my_input_func # 提前绑定参数,得到无参的可调用函数 temp_my_input_fn = get_my_input_fn(my_feature, targets) _ = linear_regressor.train(input_fn=temp_my_input_fn, steps=100)
关键原理
TensorFlow的Estimator.train()要求input_fn是一个不需要传入参数的可调用对象(函数、lambda或者实现了__call__的对象),它会在训练过程中多次调用这个对象来获取批次数据。你必须确保传递的是“函数本身”,而不是“函数执行后的结果”。
内容的提问来源于stack exchange,提问作者user12187347
相关产品推荐
相关产品推荐

