MARL实战:用Python构建集中式训练分布式执行架构

如果你已经玩过单智能体强化学习,比如让一个AI学会走迷宫或者下围棋,那么当你第一次接触多智能体强化学习时,那种感觉就像是突然从单人游戏切换到了多人联机模式。事情变得复杂,也变得更加有趣。每个智能体不再是一个孤岛,它们的观察是局部的,动作会相互影响,奖励可能彼此冲突或协同。如何让一群AI学会合作或竞争,共同完成一个任务,是MARL要解决的核心问题。

在众多MARL架构中,集中式训练分布式执行 脱颖而出,成为当前最主流、也最实用的范式。它巧妙地解决了“训练时需要全局视野”和“执行时需要独立决策”这一对矛盾。简单来说,就是训练时,我们有一个“上帝视角”的中央大脑,它能收集所有智能体的信息,进行更优的策略评估和更新;而执行时,每个智能体则依靠自己独立的“小脑”进行决策,无需实时与中央通信,保证了效率。这就像一支足球队,平时训练时教练(中央控制器)能看到所有队员的表现并进行战术指导,但真正比赛时,每个队员必须根据场上瞬息万变的情况,独立做出传球、射门等决策。

本文面向有一定Python和强化学习基础的开发者,我们将抛开理论,直接动手。从零开始,用代码一步步搭建一个CTDE架构的MARL系统。我们会涵盖环境搭建、网络结构设计、核心通信机制实现,并提供可直接运行、修改的代码示例。目标是让你不仅能理解CTDE,更能亲手把它实现出来,应用到自己的项目中。

1. 环境搭建与问题定义

在开始写智能体之前,我们首先要为它们创造一个“世界”。一个设计良好的多智能体环境是成功的一半。对于CTDE架构,环境需要满足几个关键特性:每个智能体应有局部观察,环境状态应能被中央控制器获取以用于训练,同时能提供个体奖励和全局奖励

1.1 选择与定制多智能体环境

虽然OpenAI Gym提供了丰富的单智能体环境,但对于MARL,我们通常需要更专门的平台。PettingZoo是一个优秀的库,它提供了大量标准化的多智能体环境。不过,为了彻底理解环境与智能体间的交互,我们将自己实现一个简化但功能完整的自定义环境。

我们设计一个经典的“协作导航”场景:在一个二维网格世界中,有多个智能体(点)和多个固定目标(星星)。智能体的任务是移动到任意一个目标位置。当所有目标都被至少一个智能体占据时,任务完成。这个环境包含了合作、局部观察(每个智能体只能看到周围一定范围内的区域)以及智能体间潜在的路径冲突等MARL典型元素。

import numpy as np
import gym
from gym import spaces

