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

PyTorch多智能体PPO模型价值损失未下降问题排查求助

问题描述

基于PyTorch实现多智能体PPO模型,采用Petting Zoo构建连续搜索场景环境,状态空间为标准化向量,参考单智能体PPO实现框架。目前智能体无学习迹象,价值损失始终未下降,已排除环境与奖励信号问题,需排查代码实现Bug。

动作与价值函数实现

def get_action(self, x, action=None):
    x.to(self.device)
    net = self.network(x)
    dropout = nn.Dropout(0.2)
    action_mean = self.actor_mean(dropout(net))
    action_logstd = self.actor_logstd.expand_as(action_mean)
    action_std = torch.exp(action_logstd).to(self.device)
    probs = Normal(action_mean, action_std)
    if action is None:
        action = probs.sample().to(self.device)
    return action, probs.log_prob(action).sum(1), probs.entropy().sum(1)

def get_value(self, x):
    x.to(self.device)
    return self.critic(self.network(x))

模型结构

self.network = nn.Sequential(
    layer_init(nn.Linear(12, 64)),
    layer_init(nn.Linear(64, 128)),
    layer_init(nn.Linear(128, 64)),
    nn.ReLU(),
).to(self.device)
self.actor_mean = nn.Sequential(layer_init(
    nn.Linear(64, np.prod(self.env.action_space.shape)), std=1).to(self.device),
    nn.Tanh())
self.actor_logstd = nn.Parameter(
    torch.zeros(1, np.prod(self.env.action_space.shape)))
self.critic = nn.Sequential(layer_init(
    nn.Linear(64, 32), std=1).to(self.device),
    nn.Tanh(),
    layer_init(
    nn.Linear(32, 1), std=1).to(self.device))
self.optimizer = optim.Adam(
    self.parameters(), lr=self.args.learning_rate, eps=1e-5, weight_decay=0.1)

训练函数实现

