NOTE · Reinforcement Learning

QMIX 混合网络:结构、形状与单调性

逐步解释 QMIX 混合网络的超网络、张量形状、单调性约束与适用边界。

参考资料

上述资料也曾出现第三方网盘副本。本文只列出作者或项目方的正式入口,避免重复分发来源和版本不明的文件。

问题背景

QMIX 是一种合作式多智能体值函数分解方法。每个智能体依据局部历史估计自己的动作价值 Q_a,训练时再由混合网络结合全局状态 s,输出联合动作价值 Q_tot。混合网络对每个 Q_a 保持单调,可以让集中训练得到的贪心联合动作由各智能体分别贪心选出,支持分散执行。

下面的实现来自公开项目。先看张量流,再看单调性约束;不要把这段混合网络单独等同于完整 QMIX,完整算法还包括智能体网络、回放、目标网络、TD 目标和探索策略。

从旧 Gym/Atari 环境迁移

某个历史项目锁定了 gym==0.10.5atari-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 的历史故障。它们仅用于辨认旧项目依赖,不代表当前安装步骤:

旧版 Gym Atari 安装 issue 截图

旧版 Box2D RAND_LIMIT issue 截图

迁移旧训练代码时至少检查四处接口:

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 分开返回 terminatedtruncated
  • 时间上限造成的 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_w1hyper_w2 根据全局状态生成混合网络的两组权重;hyper_b1hyper_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_w1self.hyper_w2 用于生成混合网络中的权重矩阵。
  • 虽然开关名叫 two_hyper_layersLinear → 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)
  • w1b1 由全局状态生成,形状分别为 (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)
  • w2b2 的形状分别为 (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=4T=10n_agents=3state_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