class CooperativeNavigationEnv(gym.Env):
    """
    自定义协作导航多智能体环境。
    """
    def __init__(self, world_size=10, num_agents=3, num_landmarks=3, obs_range=5):
        super(CooperativeNavigationEnv, self).__init__()
        self.world_size = world_size
        self.num_agents = num_agents
        self.num_landmarks = num_landmarks
        self.obs_range = obs_range  # 智能体的局部观察范围

        # 动作空间:每个智能体可以向上、下、左、右移动或停留
        self.action_space = spaces.Discrete(5)  # 适用于每个智能体

        # 观察空间:对于每个智能体,观察包括:
        # 1. 自身位置 (2维)
        # 2. 自身速度 (2维,由动作决定,这里简化)
        # 3. 在观察范围内的其他智能体的相对位置 (最多 (num_agents-1)*2 维,用0填充)
        # 4. 在观察范围内的地标的相对位置 (最多 num_landmarks*2 维,用0填充)
        # 我们将所有维度拼接成一个一维向量
        max_obs_dim = 2 + 2 + (self.num_agents - 1) * 2 + self.num_landmarks * 2
        self.observation_space = spaces.Box(low=-np.inf, high=np.inf, shape=(max_obs_dim,), dtype=np.float32)

        # 状态空间(全局,仅用于中央控制器训练)
        # 包含所有智能体和所有地标的绝对坐标
        self.state_space = spaces.Box(low=0, high=world_size, shape=( (num_agents+num_landmarks)*2, ), dtype=np.float32)

        self.reset()

    def reset(self):
        """重置环境,随机初始化智能体和地标位置。"""
        # 随机生成不重叠的智能体位置
        self.agent_positions = np.random.rand(self.num_agents, 2) * self.world_size
        # 随机生成不重叠的地标位置
        self.landmark_positions = np.random.rand(self.num_landmarks, 2) * self.world_size
        self.steps = 0
        self.max_steps = 50
        return self._get_global_state()

    def _get_global_state(self):
        """获取全局状态,用于中央控制器。"""
        # 将智能体和地标的位置扁平化拼接
        state = np.concatenate([self.agent_positions.flatten(), self.landmark_positions.flatten()])
        return state

    def _get_agent_obs(self, agent_id):
        """获取指定智能体的局部观察。"""
        obs = []
        agent_pos = self.agent_positions[agent_id]

        # 1. 自身位置(归一化到[0,1])
        obs.append(agent_pos / self.world_size)

        # 2. 自身速度(本例中简化,用上一动作的方向向量表示,初始为0)
        # 为简化,我们暂时用零向量占位
        obs.append(np.array([0.0, 0.0]))

        # 3. 其他智能体的相对位置(如果在观察范围内)
        other_agents_rel = []
        for i in range(self.num_agents):
            if i == agent_id:
                continue
            rel_pos = self.agent_positions[i] - agent_pos
            distance = np.linalg.norm(rel_pos)
            if distance <= self.obs_range:
                # 归一化相对位置
                other_agents_rel.append(rel_pos / self.obs_range)
            else:
                # 超出范围,用零向量表示“不可见”
                other_agents_rel.append(np.array([0.0, 0.0]))
        # 填充或截断,确保维度固定
        obs.append(np.array(other_agents_rel).flatten())

        # 4. 地标的相对位置(如果在观察范围内)
        landmarks_rel = []
        for lm_pos in self.landmark_positions:
            rel_pos = lm_pos - agent_pos
            distance = np.linalg.norm(rel_pos)
            if distance <= self.obs_range:
                landmarks_rel.append(rel_pos / self.obs_range)
            else:
                landmarks_rel.append(np.array([0.0, 0.0]))
        obs.append(np.array(landmarks_rel).flatten())

        # 拼接所有观察部分
        full_obs = np.concatenate(obs)
        return full_obs

    def step(self, actions):
        """
        执行所有智能体的动作。
        :param actions: 一个列表,包含每个智能体的动作索引(0-4)。
        :return: (observations, rewards, done, info)
        """
        assert len(actions) == self.num_agents
        self.steps += 1

        # 定义动作到位移向量的映射
        action_map = {
            0: np.array([0, 1]),   # 上
            1: np.array([0, -1]),  # 下
            2: np.array([-1, 0]),  # 左
            3: np.array([1, 0]),   # 右
            4: np.array([0, 0])    # 停留
        }

        # 更新每个智能体的位置
        new_positions = np.copy(self.agent_positions)
        for i, act in enumerate(actions):
            displacement = action_map[act]
            new_pos = self.agent_positions[i] + displacement
            # 边界检查
            new_pos = np.clip(new_pos, 0, self.world_size)
            new_positions[i] = new_pos
        self.agent_positions = new_positions

        # 计算奖励
        rewards = self._compute_rewards()

        # 获取所有智能体的局部观察
        observations = [self._get_agent_obs(i) for i in range(self.num_agents)]

        # 判断回合是否结束:所有地标都被占据或达到最大步数
        done = self._is_done()

        # 额外信息:全局状态,用于集中式训练
        info = {'global_state': self._get_global_state()}

        return observations, rewards, done, info

    def _compute_rewards(self):
        """计算每个智能体的奖励。这里设计一个合作型奖励。"""
        rewards = np.zeros(self.num_agents)
        # 计算每个智能体到最近地标的距离
        for i, agent_pos in enumerate(self.agent_positions):
            distances = np.linalg.norm(self.landmark_positions - agent_pos, axis=1)
            min_distance = np.min(distances)
            # 奖励与距离负相关,鼓励靠近地标
            rewards[i] = -0.1 * min_distance
        # 额外的全局合作奖励:如果所有地标都被占据,给予所有智能体一大笔奖励
        if self._all_landmarks_covered():
            rewards += 10.0
        return rewards.tolist()

    def _all_landmarks_covered(self):
        """检查是否每个地标都至少被一个智能体占据(距离小于阈值)。"""
        coverage_threshold = 0.5  # 占据的距离阈值
        for lm_pos in self.landmark_positions:
            distances = np.linalg.norm(self.agent_positions - lm_pos, axis=1)
            if np.min(distances) > coverage_threshold:
                return False
        return True

    def _is_done(self):
        return self._all_landmarks_covered() or (self.steps >= self.max_steps)

    def render(self, mode='human'):
        """简单可视化(文本或未来可扩展为图形)。"""
        grid = np.full((self.world_size+1, self.world_size+1), '.', dtype=str)
        # 标记地标
        for idx, (x, y) in enumerate(self.landmark_positions):
            ix, iy = int(x), int(y)
            grid[ix, iy] = '*'
        # 标记智能体
        for idx, (x, y) in enumerate(self.agent_positions):
            ix, iy = int(x), int(y)
            grid[ix, iy] = str(idx)
        for row in range(self.world_size, -1, -1):
            print(' '.join(grid[:, row]))
        print('---')