def train(self, plot=False):
    global_step = 0
    start_time = time.time()
    next_obs = self.env.reset()
    next_done = torch.zeros(self.args.num_envs).to(self.device)
    num_updates = self.args.total_timesteps // self.args.batch_size
    for update in range(1, num_updates+1):
        # Annealing the rate if instructed to do so.
        if self.args.anneal_lr:
            frac = 1.0 - (update - 1.0) / num_updates
            lrnow = self.lr(frac)
            self.optimizer.param_groups[0]['lr'] = lrnow

        for step in range(0, self.args.num_steps):
            global_step += 1 * self.args.num_envs
            self.obs[step] = next_obs
            self.dones[step] = next_done

            with torch.no_grad():
                self.values[step] = self.get_value(
                    self.obs[step]).flatten()
                action, logproba, _ = self.get_action(self.obs[step])

            self.actions[step] = action
            self.logprobs[step] = logproba

            next_obs, rs, ds, infos = self.env.step(torch.clamp(
                action, self.env.action_space.low.mean(), self.env.action_space.high.mean()))
            self.rewards[step], next_done = rs.view(
                -1), torch.Tensor(ds).to(self.device)

        # bootstrap reward if not done. reached the batch limit
        with torch.no_grad():
            last_value = self.get_value(
                next_obs.to(self.device)).reshape(1, -1)
            if self.args.gae:
                advantages = torch.zeros_like(self.rewards).to(self.device)
                lastgaelam = 0
                for t in reversed(range(self.args.num_steps)):
                    if t == self.args.num_steps - 1:
                        nextnonterminal = 1.0 - next_done
                        nextvalues = last_value
                    else:
                        nextnonterminal = 1.0 - self.dones[t+1]
                        nextvalues = self.values[t+1]
                    delta = self.rewards[t] + self.args.gamma * \
                        nextvalues * nextnonterminal - self.values[t]
                    advantages[t] = lastgaelam = delta + self.args.gamma * \
                        self.args.gae_lambda * nextnonterminal * lastgaelam
                returns = advantages + self.values
            else:
                returns = torch.zeros_like(self.rewards).to(self.device)
                for t in reversed(range(self.args.num_steps)):
                    if t == self.args.num_steps - 1:
                        nextnonterminal = 1.0 - next_done
                        next_return = last_value
                    else:
                        nextnonterminal = 1.0 - self.dones[t+1]
                        next_return = returns[t+1]
                    returns[t] = self.rewards[t] + self.args.gamma * \
                        nextnonterminal * next_return
                advantages = returns - self.values

        # flatten the batch
        b_obs = self.obs.reshape((-1,)+self.env.observation_space.shape)
        b_logprobs = self.logprobs.reshape(-1)
        b_actions = self.actions.reshape(
            (-1,)+self.env.action_space.shape)
        b_advantages = advantages.reshape(-1)
        b_returns = returns.reshape(-1)
        b_values = self.values.reshape(-1)

        # Optimizaing the policy and value network
        target_agent = Agent(self.env, self.hp_dict).to(self.device)
        inds = np.arange(self.args.batch_size,)
        stopped = 0
        for i_epoch_pi in range(self.args.update_epochs):
            np.random.shuffle(inds)
            target_agent.load_state_dict(self.state_dict())
            for start in range(0, self.args.batch_size, self.args.minibatch_size):
                end = start + self.args.minibatch_size
                minibatch_ind = inds[start:end]
                mb_advantages = b_advantages[minibatch_ind]
                if self.args.norm_adv:
                    mb_advantages = (mb_advantages - mb_advantages.mean()
                                     ) / (mb_advantages.std() + 1e-8)

                _, newlogproba, entropy = self.get_action(
                    b_obs[minibatch_ind], b_actions[minibatch_ind])
                ratio = (newlogproba - b_logprobs[minibatch_ind]).exp()

                # Stats
                approx_kl = (
                    b_logprobs[minibatch_ind] - newlogproba).mean()

                # Policy loss
                pg_loss1 = -mb_advantages * ratio
                pg_loss2 = -mb_advantages * \
                    torch.clamp(ratio, 1-self.args.clip_coef,
                                1+self.args.clip_coef)
                pg_loss = torch.max(pg_loss1, pg_loss2).mean()
                entropy_loss = entropy.mean()

                # Value loss
                new_values = self.get_value(b_obs[minibatch_ind]).view(-1)
                if self.args.clip_vloss:
                    v_loss_unclipped = (
                        (new_values - b_returns[minibatch_ind]) ** 2)
                    v_clipped = b_values[minibatch_ind] + torch.clamp(
                        new_values - b_values[minibatch_ind], -self.args.clip_coef, self.args.clip_coef)
                    v_loss_clipped = (
                        v_clipped - b_returns[minibatch_ind])**2
                    v_loss_max = torch.max(
                        v_loss_unclipped, v_loss_clipped)
                    v_loss = 0.5 * v_loss_max.mean()
                else:
                    v_loss = 0.5 * \
                        ((new_values -
                         b_returns[minibatch_ind]) ** 2).mean()

                loss = pg_loss - self.args.ent_coef * entropy_loss + v_loss * self.args.vf_coef

                self.optimizer.zero_grad()
                loss.backward()
                nn.utils.clip_grad_norm_(
                    self.parameters(), self.args.max_grad_norm)
                self.optimizer.step()

            if self.args.kle_stop:
                if approx_kl > self.args.target_kl:
                    stopped += 1
                    break
            if self.args.kle_rollback:
                if (b_logprobs[minibatch_ind] - self.get_action(b_obs[minibatch_ind], b_actions[minibatch_ind])[1]).mean() > self.args.target_kl:
                    self.load_state_dict(target_agent.state_dict())
                    break

潜在Bug排查点

  1. Dropout层实现错误:get_action中每次调用都新建nn.Dropout(0.2),该层未加入模型参数,且critic计算时未使用Dropout,导致actor与critic共享的backbone输入不一致,价值函数学习目标混乱。

    • 修复:将Dropout层移至self.network的Sequential结构中,作为固定层。
  2. 设备迁移未赋值:get_action和get_value中x.to(self.device)未赋值给x,原张量未转移到目标设备,可能导致网络输入与参数设备不匹配,梯度无法正常回传。

    • 修复:改为x = x.to(self.device)。
  3. 权重初始化不合理:layer_init设置std=1,尤其是critic最后一层初始化标准差过大,导致初始价值预测偏离实际回报范围,价值损失难以优化。

    • 修复:调整初始化策略,actor输出层用std=0.01,critic输出层用std=1e-3。
  4. 权重衰减过大:Adam优化器设置weight_decay=0.1,强L2正则压制价值函数的学习信号,导致参数无法更新到合理值。

    • 修复:将weight_decay降至1e-5或移除该参数。
  5. KL散度早停逻辑错误:kle_stop仅用最后一个minibatch的KL值判断早停,未计算整个epoch的平均KL;kle_rollback中重新计算KL时未使用torch.no_grad(),产生冗余计算图。

    • 修复:epoch内累积KL值求平均后判断早停;计算KL时添加torch.no_grad()上下文。
  6. 多智能体维度适配缺失:参考单智能体实现未处理多智能体的观测/回报维度,可能将多智能体数据混在一起训练,导致价值函数无法学习到每个智能体的状态价值。

    • 修复:确认数据维度是否包含智能体标识,计算returns和advantages时按智能体维度单独处理。

内容的提问来源于stack exchange,提问作者Herman Mostein

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 05:07:00