NOTE · Reinforcement Learning
QMIX 混合网络:结构、形状与单调性
逐步解释 QMIX 混合网络的超网络、张量形状、单调性约束与适用边界。
参考资料
- 王树森《深度强化学习》中文笔记与课程资料:GitHub 原始仓库
- QMIX 论文:Monotonic Value Function Factorisation for Deep Multi-Agent Reinforcement Learning
- Gymnasium:官方文档与从 Gym 迁移指南
上述资料也曾出现第三方网盘副本。本文只列出作者或项目方的正式入口,避免重复分发来源和版本不明的文件。
问题背景
QMIX 是一种合作式多智能体值函数分解方法。每个智能体依据局部历史估计自己的动作价值 Q_a,训练时再由混合网络结合全局状态 s,输出联合动作价值 Q_tot。混合网络对每个 Q_a 保持单调,可以让集中训练得到的贪心联合动作由各智能体分别贪心选出,支持分散执行。
下面的实现来自公开项目。先看张量流,再看单调性约束;不要把这段混合网络单独等同于完整 QMIX,完整算法还包括智能体网络、回放、目标网络、TD 目标和探索策略。
从旧 Gym/Atari 环境迁移
某个历史项目锁定了 gym==0.10.5、atari-py==0.2.9、旧 Box2D 和 OpenCV,并通过第三方网盘补 DLL/ROM。这套安装链只适合解释项目为何依赖旧接口,不应作为新环境模板。
新项目优先使用维护中的 Gymnasium,并在独立虚拟环境中按项目需要选择 extras:
python -m pip install gymnasium
# 仅在确实使用相应环境时选择:
python -m pip install 'gymnasium[classic-control]'
python -m pip install 'gymnasium[box2d]'
具体 extras、系统依赖和受支持 Python 版本会变化,应在执行时查 Gymnasium 官方安装页。Atari ROM 具有独立许可和获取流程,本页不提供 ROM、DLL 或第三方镜像。
下面两张上游 issue 截图记录了 Windows 下 atari-py 与 Box2D 的历史故障。它们仅用于辨认旧项目依赖,不代表当前安装步骤:


迁移旧训练代码时至少检查四处接口:
import gymnasium as gym
env = gym.make("CartPole-v1")
observation, info = env.reset(seed=2026)
action = env.action_space.sample()
next_observation, reward, terminated, truncated, info = env.step(action)
episode_done = terminated or truncated
if episode_done:
observation, info = env.reset()
else:
observation = next_observation
reset返回(observation, info);step分开返回terminated与truncated;- 时间上限造成的
truncated与任务终止的terminated在价值 bootstrap 中可能需要不同处理; - 随机种子通过
reset(seed=...)等入口设置,不能只调用一个全局 NumPy seed 就假定环境完全可复现; - 需要渲染时通常在
gym.make(..., render_mode=...)创建阶段声明。
如果目标是复现旧论文,应封存旧依赖、操作系统和环境资源的合法来源;如果目标是继续开发,则应迁移接口并重新验证学习曲线,不能只让代码“成功 import”。
混合网络代码
import torch
import torch.nn as nn
import torch.nn.functional as F
class QMixNet(nn.Module):
def __init__(self, args):
super(QMixNet, self).__init__()
self.args = args
if args.two_hyper_layers:
self.hyper_w1 = nn.Sequential( nn.Linear(args.state_shape, args.hyper_hidden_dim),
nn.ReLU(),
nn.Linear(args.hyper_hidden_dim, args.n_agents*args.qmix_hidden_dim))
self.hyper_w2 = nn.Sequential( nn.Linear(args.state_shape, args.hyper_hidden_dim),
nn.ReLU(),
nn.Linear(args.hyper_hidden_dim, args.qmix_hidden_dim))
else:
self.hyper_w1 = nn.Linear(args.state_shape, args.n_agents*args.qmix_hidden_dim)
self.hyper_w2 = nn.Linear(args.state_shape, args.qmix_hidden_dim)
self.hyper_b1 = nn.Linear(args.state_shape, args.qmix_hidden_dim)
self.hyper_b2 = nn.Sequential( nn.Linear(args.state_shape, args.qmix_hidden_dim),
nn.ReLU(),
nn.Linear(args.qmix_hidden_dim, 1))
def forward(self, q_values, states):
episode_num = q_values.size(0)
q_values = q_values.reshape(-1, 1, self.args.n_agents)
states = states.reshape(-1, self.args.state_shape)
w1 = torch.abs(self.hyper_w1(states))
b1 = self.hyper_b1(states)
w1 = w1.reshape(-1, self.args.n_agents, self.args.qmix_hidden_dim)
b1 = b1.reshape(-1, 1, self.args.qmix_hidden_dim)
hidden = F.elu(torch.bmm(q_values, w1) + b1)
w2 = torch.abs(self.hyper_w2(states))
b2 = self.hyper_b2(states)
w2 = w2.reshape(-1, self.args.qmix_hidden_dim, 1)
b2 = b2.reshape(-1, 1, 1)
q_total = torch.bmm(hidden, w2) + b2
q_total = q_total.reshape(episode_num, -1, 1)
return q_total
网络结构
QMixNet是混合网络,不是每个智能体的局部 Q 网络。hyper_w1和hyper_w2根据全局状态生成混合网络的两组权重;hyper_b1和hyper_b2生成偏置。- 因为权重与状态有关,同一组局部 Q 值在不同全局状态下可以得到不同的联合价值。
可选的超网络深度(hypernetwork)
if args.two_hyper_layers:
self.hyper_w1 = nn.Sequential(
nn.Linear(args.state_shape, args.hyper_hidden_dim),
nn.ReLU(),
nn.Linear(args.hyper_hidden_dim, args.n_agents * args.qmix_hidden_dim)
)
self.hyper_w2 = nn.Sequential(
nn.Linear(args.state_shape, args.hyper_hidden_dim),
nn.ReLU(),
nn.Linear(args.hyper_hidden_dim, args.qmix_hidden_dim)
)
else:
self.hyper_w1 = nn.Linear(args.state_shape, args.n_agents * args.qmix_hidden_dim)
self.hyper_w2 = nn.Linear(args.state_shape, args.qmix_hidden_dim)
self.hyper_w1和self.hyper_w2用于生成混合网络中的权重矩阵。- 虽然开关名叫
two_hyper_layers,Linear → ReLU → Linear实际是两个线性层、一个隐藏层,不是两个隐藏层。 - 关闭该开关时,权重由单个线性映射直接生成。
偏置项(biases)
self.hyper_b1 = nn.Linear(args.state_shape, args.qmix_hidden_dim)
self.hyper_b2 = nn.Sequential(
nn.Linear(args.state_shape, args.qmix_hidden_dim),
nn.ReLU(),
nn.Linear(args.qmix_hidden_dim, 1)
)
hyper_b1生成第一层偏置,hyper_b2通过一个隐藏层生成标量偏置。- 偏置不需要非负约束:它们不会改变
Q_tot对各个局部Q_a的偏导符号。
前向传播(forward)
def forward(self, q_values, states):
episode_num = q_values.size(0)
q_values = q_values.reshape(-1, 1, self.args.n_agents)
states = states.reshape(-1, self.args.state_shape)
- 常见输入约定是
q_values: (B, T, n_agents)、states: (B, T, state_shape),其中B为批量大小、T为时间步数;具体仍以调用代码为准。 - 两次
reshape把前导维展平为N = B × T,得到(N, 1, n_agents)和(N, state_shape),便于批量矩阵乘法。 - 这里使用
reshape,避免输入非连续时view失败。
生成w1和b1
w1 = torch.abs(self.hyper_w1(states))
b1 = self.hyper_b1(states)
w1 = w1.view(-1, self.args.n_agents, self.args.qmix_hidden_dim)
b1 = b1.view(-1, 1, self.args.qmix_hidden_dim)
w1和b1由全局状态生成,形状分别为(N, n_agents, qmix_hidden_dim)和(N, 1, qmix_hidden_dim)。torch.abs()令混合权重非负;配合单调递增的 ELU,使∂Q_tot/∂Q_a ≥ 0。这保证的是混合函数的单调性以及集中/分散贪心选择的一致性,不是训练收敛保证。
计算隐藏层输出
hidden = F.elu(torch.bmm(q_values, w1) + b1)
torch.bmm(q_values, w1)将(N, 1, n_agents)与(N, n_agents, qmix_hidden_dim)相乘,输出(N, 1, qmix_hidden_dim)。- 加上
b1后经过单调递增的 ELU,得到混合网络隐藏表示。
生成w2和b2
w2 = torch.abs(self.hyper_w2(states))
b2 = self.hyper_b2(states)
w2 = w2.view(-1, self.args.qmix_hidden_dim, 1)
b2 = b2.view(-1, 1, 1)
w2和b2的形状分别为(N, qmix_hidden_dim, 1)与(N, 1, 1),用于把隐藏表示映射为标量联合价值。
计算总Q值
q_total = torch.bmm(hidden, w2) + b2
q_total = q_total.view(episode_num, -1, 1)
return q_total
torch.bmm(hidden, w2) + b2得到(N, 1, 1),最后恢复为(B, T, 1);代码中的episode_num对应这里的B。
能保证什么,不能保证什么
- 超网络让混合权重随全局状态变化,混合网络因此比简单求和更有表达力。
- 非负权重约束保证
Q_tot对每个局部Q_a单调,从而支持分散贪心执行。 - 单调约束也限制了可表示的联合价值函数;并非所有协作任务的真实
Q_tot都能由这种形式精确表示。 - 该结构不提供“理论收敛”保证。训练是否稳定还取决于目标值、优化器、回放分布、函数逼近和环境非平稳性等完整链路。
张量形状的可执行检查
设 B=4、T=10、n_agents=3、state_shape=20。进入混合网络前可检查:
assert q_values.shape == (4, 10, 3)
assert states.shape == (4, 10, 20)
q_total = mixer(q_values, states)
assert q_total.shape == (4, 10, 1)
assert torch.isfinite(q_total).all()
若训练代码还包含 padding mask、可用动作 mask 或 RNN hidden state,应把时间步和 episode 有效性一起测试。单独验证 mixer 输出形状,不能证明 TD target、mask 和 replay batch 对齐。
abs 约束的工程细节
这份实现用 torch.abs 产生非负权重,能满足单调性,但零点处的梯度行为和权重参数化会影响优化。其他实现可能使用 softplus。两种方式都不能绕过 QMIX 的表达能力限制;更换参数化还会改变训练动态,复现实验时必须记录。
论文参考:QMIX: Monotonic Value Function Factorisation for Deep Multi-Agent Reinforcement Learning。