注意:这个自定义环境是一个高度简化的示例,旨在阐明MARL环境的核心接口。在实际项目中,你可能需要处理更复杂的观察空间(如图像)、连续动作空间、异步执行等问题。PettingZooSMAC等库提供了更成熟、标准化的环境。

1.2 安装必要的依赖库

我们的实现将主要依赖PyTorch来构建神经网络,以及Gym作为环境接口。确保你已经安装了它们。

# 使用pip安装核心依赖
pip install torch gym numpy

对于更复杂的可视化,你可能还需要matplotlib。但为了保持核心代码的简洁,我们的示例将主要使用文本渲染。

2. 网络结构设计:个体策略与中央价值网络

CTDE架构的核心在于其独特的网络分工:每个智能体拥有自己的策略网络,用于分布式执行;同时,存在一个中央价值网络,在训练时利用全局信息评估状态-动作对的价值。

2.1 个体策略网络

每个智能体的策略网络(Actor)输入其局部观察,输出动作的概率分布(离散动作)或动作参数(连续动作)。网络结构通常是一个多层感知机。

import torch
import torch.nn as nn
import torch.nn.functional as F

class PolicyNetwork(nn.Module):
    """
    个体策略网络(Actor)。
    输入:局部观察 (obs_dim)
    输出:离散动作的概率分布 (act_dim)
    """
    def __init__(self, obs_dim, act_dim, hidden_dim=128):
        super(PolicyNetwork, self).__init__()
        self.fc1 = nn.Linear(obs_dim, hidden_dim)
        self.fc2 = nn.Linear(hidden_dim, hidden_dim)
        self.fc3 = nn.Linear(hidden_dim, act_dim)

    def forward(self, obs):
        """
        :param obs: 局部观察张量 [batch_size, obs_dim]
        :return: 动作概率 logits [batch_size, act_dim]
        """
        x = F.relu(self.fc1(obs))
        x = F.relu(self.fc2(x))
        logits = self.fc3(x)  # 未归一化的logits
        return logits

    def get_action(self, obs, deterministic=False):
        """
        根据观察采样一个动作。
        :param obs: 单个观察向量 [obs_dim]
        :param deterministic: 如果为True,选择概率最大的动作;否则按概率采样。
        :return: 动作索引 (int), 动作的log概率 (float)
        """
        obs_tensor = torch.FloatTensor(obs).unsqueeze(0)  # [1, obs_dim]
        logits = self.forward(obs_tensor)  # [1, act_dim]
        probs = F.softmax(logits, dim=-1).squeeze(0)  # [act_dim]
        dist = torch.distributions.Categorical(probs)

        if deterministic:
            action = torch.argmax(probs).item()
            log_prob = dist.log_prob(torch.tensor(action))
        else:
            action = dist.sample().item()
            log_prob = dist.log_prob(torch.tensor(action))

        return action, log_prob.item()

这个策略网络是每个智能体私有的。在分布式执行阶段,智能体i仅使用自己的PolicyNetwork_i,根据局部观察o_i独立做出决策a_i,无需与其他智能体或中央节点通信。

2.2 中央价值网络

中央价值网络(Critic)是CTDE的训练核心。它的输入是全局状态s所有智能体的联合动作a,输出一个标量值,用于评估在状态s下采取联合动作a的好坏。这个网络只在训练时使用,并且需要收集所有智能体的信息。

