单设备TensorFlow分布式A3C训练并行化相关问题咨询
背景与代码
我目前有一台配备1块支持CUDA的GPU和1颗8核CPU的计算机,想要实现强化学习的A3C算法(该算法会并行化计算图与训练环境,并向全局图同步梯度更新),计划采用TensorFlow分布式API完成,我的代码如下:
# Hyperparameter definition ############################################################ epochs = 100000 # Global steps t_max = 1000 # Thread steps max_grad_norm = 0.5 alpha = 0.99 gamma = 0.99 ############################################################# ckpt_path = 'ckpt/checkpoints' log_path = 'logs' def test(ac, load_path, sess): ac.load(load_path) env = gym.make('Breakout-v0') while True: obs = env.reset() state_c, state_h = ac.init_lstm_c, ac.init_lstm_h done = False while not done: action, _, state_c, state_h = ac.step(sess, obs, state_c, state_h) obs, _, done, _ = env.step(action) def train(n_epochs, t_max, gamma, ac, sess): env = gym.make("Breakout-v0") save_dir = os.path.join(log_path) writer = tf.summary.FileWriter(save_dir) for ep in range(n_epochs): ep_obs, ep_disc_rew, m_rew, ep_act, ep_vals, state_c, state_h = process_episode(sess, ac, env, t_max, gamma) log = ac.learn(sess, ep_obs, state_c, state_h, ep_disc_rew, m_rew, ep_act, ep_vals) step = tf.train.get_global_step().eval(session=sess) writer.add_summary(log, global_step=step) def process_episode(sess, ac, env, t_max, gamma): ep_observations, ep_rewards, ep_actions, ep_values = [], [], [], [] done = False t = 0 observation = env.reset() state_c, state_h = ac.c_init, ac.h_init while t < t_max and not done: action, value, state_c, state_h = ac.step(sess, observation, state_c, state_h) ep_observations.append(observation) ep_values.append(value) ep_actions.append(action) observation, reward, done, _ = env.step(action) ep_rewards.append(reward) ep_disc_rewards = discount_rewards(ep_rewards, gamma) t_rew = np.sum(ep_rewards) return ep_observations, ep_disc_rewards, t_rew, ep_actions, ep_values, state_c, state_h if __name__ == '__main__': config = tf.ConfigProto(allow_soft_placement=True) config.gpu_options.allow_growth = True server = tf.train.Server.create_local_server() gs = tf.train.create_global_step(tf.get_default_graph()) env = gym.make("Breakout-v0") ac = ActorCritic(env.observation_space, env.action_space) with tf.train.MonitoredTrainingSession( master=server.target, checkpoint_dir=ckpt_path, save_summaries_steps=None, config=config) as sess: train(epochs, t_max, gamma, ac, sess) sess.stop()
现在我有两个问题需要请教:
MonitoredTrainingSession代码块内的代码是否会在多个worker间并行执行?- 我的计算图是否已在多个worker间复制?若未实现,该如何操作?
问题1解答:MonitoredTrainingSession内的代码不会自动在多个worker并行执行
你现在的代码里,tf.train.Server.create_local_server()只是搭了一个本地分布式服务器的架子,但并没有实际启动多个worker进程/线程来跑训练逻辑。MonitoredTrainingSession主要是帮你管理会话的生命周期——比如自动保存checkpoint、恢复模型、处理会话异常这些,但它不会主动帮你把任务拆分成多个并行的worker任务。
简单说,你当前的代码还是单进程单线程在运行,完全没用到A3C需要的多worker并行特性。
问题2解答:你的计算图没有在多个worker间复制,这里是实现方法
A3C的核心就是每个worker有独立的游戏环境和本地计算图副本,各自收集经验、计算梯度,然后同步更新到全局的参数服务器上。你现在的代码只创建了一个ActorCritic实例,所有计算都挤在同一个图里,完全没实现多worker的图复制。要搞定这个,你需要做这几件事:
第一步:明确分布式角色(参数服务器+多个worker)
在TensorFlow的分布式架构里,得区分「参数服务器(ps)」和「worker」节点。对你这种单机器(1GPU+8核CPU)的情况,把参数服务器放在GPU上(用来存全局参数,处理梯度更新),8个worker分别对应8个CPU核心,各自跑独立的环境和计算逻辑,这样效率最高。
第二步:给每个worker创建独立的计算图副本
你需要启动多个进程/线程作为worker,每个worker都创建自己的ActorCritic实例,并且连接到同一个参数服务器。这里推荐用Python的multiprocessing库来启动多进程,或者用TensorFlow的tf.distribute模块(更现代的方式)。
第三步:修改代码实现多worker并行
给你一个针对现有代码的调整示例,用多进程来实现8个worker的并行:
import multiprocessing def run_worker(job_name, task_index, cluster_spec): # 启动当前节点的服务器 server = tf.train.Server(cluster_spec, job_name=job_name, task_index=task_index) if job_name == 'ps': # 参数服务器只需要等待worker连接,不用做其他事 server.join() else: # worker节点创建自己的会话和模型,开始训练 config = tf.ConfigProto(allow_soft_placement=True) config.gpu_options.allow_growth = True with tf.train.MonitoredTrainingSession( master=server.target, checkpoint_dir=ckpt_path, save_summaries_steps=None, config=config) as sess: env = gym.make("Breakout-v0") ac = ActorCritic(env.observation_space, env.action_space) train(epochs, t_max, gamma, ac, sess) if __name__ == '__main__': # 定义集群:1个参数服务器,8个worker(对应你的8核CPU) cluster = tf.train.ClusterSpec({ 'ps': ['localhost:2222'], 'worker': [f'localhost:{2223+i}' for i in range(8)] }) # 启动参数服务器进程 ps_process = multiprocessing.Process(target=run_worker, args=('ps', 0, cluster)) ps_process.start() # 启动8个worker进程 worker_processes = [] for worker_idx in range(8): p = multiprocessing.Process( target=run_worker, args=('worker', worker_idx, cluster) ) p.start() worker_processes.append(p) # 等待所有进程结束 ps_process.join() for p in worker_processes: p.join()
另外,你需要确保ActorCritic类的learn方法是把本地计算的梯度同步到全局参数上,而不是更新本地参数。可以在创建模型变量时指定collections=[tf.GraphKeys.GLOBAL_VARIABLES],或者用tf.train.replica_device_setter来自动把参数分配到ps节点,计算逻辑分配到worker节点。
最后提一句:A3C是异步更新的,不用等所有worker都计算完梯度再统一更新,每个worker算完就直接更新全局参数,这也是它效率高的原因之一。
内容的提问来源于stack exchange,提问作者Javier Ventajas Hernández

