ARTICLE DETAIL

资讯详情

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

QMIX算法详解:多智能体强化学习中的单调值分解与CTDE实践

QMIX算法详解:多智能体强化学习中的单调值分解与CTDE实践 做多智能体强化学习这件事QMIX是一个绕不开的名字。它出自ICML 2018的论文《QMIX: Monotonic Value Function Factorisation for Deep Multi-Agent Reinforcement Learning》之后几年已经成为Value-Based多智能体算法里被复现最多、引用最广的baseline之一。它的核心作用听起来并不复杂在“集中训练、分布式执行”的框架下用一个带单调性约束的混合网络把多个智能体的局部Q值合并成全局Q值从而解决全局奖励如何分配到单智能体决策的问题。如果你正在做机器人调度、游戏AI、交通信号控制这类协作场景或者刚进入MARL领域想找一个能快速上手的切入点QMIX都很值得系统学一遍。下面我会把它的设计动机、核心原理、代码实现和调参避坑一次讲清楚。1. 多智能体问题到底难在哪QMIX想解决什么1.1 为什么不能直接把DQN套在每个智能体上很多同学刚接触多智能体强化学习时会想DQN已经在单智能体环境里这么成熟了给每个智能体单独装一个DQN不就行了吗答案是可以跑但很难收敛甚至经常出现“每个智能体都在进步整体reward却越来越烂”的诡异现象。原因主要有三点。第一点是环境非平稳。单智能体环境里状态转移只取决于智能体自身的动作和环境动态但在多智能体环境里某个智能体的下一帧观测会受其他智能体动作影响而其他智能体又会在训练中不断更新策略。从一个智能体的视角看整个环境就像在“变规则”经验回放里的老数据往往已经对应了旧队友策略拿来训练自然会震荡。DQN假设经验数据近似独立同分布这个前提在多智能体下被破坏得很厉害。第二点是联合动作空间爆炸。假设每个智能体只有5个动作10个智能体的联合动作空间就是5的10次方将近一千万种。直接用联合Q网络去枚举全局动作做max操作计算量会指数级增长根本没法落地。这也是Value-Based多智能体算法最核心的痛点如何在不枚举全局动作的前提下找到最优联合动作。第三点是信用分配问题。环境通常只返回一条全局奖励比如一队机器人搬箱子成功了一次系统给一个总分但每个智能体到底为这个成功贡献了多少算法是无从直接得知的。如果不做任何分解单个智能体很难判断自己某一步是好是坏梯度信号会淹没在噪声里。把这三点放在一起再看QMIX的设计就很顺了用局部Q函数回避联合动作枚举用带单调性约束的因子化保证全局和局部的最优动作一致用混合网络在训练阶段做全局信用分配同时让执行阶段不依赖全局信息。1.2 VDN先做的“求和分解”差在哪在QMIX之前VDNValue Decomposition Networks是分解路线的代表。VDN的想法非常直接全局Q_tot等于所有智能体局部Q_i加起来公式写出来就是Q_tot ΣQ_i(τ_i, u_i)。这样做的好处相当明显线性加法天然满足可分解性训练简单而且由于加法对每个分量都单调递增全局argmax一定等于各个局部argmax的组合。但VDN的问题也恰恰出在这个“简单”上。真实协作任务里的Q_tot并不是智能体价值的简单相加智能体之间交互的影响、全局状态对联合价值的影响都是非线性的。我经常举一个例子两个机器人一起抬一件重物单独一个机器人根本扛不动两个人的力量叠加会产生“112”的协作效应这种价值函数很难用线性求和表示。VDN在SMAC的不少中等难度地图上表现有限正是因为它没有建模这种非线性协作关系。QMIX的目标就是在保留VDN“可分解、易训练”优点的前提下把表达能力从线性提升到非线性同时仍然保证一个关键性质全局Q_tot的最大值与所有局部Q_i的最大值保持一致。这个性质在论文里叫Individual-Global-Max也就是IGM条件。1.3 CTDE与IGM理解QMIX价值主张的两把钥匙先展开说CTDE全称是Centralized Training with Decentralized Execution集中训练分布式执行。这也是QMIX运行的总体框架。训练的时候算法可以使用全局状态、全局奖励、所有智能体的观测信息很完整学起来准确执行的时候每个智能体只能根据自己局部观测历史和局部Q值做贪心选择不需要和其他智能体实时通信。这个特性在实际工程里价值很大因为通信延迟或带宽限制常常是系统的最大瓶颈。IGM条件用公式表达是argmax_ Q_tot(τ, ) (argmax_{u_1} Q_1(τ_1, u_1), argmax_{u_2} Q_2(τ_2, u_2), ..., argmax_{u_n} Q_n(τ_n, u_n))翻译成大白话就是先对每个智能体单独求出它最好的动作然后把这些局部最优动作拼在一起得到的联合动作同样也是全局Q_tot的最优联合动作。这样一来执行阶段每个智能体只需要算自己的Q_i并取max不需要看别人在做什么不需要知道全局状态也能保证整个团队的动作组合是最优的。QMIX通过让混合网络的权重非负来保证Q_tot对每个Q_i单调不减从而近似满足IGM这就是它在非线性分解空间里依然能保持一致性的根本原因。VDN是线性求和天然满足IGMQMIX则把满足IGM的范围扩展到了更广泛的单调非线性函数族。从VDN到QMIX再到后来QTRAN、WQMIX、QPLEX这些变体本质上都是同一个问题如何在更复杂的函数空间里满足或者放宽IGM条件。搞清楚这条线索后面读任何MARL论文都会轻松很多。2. QMIX核心原理拆解网络结构、单调约束与训练机制2.1 智能体网络、混合网络、超网络各司其职QMIX整体结构分为三大块智能体网络Agent Network、混合网络Mixing Network、超网络Hypernetwork。智能体网络负责把每个智能体的局部观测历史映射成局部Q值。它通常是一个DRQN也就是带RNN结构的DQN核心组件是GRU或LSTM因为单帧观测往往无法提供完整的马尔可夫状态必须靠循环结构记住历史信息。网络输出的是每个智能体对所有离散动作的Q值估计。混合网络负责把所有局部Q_i合并成全局Q_tot。它接收的是当前状态下每个智能体的Q_i以及全局状态s最后的输出就是Q_tot。混合网络每一层的权重不是直接训练出来的而是由超网络根据全局状态s生成。超网络接收全局状态作为输入输出混合网络每一层的权重矩阵和偏置。它的存在让“全局状态如何影响价值合并规则”这件事变成可训练的。由于全局状态只在训练时可用执行阶段根本用不到超网络这同时契合了CTDE的要求。这三个模块合在一起就是QMIX的全貌。它最大程度利用了全局状态来做训练阶段的集中学习又在网络结构上约束了单调性从而保证分布式执行阶段每个智能体各自的最优动作可以无冲突地组合成全局最优动作。2.2 单调性约束与非负权重一切的关键要理解QMIX的单调性约束可以把混合网络想象成一个评分器Q_tot如果对某个智能体的Q_i单调递增意味着其他条件不变时某个智能体对某个动作的局部评分越高全局评分只允许变高或持平不允许变低。数学上写出来就是∂Q_tot / ∂Q_i ≥ 0对于任意智能体i都成立这个条件怎么保证最简单的办法是把混合网络所有中间层的权重都限制为非负同时激活函数选用单调不减函数比如ReLU。根据链式法则非负权重在单调不减的激活函数下逐层复合最终Q_tot对Q_i的偏导必然非负。这个约束看起来很简单但它在实际中换来了一个很重要的性质团队不能通过“牺牲某个智能体的利益”来让自己的全局Q更大。换句话说一个动作对团队好不好至少不能和它对这个智能体好不好相矛盾。很多协作任务恰恰满足这个性质团队的整体奖励归根结底来自个体贡献的加总只是贡献形式比较复杂而已。这里要特别说明一点非负权重是实现单调性的一种充分条件不是必要条件。也就是说即便混合网络权重非负仍然存在一些满足IGM但本身不单调的值函数无法被表示反过来某些非单调但可分解的值函数也无法用QMIX表示。这是QMIX表达能力的理论上限也是后面WQMIX、QPLEX等变体试图突破的点。我在第4章会再展开。2.3 为什么全局状态要走超网络而不是直接输入全局状态在训练阶段包含了很丰富的信息比如SMAC里的所有单位位置、血量、技能冷却等。QMIX把全局状态交给超网络由超网络生成混合网络的权重这样做有一个非常重要的原因不能破坏单调性。如果直接把全局状态拼到混合网络的输入里它就会作为普通特征参与加权求和一旦它在某个维度上有负权重Q_tot对某个Q_i的偏导就可能变负单调性就破了。但让全局状态只出现在权重生成器里情况就不一样了。权重本身生成多少、正负如何都先经过非负激活全局状态对每个Q_i的影响就变成了“调节合并系数的上下文”不会直接改变偏导符号。用生活化的话说Q_i决定的是“某个智能体自己评估这个动作有多好”Q_tot决定的是“这个团队在这个全局局面下整体有多好”。混合网络要解决的是当我知道全局局面时应该怎么把个体的好翻译成团队的好。这个翻译规则理应随全局状态变化所以超网络动态生成权重比固定一个权重矩阵合理得多。2.4 为什么智能体网络非要用循环结构智能体网络的输入是每个智能体自己的局部观测中间经过一层DRQN。这里很多人会问为什么就不能用全连接网络因为在部分可观测环境下单帧观测往往不够。比如SMAC里一个单位能看到视野内的敌人但看不到视野外的信息要判断自己是不是被包夹必须结合过去几步的观测变化。观测历史才是真正的状态RNN就是用来编码这个历史的。具体实现中GRU比LSTM更常见原因很实际GRU参数更少训练更快在短序列任务上和LSTM差距不大而多智能体场景本身就要训练多个网络参数效率很关键。每个智能体网络接收观测序列输出当前时刻隐含状态h_i再接一个全连接层得到每个动作的Q_i值。要注意Q_i的维度通常是(batch_size, n_agents, n_actions)这个维度需要一直保留到混合网络因为混合网络要沿着智能体维度做加权。2.5 目标值与TD Loss训练机制的关键点QMIX的训练沿用了DQN的离线训练框架但细节上比单智能体更讲究。目标Q_tot写成y r γ * max_ Q_tot⁻(τ, )这里的Q_tot⁻表示target网络。有意思的是借助IGM条件和单调性约束max_ Q_tot⁻(τ, )并不需要真的去枚举联合动作而是可以拆成每个智能体各自取max之后的Q_i值再经过target混合网络聚合。这正是QMIX能在指数级联合动作空间中高效训练的根本原因。Loss就是标准的TD误差平方常见写法是L(θ) E[(Q_tot(s, ; θ) - y)²]我在实现时还会注意一个细节Q_tot的维度要从(batch, 1)调整成(batch,)否则和reward形状对齐时容易产生隐式的广播错误。PyMARL的代码里有一个很不起眼的squeeze操作作用就在这里很多人复现掉点就是栽在这个小地方。3. 动手实现环境、代码与训练全流程3.1 环境选型SMAC还是自己造一个简单环境最常见的QMIX实验环境是SMACStarCraft Multi-Agent Challenge它基于星际争霸II做微操对战支持部分可观测、异质智能体、局部视野、复杂奖励和真实场景比较接近。社区里已经有很多baseline结果可以对比这是它的最大优势。缺点是安装略麻烦需要下载SC2客户端和地图包对机器性能也有一定要求。如果你只是想快速验证算法逻辑建议先用简单网格环境。比如5x5网格里两个智能体协作把球推到目标位置每个智能体动作是上下左右。这种环境几分钟就能跑起来能直观看到全局reward涨不涨非常适合验证分解逻辑。我自己建议的路径是先在网格环境把代码逻辑调通再去SMAC的简单地图比如2s3z跑50个episode确认能正常拿到observation、reward、done。之后再上更高难度地图。这样不管是环境接口问题还是算法问题都能分得很清楚排查起来快很多。3.2 一个清晰的代码结构如果从零写QMIX仓库我建议这样组织qmix-project/ ├── config/ │ ├── env.yaml # 环境参数 │ └── algorithm.yaml # 算法超参数 ├── envs/ │ ├── smac_wrapper.py # 环境封装 ├── agents/ │ ├── rnn_agent.py # 智能体网络 │ ├── mixing_net.py # 混合网络与超网络 ├── runners/ │ ├── episode_runner.py # 采样收集 ├── learners/ │ ├── qmix_learner.py # 训练主逻辑 ├── utils/ │ └── buffer.py # 经验回放 └── main.py这个结构分环境、智能体、运行器、学习器几层和PyMARL等开源baseline的组织方式基本一致。我的建议是参考PyMARL的代码但一定不要直接复制跑一下就算完最好自己把核心模块重写一遍。多智能体代码的难点不在某个网络有多复杂而在各种维度怎么对齐、loss怎么流这些只有手写一遍才能真正搞清楚。3.3 核心网络实现与维度注意事项智能体网络部分用PyTorch写出来非常短import torch import torch.nn as nn class RNNAgent(nn.Module): def __init__(self, input_dim, action_dim, hidden_dim64): super().__init__() self.fc1 nn.Linear(input_dim, hidden_dim) self.gru nn.GRUCell(hidden_dim, hidden_dim) self.fc2 nn.Linear(hidden_dim, action_dim) def forward(self, obs, hidden_state): x torch.relu(self.fc1(obs)) h self.gru(x, hidden_state) q self.fc2(h) return q, h混合网络的实现稍复杂一点class MixingNet(nn.Module): def __init__(self, n_agents, state_dim, embed_dim32): super().__init__() self.n_agents n_agents self.state_dim state_dim self.embed_dim embed_dim self.hyper_w1 nn.Linear(state_dim, embed_dim * n_agents) self.hyper_w2 nn.Linear(state_dim, embed_dim) self.hyper_b1 nn.Linear(state_dim, embed_dim) self.hyper_b2 nn.Linear(state_dim, 1) def forward(self, q_values, states): batch q_values.size(0) w1 torch.abs(self.hyper_w1(states)) w1 w1.view(batch, self.n_agents, self.embed_dim) b1 torch.relu(self.hyper_b1(states)).view(batch, 1, self.embed_dim) q_values q_values.unsqueeze(1) hidden torch.relu(torch.bmm(q_values, w1) b1) w2 torch.abs(self.hyper_w2(states)).view(batch, self.embed_dim, 1) b2 self.hyper_b2(states).view(batch, 1, 1) q_tot torch.bmm(hidden, w2) b2 return q_tot.view(batch, 1)这份代码里有几个细节值得反复强调。第一w1和w2在计算前必须用torch.abs取绝对值这是单调性的来源。第二偏置不需要非负因为偏置对Q_i求偏导是0不影响单调性。第三b1用relu而不是abs实际对单调性无影响只是让中间层进入激活函数的工作区间。第四q_values在进入混合网络前要unsqueeze成(batch, 1, n_agents)因为后面要用bmm做批量矩阵乘法。这些维度如果看不对模型很容易“能跑但永远不收敛”。3.4 训练循环与Loss流向QMIX的训练循环分为四步。第一步用当前策略与环境交互收集完整episode存入经验回放池。第二步从池中采样一个batch的序列序列长度常见在10到25之间。第三步用target网络计算目标Q_tot公式是r γ * max_ Q_tot⁻其中max部分通过各智能体target网络取局部max再聚合得到。第四步用online网络计算当前Q_tot和target做smooth L1 loss或MSE loss反向传播更新参数。这里特别容易错的地方是计算target时不能用当前步的Q_i去取max而必须把下一个观测和各自的hidden state传给target网络重新算Q_i。很多人复现时图省事直接在同一个时间步上算了两步的Q值结果target被严重低估学习效果大打折扣。另一个容易踩的坑是target网络更新频率。QMIX默认是硬更新每隔一定episodes把online网络参数直接复制到target网络。更新太频繁target追着online跑失去稳定目标的意义更新太慢训练初期可能震荡。我一般从200个episode起步按任务复杂度上下浮动。3.5 超参数配置参考我给出在SMAC 2s3z地图上比较稳的一套配置可以作为初始参考。参数取值说明学习率0.0005Adam优化器batch_size32每次训练的序列数回放池大小5000以完整episode为单位训练序列长度10-25按episode长度截断目标网络更新间隔200 episodes硬更新折扣因子gamma0.99标准设置epsilon初值1.0随episode衰减至0.05智能体hidden_dim64GRU隐藏层维度mix嵌入维度32混合网络中间层维度探索策略上我不建议一开始就把epsilon降得很低。多智能体探索比单智能体更困难因为某个智能体的一个随机动作可能破坏整个团队策略的连续性经验池里如果没有足够的成功样本分解学习就失去了依据。前1000个episode建议把epsilon保持在0.3以上之后再逐步衰减到0.05。这个看起来反直觉的配置在实际训练里经常能救回一个死活不涨的曲线。4. 常见问题、调试技巧与算法变体4.1 训练不收敛先按顺序排查三件事遇到reward一直不涨我的排查顺序非常固定这也是一套能省很多时间的方法论。第一步检查维度。打印q_values、q_tot、reward、done的形状。确认q_values是(batch, n_agents, n_actions)q_tot是(batch, 1)reward是(batch,)。任何形状对不上都会在loss计算时产生隐性的广播错误这类错误通常不会直接报错只会让Q值越学越偏。我在复现QMIX时一大半“玄学不收敛”最后都定位到维度上。第二步看TD loss曲线。如果loss一直非常大且剧烈震荡先降低学习率或者提高target更新间隔。如果loss在下降但reward不涨问题多半出在探索上调epsilon和成功样本比例。如果loss掉到很低但reward还是很差那可能是动作选择策略出了问题检查是不是epsilon衰减太快或者Q值已经过度乐观。第三步打印Q值分布。很多多智能体实现里每个局部Q_i都有误差误差在混合过程中会累积放大最后Q_tot会变得异常大或者异常小。这种情况可以考虑改用Double-DQN形式的target或者在混合网络里加L2正则。多智能体的overestimation问题比单智能体严重得多不能掉以轻心。4.2 单调性假设何时会失效QMIX虽然很强但它不是万能的。它的单调性假设有一个隐含前提单个智能体表现变好不应该让团队整体变差。但部分任务并不满足这一点。最典型的例子是对抗场景中的策略误导。假设两个智能体分别守两个门敌人可能从任意一个门进来。如果智能体A选择“死守自己的门”它自己的局部Q值可能很高但如果团队已经确定敌人会从B门进来A的正确动作其实是去支援B。“死守自己的门”让A的局部Q值上升却会让全局Q_tot下降单调性在这里被打破了。QMIX面对这类非单调信用分配会非常吃力甚至学到完全错误的策略。从理论角度说QMIX只能表示满足IGM条件的值函数。IGM要求全局最优联合动作能被各个局部最优动作无冲突组合出来。如果任务本质上是“只有某些组合动作才有正收益”比如某些矩阵博弈场景QMIX的表现会明显下滑这时候可以考虑QTRAN、WQMIX、QPLEX这类放宽约束的算法。我的实际建议是遇到不收敛先别急着换算法先在简单地图上做一个快速判断。如果在非常简单的地图上QMIX都学不到正收益那大概率不是调参问题而是任务的信用分配结构本身不适合单调分解。这个时候再换WQMIX或QPLEX而不是盲目堆训练时间。4.3 快速问题排查速查表我把平时碰到的典型问题整理成一张表供参考。现象可能原因处理办法loss不降维度广播错、学习率过大打印shape降学习率reward不涨探索不足、成功样本太少提高epsilon增加buffer容量loss下降但reward差单调性假设不满足尝试WQMIX、QPLEXQ值爆炸目标网络更新太慢、误差累积缩短更新间隔调整target训练震荡严重学习率偏高、序列截断不当降学习率按episode索引采样hidden state错乱不同episode被拼接以完整episode存buffer最后一行我特别想说一下。QMIX的经验回放比单智能体DQN要敏感得多因为一个episode里所有智能体的观测、动作、奖励是绑在同一个时间轴上的。如果回放数据里混入了截断不当的序列把不同episode拼接在一起GRU的hidden state历史上下文立刻断开学习效果马上变差。所以在写buffer时必须要以完整episode为最小单位存储采样时也按episode索引抽取再做时间维度的截断。这个小细节很多论文复现教程不会专门讲但在实际训练里影响非常大。4.4 再往前走一步从QMIX到更广阔的MARL如果你已经把QMIX跑通并理解透了接下来有几个方向很值得关注。第一个是WQMIX。它在QMIX框架里引入权重机制对不同样本做差异化处理让目标Q_tot在部分状态下可以暂时违背单调性从而学习更复杂的信用分配。简单说它保留了QMIX的大部分优点又在表达能力上往前走了一步。第二个是QPLEX。它用duplex dueling结构在值函数分解时同时保留优势函数和状态值能够表示更广泛的一类函数同时仍然满足IGM条件。实验上QPLEX在很多SMAC难图上明显超过QMIX代价是实现复杂度更高训练也更敏感。第三个是转向Policy-Based方法比如MAPPO、MADDPG。如果你不需要严格的分布式执行或者动作空间本身是连续的这些on-policy方法会更顺手。但无论用哪个算法QMIX里面关于信用分配、单调分解、CTDE的思想都会以各种形式反复出现。我自己的习惯是离散动作、协作明确的任务优先用QMIX或WQMIX连续控制、需要细粒度协调的任务转向MAPPO。没有被某个算法绑死反而让后续选型更自由。我在实际跑QMIX的过程里印象最深的不是它的分数有多高而是它逼着我去思考“全局信号很稀疏、个体信息很局部”时到底该相信谁。这个思维模式比算法本身更重要。如果你正准备从QMIX入手学习多智能体强化学习建议先在一个小网格环境里把它跑通再去SMAC上做基准测试。踩几次坑之后你会发现很多MARL问题本质上都在回答同一个问题如何把一堆局部判断合成一个整体还不打架的策略。
返回列表