class CentralizedValueNetwork(nn.Module):
    """
    中央价值网络(Critic)。
    输入:全局状态 (state_dim) + 所有智能体的联合动作 (num_agents * act_dim)
    输出:状态-动作值 Q(s, a) 标量
    """
    def __init__(self, state_dim, num_agents, act_dim, hidden_dim=256):
        super(CentralizedValueNetwork, self).__init__()
        self.num_agents = num_agents
        total_action_dim = num_agents * act_dim
        input_dim = state_dim + total_action_dim

        self.fc1 = nn.Linear(input_dim, hidden_dim)
        self.fc2 = nn.Linear(hidden_dim, hidden_dim)
        self.fc3 = nn.Linear(hidden_dim, 1)  # 输出一个Q值

    def forward(self, state, actions):
        """
        :param state: 全局状态张量 [batch_size, state_dim]
        :param actions: 联合动作张量 [batch_size, num_agents, act_dim] (如果是离散动作,需要是one-hot)
                        或 [batch_size, num_agents] (动作索引,需转换为one-hot)
        :return: Q值 [batch_size, 1]
        """
        batch_size = state.shape[0]
        # 确保actions是one-hot形式
        if actions.dim() == 2 and actions.shape[-1] != self.num_agents * self.act_dim:
            # 假设actions是动作索引 [batch_size, num_agents]
            # 将其转换为one-hot [batch_size, num_agents, act_dim],再展平
            actions_one_hot = F.one_hot(actions.long(), num_classes=self.act_dim).float()  # [batch_size, num_agents, act_dim]
            actions_flat = actions_one_hot.view(batch_size, -1)  # [batch_size, num_agents * act_dim]
        else:
            # 假设已经是展平的形式或合适的形状,这里简化处理
            actions_flat = actions.view(batch_size, -1)

        # 拼接状态和联合动作
        x = torch.cat([state, actions_flat], dim=-1)  # [batch_size, state_dim + num_agents*act_dim]
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        q_value = self.fc3(x)  # [batch_size, 1]
        return q_value

提示:对于离散动作,在输入中央价值网络前,通常需要将动作索引转换为one-hot向量,以便网络能更好地处理。对于连续动作,则直接拼接动作值即可。

中央价值网络的设计允许在训练时利用全局信息(s, a)来更准确地评估动作价值,从而引导各个独立策略网络的学习。这是CTDE比完全分布式方法性能更好的关键。

3. 通信与数据流:实现CTDE的训练循环

CTDE的训练流程可以概括为:分布式采集数据,集中式优化策略。我们需要一个协调器来管理这个过程。

3.1 经验回放缓冲区

为了稳定训练,我们使用经验回放缓冲区来存储智能体与环境交互产生的轨迹数据。对于CTDE,我们需要存储的数据包括:

  • 全局状态 s_t
  • 所有智能体的联合动作 a_t (由各自策略网络产生)
  • 奖励向量 r_t (每个智能体的奖励)
  • 下一个全局状态 s_{t+1}
  • 终止标志 done
from collections import deque
import random

class CTDE_ReplayBuffer:
    """
    为CTDE架构设计的经验回放缓冲区。
    存储 (s, a, r, s', done) 元组,其中a和r是向量形式。
    """
    def __init__(self, capacity):
        self.buffer = deque(maxlen=capacity)

    def push(self, state, joint_actions, rewards, next_state, done):
        """
        存储一条经验。
        :param state: 全局状态 (np.array)
        :param joint_actions: 所有智能体的动作列表 [act1, act2, ...]
        :param rewards: 所有智能体的奖励列表 [r1, r2, ...]
        :param next_state: 下一个全局状态 (np.array)
        :param done: 布尔值
        """
        experience = (state.copy(),
                      joint_actions.copy(),
                      rewards.copy(),
                      next_state.copy(),
                      done)
        self.buffer.append(experience)

    def sample(self, batch_size):
        """
        随机采样一批经验。
        :return: 元组 (states, actions, rewards, next_states, dones)
                 所有数据都转换为PyTorch张量。
        """
        if len(self.buffer) < batch_size:
            return None

        batch = random.sample(self.buffer, batch_size)
        states, actions, rewards, next_states, dones = zip(*batch)

        # 转换为张量
        states_t = torch.FloatTensor(np.array(states))
        # 注意:actions是一个列表的列表,需要转换为张量 [batch, num_agents]
        actions_t = torch.LongTensor(np.array(actions))  # 假设是离散动作索引
        rewards_t = torch.FloatTensor(np.array(rewards))
        next_states_t = torch.FloatTensor(np.array(next_states))
        dones_t = torch.FloatTensor(np.array(dones))

        return states_t, actions_t, rewards_t, next_states_t, dones_t

    def __len__(self):
        return len(self.buffer)

3.2 训练算法:MAPPO示例

我们将实现一个简化版的多智能体近端策略优化算法,这是CTDE架构下非常流行且有效的算法。其核心思想是:利用中央价值网络计算的优势函数来更新各个智能体的策略网络,同时通过裁剪等技巧保证策略更新的稳定性。

首先,定义我们的多智能体PPO(MAPPO)智能体类,它将管理所有个体策略网络和一个中央价值网络。

