TensorFlow Estimator训练报错TypeError: DatasetV1Adapter is not a callable object问题排查
TensorFlow Estimator训练报错:"unsupported callable"解决方案
问题背景
我在使用TensorFlow Estimator训练模型时遇到了一个类型错误,先给大家展示我的代码和报错详情:
我的Estimator定义
estimator = tf.estimator.Estimator( model_fn=model_fn, model_dir=model_dir, params=None, warm_start_from=warm_start_from, config=tf.estimator.RunConfig( model_dir=model_dir, tf_random_seed=1, save_summary_steps=100, save_checkpoints_secs=1600, # 30分钟 session_config=tf.ConfigProto(allow_soft_placement=True, log_device_placement=True), keep_checkpoint_max=keep_checkpoint_max, keep_checkpoint_every_n_hours=10000, log_step_count_steps=100, experimental_max_worker_delay_secs=100, session_creation_timeout_secs=7200 ) )
训练启动代码
我尝试用以下代码启动训练:
estimator.train(input_fn=train_input_fn, max_steps=max_steps)
这里的train_input_fn是一个已经创建好的数据集实例:
- 实例信息:
<DatasetV1Adapter shapes: {image: (128, 32, 32, 3), label: (128,)}, types: {image: tf.int32, label: tf.int32}> - 类型:
<class 'tensorflow.python.data.ops.dataset_ops.DatasetV1Adapter'>
报错信息
运行后直接抛出类型错误:
Traceback (most recent call last): File "/usr/lib/python3.8/inspect.py", line 1135, in getfullargspec sig = _signature_from_callable(func, File "/usr/lib/python3.8/inspect.py", line 2228, in _signature_from_callable raise TypeError('{!r} is not a callable object'.format(obj)) TypeError: <DatasetV1Adapter shapes: {image: (128, 32, 32, 3), label: (128,)}, types: {image: tf.int32, label: tf.int32}> is not a callable object The above exception was the direct cause of the following exception: Traceback (most recent call last): File "scripts/run_cifar.py", line 182, in <module> fire.Fire() File "/usr/local/lib/python3.8/dist-packages/fire/core.py", line 138, in Fire component_trace = _Fire(component, args, parsed_flag_args, context, name) File "/usr/local/lib/python3.8/dist-packages/fire/core.py", line 466, in _Fire component, remaining_args = _CallAndUpdateTrace( File "/usr/local/lib/python3.8/dist-packages/fire/core.py", line 675, in _CallAndUpdateTrace component = fn(*varargs, **kwargs) File "scripts/run_cifar.py", line 155, in train gpu_utils.run_training( File "/workspace/diffusion_tf/gpu_utils/gpu_utils.py", line 280, in run_training estimator.train(input_fn=train_input_fn, max_steps=max_steps) File "/usr/local/lib/python3.8/dist-packages/tensorflow_estimator/python/estimator/estimator.py", line 370, in train loss = self._train_model(input_fn, hooks, saving_listeners) File "/usr/local/lib/python3.8/dist-packages/tensorflow_estimator/python/estimator/estimator.py", line 1161, in _train_model return self._train_model_default(input_fn, hooks, saving_listeners) File "/usr/local/lib/python3.8/dist-packages/tensorflow_estimator/python/estimator/estimator.py", line 1187, in _train_model_default self._get_features_and_labels_from_input_fn( File "/usr/local/lib/python3.8/dist-packages/tensorflow_estimator/python/estimator/estimator.py", line 1025, in _get_features_and_labels_from_input_fn self._call_input_fn(input_fn, mode)) File "/usr/local/lib/python3.8/dist-packages/tensorflow_estimator/python/estimator/estimator.py", line 1103, in _call_input_fn input_fn_args = function_utils.fn_args(input_fn) File "/usr/local/lib/python3.8/dist-packages/tensorflow_core/python/util/function_utils.py", line 57, in fn_args args = tf_inspect.getfullargspec(fn).args File "/usr/local/lib/python3.8/dist-packages/tensorflow_core/python/util/tf_inspect.py", line 257, in getfullargspec return _getfullargspec(target) File "/usr/lib/python3.8/inspect.py", line 1144, in getfullargspec raise TypeError('unsupported callable') from ex TypeError: unsupported callable
问题分析
这个错误的核心原因很简单:TensorFlow Estimator的train()方法要求input_fn必须是一个可调用的函数,但我直接传入了一个已经实例化好的Dataset对象——它是数据集实例,不是函数,系统自然会报错说它不是可调用对象。
Estimator这么设计是有考量的:它需要在合适的时机(比如创建计算图、初始化会话时)调用input_fn生成数据集,这样能保证在分布式训练、重启训练等场景下,数据集可以正确地重新初始化和加载,而不是依赖一个固定的实例。
解决方案
只需要把这个Dataset包装成一个无参的可调用函数就行,有两种简单的写法:
方法1:定义包装函数
# 把已有的train_input_fn包装成无参函数 def train_input_fn_wrapper(): return train_input_fn # 然后传给estimator.train estimator.train(input_fn=train_input_fn_wrapper, max_steps=max_steps)
方法2:用lambda表达式(更简洁)
# 直接用lambda把数据集包装成可调用对象 estimator.train(input_fn=lambda: train_input_fn, max_steps=max_steps)
这样修改后,Estimator就能正确调用input_fn获取数据集,训练就能正常启动了。
内容的提问来源于stack exchange,提问作者Alessandro
相关产品推荐
相关产品推荐

