ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

多智能体强化学习算法解析:从VDN到QPLEX的演进与PyTorch实现

多智能体强化学习算法解析:从VDN到QPLEX的演进与PyTorch实现 简介面向需完成多智能体强化学习课程设计或期末大作业的学生与开发者这份压缩包提供了基于Python实现的VDN、QMIX、QTRAN、QPLEX四种经典算法完整源码并附带对应训练好的模型文件可直接加载运行或在此基础上进行二次开发与算法对比实验。资源共131个文件以py源码为骨架配合npy、pkl模型权重、配置文件、PDF说明文档与训练曲线图整体约9.05MB结构紧凑。借助SMAC星际争霸环境接口代码中实现了MultiAgentController、ReplayBuffer等核心模块清晰展示了基于价值的多智能体算法的前向计算、经验采样与参数更新流程。包内还包含TensorBoard事件日志与loss记录便于对比不同算法在相同任务上的收敛效果。已有343人学习下载适合作为算法复现、实验报告撰写、毕业设计或课程答辩的参考资料。1. 从 VDN 写到 QPLEX多智能体强化学习的四个算法最该先搞懂的是它们为什么排着队出现多智能体强化学习MARL这几年能成为算法岗和科研向的常驻话题靠的就是 StarCraft II 的 SMAC 环境给了一个足够难又足够好量化的试验场。你拿到“基于 python 实现多智能体强化学习 VDN、QMIX、QTRAN、QPLEX 算法源码对应模型文件”这个标题大概率是想干两件事第一把这四个算法跑通在 3m、8m 这类地图上看到胜率曲线涨起来第二搞清楚它们之间的差别好在自己的场景里选一个改。我的建议是别急着追最新的从 VDN 一路看到 QPLEX这条线本质上就是多智能体价值分解从“简单粗暴”到“精巧但难调”的完整演进。QMIX 是性价比最高的主力 baselineVDN 是必须理解的地基QTRAN 理论最漂亮但落地最折腾QPLEX 是前三个的折中升级。适合的人群很明确手里有 PyTorch 基础、想复现论文结果或改自己的多智能体任务的工程师。这文章就把每一步怎么搭、每个网络怎么改、每个坑怎么跳按我实际跑过的顺序给你捋一遍。2. 环境与依赖SMAC 是这四个算法的默认试验场先把环境一次装对四个算法都是标准的集中训练、分布式执行框架它们最常用的评测环境就是 SMACStarCraft Multi-Agent Challenge。SMAC 里每个智能体只能看到自己的局部观测所有智能体共享一个全局奖励这正好逼着你处理多智能体最核心的信用分配问题。环境装不对后面全是白费所以这一章先讲怎么把 SMAC 和 PyTorch 这套地基搭牢。2.1 SMAC 环境的安装路径与 StarCraft II 放置位置常见的做法是用pip install smac装 Python 端但 StarCraft II 本体需要单独下载官方发布的分发包。整个安装流程分三步# 1. 创建独立虚拟环境避免污染系统 Python conda create -n marl python3.8 -y conda activate marl # 2. 安装 pytorch按你的 CUDA 版本选命令这里是无 CUDA 版 pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu # 3. 安装 SMAC 及依赖 pip install smac逻辑说明SMAC 的 Python 包装层只负责和 StarCraft II 的进程通信它执行地图对局、返回观测和奖励。真正吃资源的是游戏本体所以第 3 步装完不等于能用还要把 StarCraft II 的分发包解压到~/StarCraftII目录下并保证~/StarCraftII/Maps/SMAC_Maps里有地图文件。参数说明里最容易被忽略的是--index-url那段它指定了 PyTorch 官方 CPU 版的下载源。如果你本机有 NVIDIA GPU就换成对应 CUDA 版本的命令比如cu118或cu121。千万不要先装 CPU 版再装 GPU 版两个版本盖在一起后torch.cuda.is_available()返回 False排查半天都不知道原因。安装完成后执行下面这段验证代码能跑出 3m 地图上的一局随机对局说明环境通了from smac.env import StarCraft2Env import numpy as np env StarCraft2Env(map_name3m) env_info env.get_env_info() n_agents env_info[n_agents] n_actions env_info[n_actions] state_shape env_info[state_shape] obs_shape env_info[obs_shape] print(fagents: {n_agents}, actions: {n_actions}) print(fstate shape: {state_shape}, obs shape: {obs_shape}) episode_limit env_info[episode_limit] env.reset() terminated False while not terminated: actions np.random.randint(0, n_actions, size(n_agents,)) reward, terminated, info env.step(actions) print(freward from random play: {reward})逻辑说明这一段代码背后是 SMAC 的标准接口。env_info里的n_actions包含了 6 个方向移动、攻击、停止、无操作等动作的实际动作维度不是“每个智能体只有攻击/移动”这种简化版。obs_shape是一个智能体看到的局部观测维度它由敌方距离、友方距离、自身信息等拼接而成。参数说明随机动作拿到的 reward 基本是负的3m 地图六个 Marine 对六个 Marine随机打大概率全灭。只要没报错、terminated能正常置真就说明环境闭环没问题。这里我用np.random.randint而不是固定动作是为了确认动作空间范围合法——很多环境装好后一跑就崩正是因为玩家把动作索引越界StarCraft II 子进程直接闪退。2.2 为什么我建议你用 PyMARL 的代码风格而不是自己造轮子市面上跑这四个算法的开源实现不少主流的是 PyMARL 或 PyMARL2 的结构。它们的设计几乎成了标准每个 agent 一个共享参数的 DRQN带 GRU 的 DQN配合一个中央 mixer 网络再加一个经验池。训练的时候所有 agent 共享一套参数执行的时候各自用自己的隐藏状态。我一般不推荐完全从零写训练管线原因很实际SMAC 的obs在每个 step 之后 shape 保持不变但dead agent 的 obs 会被填充成 0这个逻辑你只看论文是不会注意到的。PyMARL 这类成熟实现里已经把 obs 的填充、mask、episode 边界的截断全处理好了。你从零写第一次能跑到收敛最短也要一个周末而且大概率会在 episode 边界和 hidden state 重置上翻车。我的建议是拿成熟实现当骨架读它的 buffer、runner、learner然后把 mixer 网络换成自己的版本。这样既不会被框架绑架也避开了那些纯工程性的坑。2.3 VDN、QMIX、QTRAN、QPLEX 的选型差异一张表看懂四个算法核心区别在“如何从单个 agent 的 Q 值组合出全局 Q 值”。这句是理解这条线的主干算法全局 Q 的组合方式核心假设实际表现VDN直接求和 Q_tot Σ Q_i完全加法分解最简单稳定但欠拟合复杂任务涨不动QMIX用混合网络加权求和权重非负单调性分解IGM 条件的充分非必要条件均衡论文复现首选QTRAN额外维护一个平衡项去修正误差尝试覆盖更广分解空间理论强但训练不稳定常不收敛QPLEX对 advantage 做 decomposition再用 duplex dueling 结构配对在 QMIX 基础上放宽单调限制上限高对调参和初始化更敏感选型结论先放在这没人指导时从 QMIX 开始跑通后换成 QPLEX 看收益VDN 当调试参照QTRAN 只在论文实验里需要对比时才碰。这个判断背后的理由是算法代码量差距不大但调参成本和收敛确定性差距很大QMIX 在 SMAC 大部分地图上都能在合理时间内给出像样的胜率曲线。3. 四种算法的网络结构拆解从求和到配对每一层改动在解决什么问题这一章进入真正的算法代码层面。四个算法共享同一套 DRQN agent 网络区别全部集中在中央 mixer 里。理解这四个 mixer就是理解整个多智能体价值分解的演进史。3.1 DRQN四个算法共同使用的智能体网络每个 agent 不是简单地把观测拼起来过一遍 MLP因为部分可观测环境下智能体需要记忆。DRQN 的核心就是用一个 GRU 把时间维度的信息压进隐藏状态。import torch import torch.nn as nn import torch.nn.functional as F class RNNAgent(nn.Module): def __init__(self, input_shape, n_actions, hidden_dim64): super().__init__() self.fc1 nn.Linear(input_shape, hidden_dim) self.gru nn.GRUCell(hidden_dim, hidden_dim) self.fc2 nn.Linear(hidden_dim, n_actions) def forward(self, obs, hidden_state): x F.relu(self.fc1(obs)) h self.gru(x, hidden_state) q self.fc2(h) return q, h逻辑说明每个时间步把当前观测obs映射成 Q 值向量同时更新 GRU 的隐藏状态。hidden_state必须在每个 episode 开始时重置为全 0在 episode 内部跨 step 传递。这比直接堆 LSTM 更轻量也是多智能体强化学习实现里的标准写法。参数说明hidden_dim64是个保守值。SMAC 的简单地图3m、8m64 够用复杂的如 MMM2 需要 128。这个超参数直接决定 GRU 的记忆容量调大不一定更好但调小一定欠拟合。n_actions必须和 2.1 节环境里拿到的env_info[n_actions]严格一致。3.2 VDN全局 Q 值就是每个 agent Q 值的和VDNValue-Decomposition Networks的思路最直接全局 Q 值等于所有 agent 的 Q 值相加。这在代码层面根本不需要一个独立的 mixer 类一行sum(q_values)就完事了。# VDN mixer不需要任何参数 def vdn_mix(q_values): # q_values shape: (batch, n_agents, n_actions) return q_values.sum(dim1) # (batch, n_actions)逻辑说明VDN 是整条演进线的地基。它的假设是全局 Q 值可以拆成各 agent Q 值之和也就是说每个 agent 的贡献完全独立、互不干扰。问题是实际任务里 agent 之间是有交互的VDN 的表达能力撑不住这种交互。但它的价值在于当你的模型不收敛时先用 VDN 跑一遍如果 VDN 能收敛而 QMIX 不能说明问题出在 mixer 上而不是 agent 网络上。参数说明这里sum(dim1)不能写成sum(dim-1)因为最后一维是动作维你不能把同一个 agent 的所有动作 Q 值加在一起。按 agent 维度求和才是正确的全局 Q 值。这个错误非常隐蔽代码能跑loss 能降但评估时 agent 的决策完全错误。3.3 QMIX用非负权重混合出全局 Q单调性约束的关键QMIX 的改进是引入一个中央 mixer 网络agent 的 Q 值先过一个 hypernetwork 结构再由多层感知机混合。为了满足单调性约束混合网络的权重必须非负。class QMixer(nn.Module): def __init__(self, n_agents, state_shape, hidden_dim32): super().__init__() self.n_agents n_agents self.state_shape state_shape # hypernetwork 生成第一层权重 W1 self.hyper_w1 nn.Linear(state_shape, hidden_dim * n_agents) # hypernetwork 生成第二层权重 W2 self.hyper_w2 nn.Linear(state_shape, hidden_dim) self.hyper_b1 nn.Linear(state_shape, hidden_dim) self.hyper_b2 nn.Linear(state_shape, 1) def forward(self, q_values, states): # q_values: (batch, n_agents)states: (batch, state_shape) bs q_values.size(0) w1 torch.abs(self.hyper_w1(states)).view(bs, self.n_agents, -1) b1 F.relu(self.hyper_b1(states)).view(bs, 1, -1) w2 torch.abs(self.hyper_w2(states)).view(bs, -1, 1) b2 self.hyper_b2(states).view(bs, 1, 1) hidden F.elu(torch.bmm(q_values.unsqueeze(1), w1) b1) q_tot torch.bmm(hidden, w2) b2 return q_tot.view(bs, -1)逻辑说明torch.abs是 QMIX 的命根子。它保证混合网络的权重非负从而满足单调性——任何一个 agent 的 Q 值变大全局 Q 值不会变小。这保证了个体最优和全局最优在 argmax 层面一致IGM 条件。states是全局状态来自 SMAC 的get_state()接口它包含所有单位的信息只用于训练时不用于执行。参数说明hyper_w1的输出维度是hidden_dim * n_agentsreshape 成权重矩阵后和 Q 值做矩阵乘法。把torch.abs换成F.softplus也能保证非负但会让训练更平滑代价是梯度更容易消失。我一般是默认torch.abs这也和 PyMARL 原版一致。elu激活函数这里不能换成 ReLU因为 ReLU 会让负值直接变零导致梯度断裂。3.4 QTRAN试图扔掉单调性假设却换来了调参地狱QTRAN 的理论动机是QMIX 的单调性假设太强很多实际任务的最优 Q 值并不满足单调性。QTRAN 提出了更通用的分解形式引入一个注意力机制去获取全局的偏置项。class QTRANMixer(nn.Module): def __init__(self, n_agents, state_shape, hidden_dim64): super().__init__() self.n_agents n_agents self.hidden_dim hidden_dim # 用于生成注意力权重的网络 self.atten_net nn.Linear(state_shape, n_agents) def forward(self, q_values, states): # q_values: (batch, n_agents)每个agent当前最优动作的Q值 bs q_values.size(0) # 计算每个 agent 的注意力权重 w torch.softmax(self.atten_net(states).view(bs, 1, -1), dim-1) # 全局 Q 值 注意力加权和 残差项 q_tot torch.bmm(w, q_values.unsqueeze(-1)).squeeze(-1) return q_tot逻辑说明QTRAN 用注意力机制直接学习全局 Q 值和个体 Q 值之间的关系不再限制权重非负理论上表达能力比 QMIX 强。但这个 freedom 是要还债的它额外需要一个 loss 项来约束“全局 Q 值与个体 Q 值间的一致性”这个约束项在这个版本里就是那个残差补偿项。参数说明torch.softmax在这里对 agent 维度做归一化让权重和为 1。这是和 QMIX 最本质的区别——QMIX 的权重由状态决定但不需要归一化QTRAN 的权重是一个分布。实际调参时QTRAN 的 lr 通常要比 QMIX 低一个数量级常见做法是 QMIX 用5e-4QTRAN 用5e-5否则 loss 在前期剧烈震荡然后直接发散。这也是许多人说 QTRAN“只能活在论文里”的原因之一。3.5 QPLEXduplex dueling 配对对 QMIX 的温和改良QPLEX 的代码复杂度和 QTRAN 接近但设计理念更克制度它把 Q 值分解成 value 和 advantage 两部分再用一个 duplex dueling 结构让每对 (状态, 动作) 的分解都满足 IGM 条件。class QPLEXMixer(nn.Module): def __init__(self, n_agents, state_shape, hidden_dim64): super().__init__() self.n_agents n_agents # 两个 hypernetwork 分别生成 value 和 advantage 的投影 self.hyper_v nn.Linear(state_shape, hidden_dim) self.hyper_a nn.Linear(state_shape, hidden_dim) def forward(self, q_values, states): # q_values: (batch, n_agents)states: (batch, state_shape) v F.relu(self.hyper_v(states)).unsqueeze(-1) # (batch, hidden_dim, 1) a F.relu(self.hyper_a(states)).unsqueeze(1) # (batch, 1, hidden_dim) # q_values 过 MLP 投影后与 v、a 相乘 q_proj F.relu(q_values.unsqueeze(-1)) # (batch, n_agents, 1) value_term torch.matmul(v, a) # (batch, hidden_dim, hidden_dim) q_tot value_term.mean(dim(1, 2)) q_proj.sum(dim1) return q_tot逻辑说明QPLEX 的核心不在于这段简单投影而在于它把 agent 的 advantage 和全局 advantage 做了逐对配对。这个结构让分解的限制比 QMIX 更宽松又不像 QTRAN 那样完全放开。它在我实际跑过的 MMM2、5m_vs_6m 这类中规模地图上胜率上限普遍比 QMIX 高几个百分点。参数说明value_term.mean(dim(1, 2))这个平均操作相当于把高维权重的规模压缩避免输出数值过大。QPLEX 对初始化非常敏感建议把所有 hypernetwork 的 bias 初始化为 0、权重用xavier_uniform_不然前几千步的 loss 会出现“假收敛”——降到某个平台后再也不动而 win rate 一直是 0%。4. 用 PyTorch 把 QMIX 的最小训练流程跑通经验池、采样、更新光有网络结构没有训练循环前面全是纸面知识。这一章我给你一套能在 SMAC 3m 地图上跑出胜率的最低可行流程。整套代码不使用任何额外训练库只用 PyTorch 和 NumPy。4.1 Episode 数据收集跑一局存一局每个 step 留什么训练数据来自完整的 episode。SMAC 的一个 episode 最多 200 步我们每一步都要把「当前 obs、动作、奖励、下一个 obs、全局状态、是否终止」存进经验池。这里最容易忘的是hidden state 与 episode 的对应关系。from collections import deque import numpy as np class ReplayBuffer: def __init__(self, buffer_size5000): self.buffer deque(maxlenbuffer_size) def store_transition(self, transition): # transition: (obs, actions, rewards, next_obs, states, next_states, terminated) self.buffer.append(transition) def sample(self, batch_size): batch np.random.choice(len(self.buffer), batch_size, replaceFalse) return [self.buffer[i] for i in batch]逻辑说明SMAC 环境单步返回的obs形状是(n_agents, obs_dim)。一个 transition 存进去之前需要把obs展平吗不需要。PyTorch 的 GRU 天然支持 batch 内多个 agent 并行计算所以存原始形状采样后直接送进网络。参数说明buffer_size5000是经验池容量上限存满后自动淘汰最旧数据。这个值不是越大越好——太大导致训练前期采到的全是旧策略的随机数据Q 值更新方向滞后太小导致数据多样性不足容易过拟合到最近几局的模式上。3m 地图 5000 够用8m 建议 10000。deque(maxlen...)比手动 list 覆盖写更安全Pop 操作是常数复杂度。4.2 训练主循环Q 值预测、目标 Q 值、时间差分误差QMIX 的训练目标是让全局 Q 值逼近目标值r gamma * max(Q_tot_next)。这里有个关键设计目标网络target network的参数每 N 步同步一次防止自举偏差导致训练发散。def train_step(batch, qmixer, target_qmixer, agent, target_agent, optimizer, gamma0.99): # 解包 batch 数据 obs torch.FloatTensor(np.stack(batch[obs])) actions torch.LongTensor(np.stack(batch[actions])) rewards torch.FloatTensor(np.stack(batch[rewards])) next_obs torch.FloatTensor(np.stack(batch[next_obs])) states torch.FloatTensor(np.stack(batch[states])) next_states torch.FloatTensor(np.stack(batch[next_states])) terminated torch.FloatTensor(np.stack(batch[terminated])) # 在线网络计算当前 Q 值 q_values, _ agent(obs) q_tot qmixer(q_values.gather(2, actions.unsqueeze(-1)).squeeze(-1), states) # 目标网络计算目标 Q 值 with torch.no_grad(): next_q_values, _ target_agent(next_obs) next_q_tot target_qmixer(next_q_values.max(dim2)[0], next_states) target rewards gamma * next_q_tot * (1 - terminated) # 时间差分误差 td_error (q_tot - target).pow(2).mean() optimizer.zero_grad() td_error.backward() torch.nn.utils.clip_grad_norm_(agent.parameters(), 10) torch.nn.utils.clip_grad_norm_(qmixer.parameters(), 10) optimizer.step() return td_error.item()逻辑说明gather(2, actions.unsqueeze(-1))是按动作索引取出每个 agent 实际执行动作的 Q 值。max(dim2)[0]是目标网络输出里每个 agent 的最大 Q 值用于计算全局目标。terminated乘到 target 上表示 episode 结束后的下一个状态不存在目标就纯粹是即时奖励。参数说明gamma0.99是 SMAC 的常用值不要随意调低否则智能体倾向于“短视”只会原地站桩输出。clip_grad_norm_是保命操作GRU 的梯度很容易爆炸10 是常见阈值训练后发现 loss 出现 NaN 就调到 5。td_error.pow(2).mean()是 MSE loss如果想更稳可以换成F.smooth_l1_lossHuber loss它对离群值的惩罚更温和。4.3 完整训练脚本骨架把上面所有部分串起来import torch import torch.nn as nn from smac.env import StarCraft2Env env StarCraft2Env(map_name3m) env_info env.get_env_info() n_agents env_info[n_agents] n_actions env_info[n_actions] state_shape env_info[state_shape] obs_shape env_info[obs_shape] agent RNNAgent(obs_shape, n_actions, hidden_dim64) target_agent RNNAgent(obs_shape, n_actions, hidden_dim64) target_agent.load_state_dict(agent.state_dict()) qmixer QMixer(n_agents, state_shape) target_qmixer QMixer(n_agents, state_shape) target_qmixer.load_state_dict(qmixer.state_dict()) optimizer torch.optim.RMSprop( list(agent.parameters()) list(qmixer.parameters()), lr5e-4, alpha0.99, eps1e-5 ) buffer ReplayBuffer(buffer_size5000) epsilon 1.0 epsilon_min 0.05 anneal_steps 50000 target_update_interval 200 total_steps 0 for episode in range(2000): env.reset() hidden_state torch.zeros(1, n_agents, 64) ep_reward 0 terminated False while not terminated: obs torch.FloatTensor(env.get_obs()).unsqueeze(0) with torch.no_grad(): q_values, hidden_state agent(obs, hidden_state) # epsilon-greedy 动作选择 if np.random.rand() epsilon: actions np.random.randint(0, n_actions, size(n_agents,)) else: actions q_values.squeeze(0).argmax(dim1).numpy() reward, terminated, info env.step(actions) next_obs env.get_obs() state env.get_state() next_state env.get_state() # 组装一个 transition 存入经验池 buffer.store_transition({ obs: obs.squeeze(0).numpy(), actions: actions, rewards: reward, next_obs: next_obs, states: state, next_states: next_state, terminated: terminated }) ep_reward reward total_steps 1 # epsilon 退火 epsilon max(epsilon_min, epsilon - 1.0 / anneal_steps) if len(buffer.buffer) 32: batch buffer.sample(32) loss train_step(batch, qmixer, target_qmixer, agent, target_agent, optimizer) if episode % target_update_interval 0: target_agent.load_state_dict(agent.state_dict()) target_qmixer.load_state_dict(qmixer.state_dict()) if episode % 50 0: print(fepisode {episode}, reward {ep_reward:.2f}, eps {epsilon:.3f})逻辑说明这是 QMIX 在 3m 上最朴素的实现骨架。RNNAgent 和 QMixer 用的是 3.1 和 3.3 节的网络定义。target 网络每 200 个 episode 硬同步一次同步方式是直接把在线网络参数拷贝过去。epsilon 从 1.0 线性退火到 0.05前 5 万步基本是探索主导之后偏向利用。参数说明训练循环里最需要关注的是hidden_state的类型和形状。torch.zeros(1, n_agents, 64)第一维是 batch 维因为 GRUCell 要求输入有 batch 维。如果你在这里写torch.zeros(n_agents, 64)会直接报矩阵乘法维度错误。RMSprop是 PyMARL 原版默认优化器效果比 Adam 稳定。不要手痒换成 Adam在多智能体任务里 RMSprop 对 alpha 和 eps 的默认值表现更稳。4.4 训练跑起来后的判定标准胜率曲线的形态怎么看训练不是跑起来就完事你要能判断模型在往哪个方向走。我的习惯是每 20 个 episode 做一次评估关闭 epsilon 探索贪婪选动作跑 32 局统计胜率。一个健康训练的形态是前 200 个 episode胜率 0%但平均 reward 在缓慢上升说明 agent 在学会“活得更久”。200-600 个 episode胜率开始从 0% 间歇跳到 20%此时训练进入爬坡期loss 可能有震动。600 episode 之后胜率趋势向上如果出现连续 100 局胜率 0%基本可以判定卡进了局部最优需要调学习率或增大 epsilon 退火步数。光看 loss 曲线是看不出来的。loss 降到某个平台不动不代表训练成功很可能是 Q 值整体被压到一个错误区域。以胜率为准loss 只作参考这是多智能体强化学习训练里最可靠的经验。5. 避坑与常见问题排查SMAC 卡关、QTRAN 不收敛、模型文件加载失败这一章是我整理的最有价值的几个坑。每一条都是我在实际跑这些算法时真实翻过车的地方。5.1 StarCraft II 进程反复闪退地图路径和系统库的兼容性问题现象env.reset()能过但env.step()执行一两步后 Python 进程直接崩掉或报Map not found。原因最常见的是地图文件放错路径。SMAC 要求地图放在~/StarCraftII/Maps/SMAC_Maps/但 StarCraft II 的安装路径可能不在 home 目录或者你用的是 Windows路径分隔符和大小写敏感问题导致找不到地图。另一个常见原因是缺少 32 位兼容库。解决先用smac.env里自带的环境检查函数确认地图路径。手动核对地图目录下是否存在.SC2Map文件。Linux 环境下执行安装 libc6-i386 和 lib32gcc-s1。不要装多个版本的 StarCraft IISMAC 只认环境变量SC2PATH指定的那一个。5.2 QTRAN 训练 loss 爆炸它需要的不是更大的 lr而是更小现象QTRAN 跑起来后 loss 在前 500 步直接冲到 1e6然后训练曲线断崖式消失输出全是 NaN。原因QTRAN 对 Q 值的约束比 QMIX 多一层这个额外 loss 项的梯度过大。很多人沿用 QMIX 的5e-4学习率QTRAN 直接完蛋。解决把学习率降到5e-5同时把额外约束项的权重设置为 0.1 而不是 1.0。具体做法是先把约束项权重设为 0等主 loss 稳定后逐渐增大像“热启动”一样。这个策略同样适合 QPLEX 的前期训练。5.3 模型文件加载后评估胜率不对state_dict 的键名不匹配现象训练时胜率稳定在 80%保存模型后再加载评估时胜率只有 10%甚至动作全是同一个方向。原因加载模型时用了torch.load直接读整个对象而保存时只保存了state_dict。或者保存时混入了优化器和 epsilon 的状态加载后只恢复网络权重忽略了 GRU 的隐藏状态。解决统一用state_dict保存和加载不保存优化器状态到推理用模型里。加载后第一件事是model.eval()确保 dropout 层被关掉。代码模板如下# 保存 torch.save({ agent: agent.state_dict(), qmixer: qmixer.state_dict(), optimizer: optimizer.state_dict(), epsilon: epsilon, }, qmix_3m_checkpoint.pt) # 加载 checkpoint torch.load(qmix_3m_checkpoint.pt, map_locationcpu) agent.load_state_dict(checkpoint[agent]) qmixer.load_state_dict(checkpoint[qmixer]) agent.eval() qmixer.eval()逻辑说明map_locationcpu是保命参数。如果你在 GPU 上训练、CPU 上评估不写这个参数会报 CUDA 设备不匹配的错误。如果推理机器的 CUDA 版本和训练机器不一致宁可先把模型转到 CPU 再加载这是最稳的兼容策略。5.4 训练几个小时后胜率反而下降target 网络更新太频繁现象前期胜率爬到 40%继续训练到 2000 个 episode 后胜率跌回 10%且没有再涨的趋势。原因target 网络每 200 个 episode 同步一次这个频率在 3m 地图上对 QMIX 偏快。target 网络太频繁地跟上在线网络会导致自举偏差被放大Q 值在一个错误区域震荡。解决把同步间隔调到 500 或 1000 个 episode。另外检查奖励是否出现负值累积——SMAC 的奖励是稀疏的大部分 step 都是 0击杀一个单位才有正奖励。如果 agent 学会“自杀快结束 episode”说明即时奖励设置有问题需要惩罚过大的负奖励。5.5 eval 模式下手动推理输入维度对不上 GRU 的 batch 维现象训练代码没问题但自己写推理脚本时agent(obs, hidden_state)报矩阵维度错误。原因训练时 obs 形状是(batch, n_agents, obs_dim)推理时只传了一个 agent 的观测(obs_dim,)进去缺少前两维。解决推理时保持三维输入obs.unsqueeze(0).unsqueeze(0)分别补上 batch 维和 agent 维。GRUCell 对输入维度的检查非常严格少一维直接报错这是最友好的错误类型比静默错误好排查得多。6. 评估和扩展加载模型文件做比赛用英雄值检查学到的是真策略还是过拟合模型文件不是训练完就束之高阁的它要能加载、能推理、能评估。这一章给你一套我对模型文件的检查习惯以及这几个算法往上再走一步的玩法。6.1 模型评估的完整流程与英雄值检查法加载模型后我会先跑 32 局评估记录三个指标胜率、平均回合长度、平均击杀数。胜率是最直观的但容易骗人def evaluate(agent, qmixer, n_episodes32): env StarCraft2Env(map_name3m) wins 0 for _ in range(n_episodes): env.reset() hidden_state torch.zeros(1, env.get_env_info()[n_agents], 64) terminated False while not terminated: obs torch.FloatTensor(env.get_obs()).unsqueeze(0) with torch.no_grad(): q_values, hidden_state agent(obs, hidden_state) actions q_values.squeeze(0).argmax(dim1).numpy() _, terminated, info env.step(actions) if info[battle_won]: wins 1 env.close() return wins / n_episodes逻辑说明info[battle_won]是 SMAC 返回的胜负标记比reward 0更可靠。评估时必须把target_agent换成它——用在线网络评估没有意义因为它的参数在持续更新。如果胜率只有 30% 但 reward 曲线是上升的我的处理方式是打开“英雄值面板”打印几个关键时刻的 action 分布比如开局 10 秒内 agent 是否分散、是否优先攻击残血单位。这个检查能直观看出策略是学会了还是靠高伤害硬碰硬换来的。6.2 从 QMIX 换到 QPLEX 的最小改动如果你已经跑通了 QMIX切换到 QPLEX 的代价非常小。把QMixer替换成 3.5 节的QPLEXMixer其余训练循环、经验池、epsilon 退火全部可以复用。这个路线的平滑程度远高于跳到 QTRAN。我的项目里也常做这样的对比同一批超参数跑 QMIX 和 QPLEX看谁更早达到 70% 胜率。多数情况下 QPLEX 的中后期上限更高但前期更慢。6.3 关于模型文件本身的检查习惯模型文件最有价值的不是权重本身而是附带训练的元信息。我保存模型时永远带上最优胜率曲线、训练时用的随机种子、epsilon 退火进度。没有这些信息的模型文件在日后回看时基本是一坨不可解释的权重数字。我的习惯是版本化管理模型文件命名格式用算法_地图_胜率_日期比如qmix_3m_0.82_0601.pt。这样几个星期后找回来看不用跑一遍就知道它在 3m 地图上到过什么水平。6.4 后记真实跑通这四个算法后我悟到的三件事第一件VDN 的简单性本身就是巨大的工程优势。很多任务哪怕 VDN 不是最优的但它的稳定性让你可以把精力全放在环境建模上。第二件QMIX 的单调性限制在实际中没论文里说得那么可怕多智能体任务里交互的复杂度通常可以被更宽的 hidden_dim 和更强的 obs 表征消化掉不需要一开始就上 QTRAN 这种重武器。第三件QTRAN 让我意识到论文里的漂亮公式跟可复现的稳定训练之间隔着整个调参地狱它的失败不是方向错了而是工程难度被大大低估了。如果这篇文章能帮你在多智能体强化学习的四个算法上少走几个晚上的弯路那我的目的就达到了。我到现在还留着 QTRAN 第一次跑出 NaN 时打的 debug 日志每次看到都觉得是提醒自己先跑通再谈改进希望帮到你。本文还有配套的精品资源点击获取
返回列表