class MAPPO_Agent:
    def __init__(self, state_dim, obs_dim_list, act_dim_list, num_agents, lr_actor=3e-4, lr_critic=1e-3, gamma=0.99, clip_param=0.2):
        """
        :param state_dim: 全局状态维度
        :param obs_dim_list: 每个智能体的观察维度列表
        :param act_dim_list: 每个智能体的动作维度列表(假设所有智能体动作空间相同)
        :param num_agents: 智能体数量
        """
        self.num_agents = num_agents
        self.gamma = gamma
        self.clip_param = clip_param

        # 创建每个智能体的策略网络
        self.actors = []
        self.actor_optimizers = []
        for i in range(num_agents):
            actor = PolicyNetwork(obs_dim_list[i], act_dim_list[i])
            self.actors.append(actor)
            self.actor_optimizers.append(torch.optim.Adam(actor.parameters(), lr=lr_actor))

        # 创建中央价值网络
        # 假设所有智能体动作空间相同
        act_dim = act_dim_list[0]
        self.critic = CentralizedValueNetwork(state_dim, num_agents, act_dim)
        self.critic_optimizer = torch.optim.Adam(self.critic.parameters(), lr=lr_critic)

        # 经验缓冲区
        self.buffer = CTDE_ReplayBuffer(capacity=10000)

    def choose_actions(self, observations):
        """
        分布式执行:每个智能体根据自身观察选择动作。
        :param observations: 列表,包含每个智能体的观察 (np.array)
        :return: 动作列表,每个动作的log概率列表
        """
        actions = []
        log_probs = []
        for i in range(self.num_agents):
            act, log_prob = self.actors[i].get_action(observations[i])
            actions.append(act)
            log_probs.append(log_prob)
        return actions, log_probs

    def store_transition(self, state, joint_actions, rewards, next_state, done):
        """存储一条经验到缓冲区。"""
        self.buffer.push(state, joint_actions, rewards, next_state, done)

    def update(self, batch_size=64, update_epochs=10):
        """
        集中式训练:从缓冲区采样,更新中央价值网络和所有策略网络。
        这是一个简化的PPO更新步骤。
        """
        if len(self.buffer) < batch_size:
            return

        for _ in range(update_epochs):
            # 采样一批数据
            batch = self.buffer.sample(batch_size)
            if batch is None:
                break
            states, actions, rewards, next_states, dones = batch
            # rewards形状: [batch, num_agents]
            # 我们使用团队平均奖励作为中央Critic的奖励信号,这是一种常见简化
            team_rewards = rewards.mean(dim=1, keepdim=True)  # [batch, 1]

            # --- 更新中央价值网络 (Critic) ---
            # 计算当前Q值
            current_q_values = self.critic(states, actions)  # [batch, 1]
            # 计算目标Q值 (使用目标网络会更好,这里为简化省略)
            with torch.no_grad():
                # 注意:这里需要下一个状态和下一个动作。下一个动作应由目标策略网络产生。
                # 为简化,我们使用当前策略网络,并假设下一个动作是贪婪动作(或从当前策略采样)。
                # 更严谨的做法是使用目标网络和贝尔曼方程。
                next_actions = []
                # 这里需要下一个状态的观察来产生下一个动作,但我们只有全局状态。
                # 这暴露了CTDE的一个实现细节:Critic训练需要s和a,但a的生成依赖于o。
                # 在实际算法如MAPPO中,Critic通常只使用状态值函数V(s),或者需要存储额外的数据。
                # 为了示例的完整性,我们假设可以获取下一个观察,但实际代码中需要调整。
                # 此处我们跳过目标Q值的精确计算,仅演示更新流程。
                target_q_values = team_rewards + self.gamma * (1 - dones.unsqueeze(1)) * self.critic(next_states, actions)  # 注意:这是不准确的占位

            critic_loss = F.mse_loss(current_q_values, target_q_values)
            self.critic_optimizer.zero_grad()
            critic_loss.backward()
            torch.nn.utils.clip_grad_norm_(self.critic.parameters(), 0.5)
            self.critic_optimizer.step()

            # --- 更新每个智能体的策略网络 (Actor) ---
            # 计算优势函数 A(s, a) ≈ Q(s, a) - V(s)
            # 我们使用当前Critic计算的Q值作为基线,减去一个状态值函数V(s)的估计。
            # 为简化,假设V(s)由Critic网络输出(如果Critic输出的是Q,则需要另一个网络输出V,或使用Q的某种基线)。
            # 在PPO中,优势函数通常用GAE估计。这里我们用Q值减去一个常数基线做演示。
            with torch.no_grad():
                # 计算状态值函数V(s)的近似值。一个常见技巧是Critic输入全零动作。
                zero_actions = torch.zeros_like(actions)  # 占位,实际需要根据网络输入调整
                state_values = self.critic(states, zero_actions)  # [batch, 1]
                advantages = current_q_values - state_values  # [batch, 1]
                # 将优势广播到每个智能体 [batch, num_agents]
                advantages_expanded = advantages.expand(-1, self.num_agents)

            # 对每个智能体分别更新
            for i in range(self.num_agents):
                # 重新计算当前动作的log概率(因为参数可能已更新)
                # 注意:我们需要存储每个智能体的观察,但缓冲区只存了全局状态和动作。
                # 因此,在标准的PPO实现中,我们需要在存储经验时也存储观察和log概率。
                # 这里为了流程展示,我们假设可以从缓冲区获取观察(实际不行)。
                # 正确的做法是在`store_transition`时也存储每个智能体的观察和旧log概率。
                # 我们调整设计:在存储经验时,额外存储每个智能体的观察和log概率。
                # 由于代码结构限制,此处省略这部分,仅展示更新逻辑。
                # 假设我们通过某种方式获得了旧log概率 `old_log_probs_i` [batch]
                # 以及当前观察 `obs_i` [batch, obs_dim]
                # new_log_probs_i = 根据obs_i和actions[:, i]计算的新log概率
                # ratio = exp(new_log_probs_i - old_log_probs_i)
                # surr1 = ratio * advantages_expanded[:, i]
                # surr2 = torch.clamp(ratio, 1-clip_param, 1+clip_param) * advantages_expanded[:, i]
                # actor_loss = -torch.min(surr1, surr2).mean()
                # self.actor_optimizers[i].zero_grad()
                # actor_loss.backward()
                # self.actor_optimizers[i].step()
                pass  # 实际实现需要完整的数据流

        # 清空缓冲区(on-policy算法通常每轮更新后清空)
        self.buffer.buffer.clear()

