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

单设备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()

现在我有两个问题需要请教:

  1. MonitoredTrainingSession代码块内的代码是否会在多个worker间并行执行?
  2. 我的计算图是否已在多个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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 04:35:30