You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

TensorFlow ModelAverageOptimizer初始化遇dtype参数重复值错误

解决TensorFlow 1.8.0分布式训练中ModelAverageOptimizer的_true_getter() got multiple values for keyword argument 'dtype'错误

我之前在维护TensorFlow 1.8.x版本的分布式训练任务时,也碰到过一模一样的问题——用ModelAverageOptimizer时初始化local_step触发这个TypeError,确实和分布式场景下的device_setter脱不了干系。下面给你两个直接能用的解决方案,再帮你把原因理清楚:

快速解决方案

方案1:修改ModelAverageOptimizer的local_step创建逻辑

这个问题的核心是ModelAverageOptimizer创建local_step时同时传了dtype参数和带dtype的初始化器,而device_setter的拦截逻辑会再自动注入一次dtype,导致重复传递。你可以找到tensorflow/contrib/opt/python/training/average_optimizer.py(或者你自己实现的ModelAverageOptimizer代码),修改创建local_step的代码:

原代码大概是这样:

self._local_step = variable_scope.get_variable(
    "local_step",
    shape=[],
    dtype=dtypes.int64,
    initializer=init_ops.zeros_initializer(),
    trainable=False,
    collections=[ops.GraphKeys.LOCAL_VARIABLES])

改成:

self._local_step = variable_scope.get_variable(
    "local_step",
    shape=[],
    initializer=init_ops.zeros_initializer(dtype=dtypes.int64),
    trainable=False,
    collections=[ops.GraphKeys.LOCAL_VARIABLES])

把dtype参数移到初始化器内部,这样get_variable就不会同时收到两个dtype参数,冲突自然就解决了。

方案2:临时绕过device_setter创建local_step

如果不想修改TensorFlow的源码,也可以在初始化ModelAverageOptimizer之前,临时取消device_setter的作用,创建完优化器再恢复:

# 先保存当前的device函数
original_device = tf.get_default_graph()._device_function

# 临时关闭device setter
tf.get_default_graph()._device_function = None

# 初始化ModelAverageOptimizer
avg_optimizer = tf.contrib.opt.ModelAverageOptimizer(
    opt=tf.train.GradientDescentOptimizer(learning_rate=0.01),
    average_decay=0.999)

# 恢复原来的device setter
tf.get_default_graph()._device_function = original_device

这种方式让local_step的创建绕过device_setter的拦截,也就不会触发参数冲突了。

错误原因拆解

在TensorFlow 1.8.0的分布式device_setter实现里,_true_getter函数会在变量创建时自动注入dtype等参数。而ModelAverageOptimizer内部创建local_step时,同时显式指定了dtype和带dtype的初始化器,导致_true_getter接收到两次dtype参数,直接触发了TypeError。这个bug在TensorFlow 1.10及以后的版本里已经被官方修复了,但1.8.0作为旧版本还得手动处理。

附:报错栈示例

TypeError: _true_getter() got multiple values for keyword argument 'dtype'
  File "train.py", line 123, in <module>
    avg_optimizer = tf.contrib.opt.ModelAverageOptimizer(base_optimizer)
  File "/path/to/anaconda/envs/tf18/lib/python3.6/site-packages/tensorflow/contrib/opt/python/training/average_optimizer.py", line 104, in __init__
    self._local_step = variable_scope.get_variable(...)
  File "/path/to/anaconda/envs/tf18/lib/python3.6/site-packages/tensorflow/python/ops/variable_scope.py", line 1292, in get_variable
    constraint=constraint)
  File "/path/to/anaconda/envs/tf18/lib/python3.6/site-packages/tensorflow/python/ops/variable_scope.py", line 1094, in get_variable
    constraint=constraint)
  File "/path/to/anaconda/envs/tf18/lib/python3.6/site-packages/tensorflow/python/ops/variable_scope.py", line 425, in get_variable
    constraint=constraint)
  File "/path/to/anaconda/envs/tf18/lib/python3.6/site-packages/tensorflow/python/ops/variable_scope.py", line 394, in _true_getter
    caching_device=caching_device, constraint=constraint)
TypeError: _true_getter() got multiple values for keyword argument 'dtype'

附:相关代码片段示例

# 分布式集群配置
cluster_spec = tf.train.ClusterSpec({
    "ps": ["ps0:2222"],
    "worker": ["worker0:2223", "worker1:2224"]
})
server = tf.train.Server(cluster_spec, job_name=job_name, task_index=task_index)

if job_name == "ps":
    server.join()
else:
    with tf.device(tf.train.replica_device_setter(
        worker_device="/job:worker/task:%d" % task_index,
        cluster=cluster_spec)):
        # 构建模型结构
        x = tf.placeholder(tf.float32, shape=[None, 784])
        y = tf.placeholder(tf.float32, shape=[None, 10])
        logits = build_model(x)
        loss = tf.losses.softmax_cross_entropy(onehot_labels=y, logits=logits)
        
        # 初始化基础优化器
        base_opt = tf.train.GradientDescentOptimizer(0.01)
        # 此处触发错误
        avg_opt = tf.contrib.opt.ModelAverageOptimizer(base_opt, average_decay=0.999)
        train_op = avg_opt.minimize(loss, global_step=tf.train.get_or_create_global_step())
        
        # 后续训练流程
        ...

内容的提问来源于stack exchange,提问作者吴培昊

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.29 07:56:00