上面的update函数是一个高度简化的示意,重点展示了CTDE中Critic使用全局信息更新Actor需要基于优势函数独立更新的逻辑。一个完整的MAPPO实现涉及更多细节,如:

  • 使用广义优势估计 计算优势函数。
  • 在经验缓冲区中存储观察、动作、奖励、下一个观察、状态、下一个状态、log概率、价值估计等完整数据。
  • 引入目标网络 来稳定Critic的训练。
  • 对策略更新进行多轮次优化

3.3 主训练循环

最后,我们将所有部分串联起来,形成完整的训练循环。

def train_mappo(env, agent, num_episodes=1000, max_steps_per_episode=50, batch_size=32):
    episode_rewards_history = []

    for episode in range(num_episodes):
        state = env.reset()  # 获取初始全局状态
        episode_rewards = np.zeros(env.num_agents)
        step = 0
        done = False

        while not done and step < max_steps_per_episode:
            # 1. 分布式执行:获取每个智能体的观察并选择动作
            observations = [env._get_agent_obs(i) for i in range(env.num_agents)]
            joint_actions, log_probs = agent.choose_actions(observations)

            # 2. 环境执行一步
            next_observations, rewards, done, info = env.step(joint_actions)
            next_state = info['global_state']

            # 3. 存储经验(需要存储观察、log概率等,这里简化)
            agent.store_transition(state, joint_actions, rewards, next_state, done)

            # 4. 准备下一轮
            state = next_state
            episode_rewards += np.array(rewards)
            step += 1

            # 5. 定期更新策略(例如,每收集一定步数或每回合结束后)
            if len(agent.buffer) >= batch_size:
                agent.update(batch_size=batch_size)

        # 回合结束,记录总奖励
        total_team_reward = episode_rewards.sum()
        episode_rewards_history.append(total_team_reward)
        print(f"Episode {episode+1}, Total Team Reward: {total_team_reward:.2f}, Steps: {step}")

        # 可选:每N回合评估一次策略
        if (episode + 1) % 50 == 0:
            evaluate_policy(env, agent, eval_episodes=5)

    return episode_rewards_history

def evaluate_policy(env, agent, eval_episodes=5):
    """评估当前策略在环境中的表现。"""
    total_eval_reward = 0
    for _ in range(eval_episodes):
        state = env.reset()
        done = False
        episode_reward = 0
        while not done:
            observations = [env._get_agent_obs(i) for i in range(env.num_agents)]
            joint_actions, _ = agent.choose_actions(observations)
            next_obs, rewards, done, info = env.step(joint_actions)
            episode_reward += sum(rewards)
        total_eval_reward += episode_reward
    avg_reward = total_eval_reward / eval_episodes
    print(f">>> Evaluation over {eval_episodes} episodes: Average Team Reward = {avg_reward:.2f}")
    return avg_reward

