如何在多优化器场景下使用tf.train.SyncReplicasOptimizer?
我之前也碰到过类似的多优化器下同步分布式训练的问题,给你整理一套可行的实现方案,核心思路是给每个独立的基础优化器都套上SyncReplicasOptimizer,统一处理梯度聚合,同时合理管理global_step更新和chief节点的队列运行器:
步骤1:为每个基础优化器包装同步优化器
因为你需要给网络不同部分用不同学习率,所以每个基础优化器都需要单独被SyncReplicasOptimizer包裹,这样不同部分的梯度能分别在所有worker间完成聚合:
# 定义基础优化器(和你原来的一致) base_opt_conv = tf.train.MomentumOptimizer(learning_rate, args.momentum) base_opt_fc_w = tf.train.MomentumOptimizer(learning_rate * 10.0, args.momentum) base_opt_fc_b = tf.train.MomentumOptimizer(learning_rate * 20.0, args.momentum) # 为每个基础优化器套上SyncReplicasOptimizer sync_opt_conv = tf.train.SyncReplicasOptimizer( base_opt_conv, replicas_to_aggregate=num_replicas_to_aggregate, total_num_replicas=num_workers, # 如果需要对卷积层变量做移动平均,可添加以下参数 # variable_averages=exp_moving_averager, # variables_to_average=conv_trainable ) sync_opt_fc_w = tf.train.SyncReplicasOptimizer( base_opt_fc_w, replicas_to_aggregate=num_replicas_to_aggregate, total_num_replicas=num_workers, # variables_to_average=fc_w_trainable ) sync_opt_fc_b = tf.train.SyncReplicasOptimizer( base_opt_fc_b, replicas_to_aggregate=num_replicas_to_aggregate, total_num_replicas=num_workers, # variables_to_average=fc_b_trainable )
步骤2:处理梯度应用与global_step更新
这里要注意只在其中一个梯度应用操作中更新global_step,避免每步训练多次更新步数:
# 计算梯度(和你原来的逻辑一致) grads = tf.gradients(reduced_loss, conv_trainable + fc_w_trainable + fc_b_trainable) grads_conv = grads[:len(conv_trainable)] grads_fc_w = grads[len(conv_trainable) : len(conv_trainable)+len(fc_w_trainable)] grads_fc_b = grads[len(conv_trainable)+len(fc_w_trainable):] # 应用梯度,仅在卷积层的操作中更新global_step train_op_conv = sync_opt_conv.apply_gradients( zip(grads_conv, conv_trainable), global_step=global_step # 唯一更新global_step的地方 ) train_op_fc_w = sync_opt_fc_w.apply_gradients(zip(grads_fc_w, fc_w_trainable)) train_op_fc_b = sync_opt_fc_b.apply_gradients(zip(grads_fc_b, fc_b_trainable)) # 合并所有训练操作为一个整体 train_op = tf.group(train_op_conv, train_op_fc_w, train_op_fc_b)
步骤3:配置chief节点的队列运行器与同步初始化
每个SyncReplicasOptimizer都会生成自己的梯度聚合队列,我们可以用Session Hook来自动管理这些队列的启动,同时在chief节点初始化同步所需的token变量:
# 收集所有同步优化器的Session Hook hooks = [] hooks.append(sync_opt_conv.make_session_run_hook(is_chief)) hooks.append(sync_opt_fc_w.make_session_run_hook(is_chief)) hooks.append(sync_opt_fc_b.make_session_run_hook(is_chief)) # 在chief节点初始化同步所需的token队列 if is_chief: sync_init_ops = [ sync_opt_conv.get_init_tokens_op(), sync_opt_fc_w.get_init_tokens_op(), sync_opt_fc_b.get_init_tokens_op() ] # 添加初始化hook class SyncInitHook(tf.train.SessionRunHook): def after_create_session(self, session, coord): session.run(sync_init_ops) hooks.append(SyncInitHook()) # 启动训练会话 with tf.train.MonitoredTrainingSession( master=master, is_chief=is_chief, hooks=hooks, checkpoint_dir="./checkpoints" # 根据你的需求设置 ) as sess: while not sess.should_stop(): sess.run(train_op)
关键注意点
- 每个优化器单独包装:不同网络部分的梯度需要独立聚合,才能保证各自的学习率策略生效;
- global_step唯一更新:多个操作同时更新
global_step会导致训练步数统计错误,只在一个操作中指定即可; - 同步变量初始化:chief节点必须初始化token队列,否则worker会一直等待聚合信号无法启动训练;
- 移动平均按需配置:如果需要对变量做指数移动平均,要给每个
SyncReplicasOptimizer指定对应的variables_to_average,避免交叉更新。
内容的提问来源于stack exchange,提问作者Ash
相关产品推荐
相关产品推荐