4. 高级话题与实战优化

实现了一个基础版本后,我们可以探讨一些让CTDE架构更强大、更实用的高级技术和优化点。

4.1 参数共享与个性化

在我们的示例中,每个智能体有自己独立的策略网络。但在许多合作场景中,智能体是同质的(具有相同的观察和动作空间)。这时,可以使用参数共享,即所有智能体共享同一个策略网络,这能大幅减少参数量,加速训练,并促进知识迁移。

class SharedPolicyNetwork(nn.Module):
    """所有智能体共享的策略网络。"""
    def __init__(self, obs_dim, act_dim, hidden_dim=128):
        super(SharedPolicyNetwork, self).__init__()
        # 网络结构与之前的PolicyNetwork相同
        self.fc1 = nn.Linear(obs_dim, hidden_dim)
        self.fc2 = nn.Linear(hidden_dim, hidden_dim)
        self.fc3 = nn.Linear(hidden_dim, act_dim)

    def forward(self, obs):
        # obs: [batch_size * num_agents, obs_dim] 或 [batch_size, num_agents, obs_dim]
        # 需要处理批量维度
        original_shape = obs.shape
        if obs.dim() == 3:
            batch_size, num_agents, obs_dim = original_shape
            obs = obs.view(-1, obs_dim)  # 展平为 [batch*agents, obs_dim]
        x = F.relu(self.fc1(obs))
        x = F.relu(self.fc2(x))
        logits = self.fc3(x)
        if len(original_shape) == 3:
            logits = logits.view(batch_size, num_agents, -1)  # 恢复形状
        return logits

使用参数共享后,在训练时,我们可以将一批数据中所有智能体的观察堆叠起来,一次性通过共享网络前向传播,计算效率更高。但要注意,这要求所有智能体的策略在参数上完全一致。为了在共享中保留一定的个性化,可以引入智能体身份编码,将唯一的ID(如one-hot向量)与观察拼接后输入网络,使网络能区分不同智能体。

4.2 信用分配问题

在多智能体合作任务中,团队获得一个全局奖励时,如何将功劳合理地分配给每个智能体,即信用分配,是一个关键问题。CTDE架构中的中央价值网络Q(s, a)本身就在一定程度上解决了这个问题,因为它评估的是联合动作a的价值。我们可以通过分析Q函数对单个动作a_i的依赖来理解每个智能体的贡献。

更高级的方法包括:

  • Counterfactual Baseline: 计算智能体i采取实际动作时的Q值,与假设i采取默认动作(或其他动作)时的Q值之差,作为该智能体的优势函数。
  • Q值分解: 将中央Q函数Q(s, a)分解为每个智能体贡献的和,即Q(s, a) ≈ Σ_i Q_i(s, a_i),然后分别优化每个Q_iVDNQMIX算法就是这类方法的代表。

下面是一个极度简化的QMIX思想演示,它要求每个智能体的个体Q值Q_i与全局Q值Q_tot满足单调性约束:∂Q_tot/∂Q_i ≥ 0。

# 简化的混合网络,用于将个体Q值合并为全局Q值
class MixingNetwork(nn.Module):
    def __init__(self, num_agents, state_dim, mixing_hidden_dim=32):
        super(MixingNetwork, self).__init__()
        self.num_agents = num_agents
        # 混合网络以全局状态s为条件,将个体Q值混合
        self.hyper_w1 = nn.Linear(state_dim, num_agents * mixing_hidden_dim)
        self.hyper_b1 = nn.Linear(state_dim, mixing_hidden_dim)
        self.hyper_w2 = nn.Linear(state_dim, mixing_hidden_dim)
        self.hyper_b2 = nn.Linear(state_dim, 1)

    def forward(self, individual_qs, state):
        """
        :param individual_qs: 个体Q值 [batch, num_agents]
        :param state: 全局状态 [batch, state_dim]
        :return: 全局Q值 [batch, 1]
        """
        bs = individual_qs.shape[0]
        # 第一层
        w1 = torch.abs(self.hyper_w1(state))  # 使用绝对值保证单调性
        w1 = w1.view(bs, self.num_agents, -1)  # [batch, num_agents, mixing_hidden]
        b1 = self.hyper_b1(state).unsqueeze(1)  # [batch, 1, mixing_hidden]
        # 将个体Q值扩展维度并与权重相乘
        individual_qs = individual_qs.unsqueeze(-1)  # [batch, num_agents, 1]
        hidden = F.elu(torch.bmm(w1.transpose(1,2), individual_qs) + b1)  # [batch, mixing_hidden, 1] -> [batch, mixing_hidden]
        hidden = hidden.squeeze(-1)
        # 第二层
        w2 = torch.abs(self.hyper_w2(state))  # [batch, mixing_hidden]
        b2 = self.hyper_b2(state)  # [batch, 1]
        # 计算最终Q_tot
        q_tot = torch.bmm(hidden.unsqueeze(1), w2.unsqueeze(-1)).squeeze(-1) + b2  # [batch, 1]
        return q_tot

4.3 处理部分可观察性与通信

在真正的部分可观察环境中,每个智能体的观察极其有限。除了依靠中央Critic在训练时提供全局信息外,我们还可以在智能体策略网络中引入循环神经网络(如LSTM或GRU)来处理观察序列,让智能体能够记忆历史信息,从而更好地推断全局状态。

此外,虽然CTDE在执行时不需要显式通信,但我们可以设计可学习的通信协议,让智能体在训练和/或执行时交换少量信息。例如,每个智能体可以生成一个消息向量,广播给其他智能体,其他智能体在决策时会将接收到的消息与自身观察结合。这种通信学习是MARL中一个活跃的研究方向。

class AgentWithCommunication(nn.Module):
    """带通信模块的智能体策略网络。"""
    def __init__(self, obs_dim, comm_dim, act_dim, hidden_dim=128):
        super(AgentWithCommunication, self).__init__()
        # 编码观察
        self.obs_encoder = nn.Linear(obs_dim, hidden_dim)
        # 通信编码器(生成要发送的消息)
        self.comm_encoder = nn.Linear(hidden_dim, comm_dim)
        # 通信解码器(处理接收到的消息)
        self.comm_decoder = nn.Linear(comm_dim * 2, hidden_dim)  # 假设接收自己和其他一个智能体的消息
        # 动作决策器
        self.action_head = nn.Linear(hidden_dim * 2, act_dim)  # 结合自身编码和通信信息

    def forward(self, obs, received_messages):
        # obs: 自身观察
        # received_messages: 从其他智能体接收的消息 [batch, num_other_agents * comm_dim]
        x = F.relu(self.obs_encoder(obs))
        # 生成要发送的消息
        message_to_send = torch.tanh(self.comm_encoder(x))  # 限制消息范围
        # 处理接收到的消息
        comm_info = F.relu(self.comm_decoder(received_messages))
        # 结合自身信息和通信信息做决策
        combined = torch.cat([x, comm_info], dim=-1)
        logits = self.action_head(combined)
        return logits, message_to_send

实现这样一个带通信的MARL系统更为复杂,需要定义通信拓扑(谁发给谁)、通信频率(每一步都发吗?)以及如何训练通信内容(通常通过端到端的梯度传播)。

4.4 工程实践:并行化与可扩展性

当智能体数量增多或环境复杂度增加时,训练速度会成为瓶颈。以下是一些工程优化方向:

  • 环境并行化: 使用SubprocVecEnvRay等工具同时运行多个环境实例,并行收集数据,极大提高数据吞吐量。
  • 策略网络批量前向传播: 如参数共享部分所述,将多个智能体的观察堆叠成批次进行一次性前向传播,充分利用GPU的并行计算能力。
  • 分布式训练框架: 对于超大规模MARL问题,可以考虑使用RLlibPyMARL等框架,它们内置了分布式采样、参数服务器等高级特性。

下表对比了不同规模MARL项目的实现选择:

项目规模 智能体数量 环境复杂度 推荐工具/策略 关键考量
小型实验 2-5 简单(如网格世界) 纯PyTorch + 自定义环境 快速原型,代码完全可控,便于理解算法细节。
中型研究 5-20 中等(如StarCraft微操单元) PyTorch + PettingZoo/SMAC + 环境并行 需要标准化环境接口和一定并行度,平衡灵活性与效率。
大型应用 数十上百 复杂(如交通灯控制、无人机集群) RLlib / PyMARL / Seed RL 依赖成熟的分布式训练框架,关注采样效率、可扩展性和部署便利性。

在实现自己的MARL项目时,建议从小型实验开始,确保基础CTDE流程跑通,再逐步引入参数共享、信用分配等高级模块,最后考虑并行化扩展。调试MARL系统比单智能体RL更具挑战性,因为问题的复杂性来自智能体之间的交互。善用可视化工具来观察智能体的行为、奖励曲线以及通信内容(如果有),对于诊断问题至关重要。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