news 2026/9/20 14:25:00

AQuaDem 源码实战:基于演示动作量化的连续控制算法解析与运行指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
AQuaDem 源码实战:基于演示动作量化的连续控制算法解析与运行指南
  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/google-research
点击查看免费下载

AQuaDem(ActionQuantizationandDemonstrations,全称 Continuous Control with Action Quantization from Demonstrations)是 Google Research 团队提出的从演示中学习连续控制任务的算法:它先把连续动作空间"量化"为一组由多模态行为克隆(Multi-head BC)学到的离散候选动作,再在这些候选动作上运行标准离散强化学习算法进行决策。本文以 aquadem 目录源码为主线,完整介绍该算法的依赖安装、运行命令、核心模块结构与配置参数,并深入 learning.py、networks.py、builder.py 等源码解释其"多分类 BC 预训练 + 离散 RL 微调"的两阶段机制。读完本文,你将能够独立安装并运行 AQuaDem,理解每个命令行参数与配置文件项的含义,并掌握将其迁移到自定义连续控制任务的要点。

一、AQuaDem 算法思想概览

AQuaDem 解决的核心问题是:如何在连续动作空间上利用专家演示提升强化学习(RL)的训练效率与最终性能。直接在高维连续动作空间做 RL 通常样本效率低,而纯粹的行为克隆(BC)又受限于演示质量、无法超越专家。

AQuaDem 的做法是"折中":

  1. 动作量化:训练一个多模态动作编码器(multi-modal encoder,即多分类 BC),给定观测时同时输出num_actions个候选连续动作,覆盖专家演示中出现的多种行为模式;
  2. 离散化决策:把原始连续动作空间替换为这N个候选动作,环境的原始连续动作 spec 被改写为离散 spec(见 discretize_spec);
  3. 离散 RL:在候选动作上运行标准的离散动作 RL 算法(本实现使用基于 Munchausen Q-Learning 的 DQN),让 agent 学会"选哪一个候选动作",从而规避连续动作空间探索的困难。

该目录为论文Continuous Control with Action Quantization from Demonstrations(Robert Dadashi、Leonard Hussenot 等,arXiv 2110.10149)的官方源码实现,基于 DeepMind Acme 框架构建,采用 JAX 实现网络与训练逻辑。

二、环境准备与依赖安装

AQuaDem 的依赖通过 requirements.txt 管理,核心依赖包括:

  • Acme 生态dm-acme==0.3.0dm-reverb==0.6.1(回放缓冲区)、dm-envdm-haikudm-sonnet
  • 深度学习jax==0.2.28flax==0.4.0optax==0.1.1tensorflow==2.7.0
  • 环境与数据gym==0.21.0dm-controltensorflow_datasets==4.4.0rlds==0.1.3,以及通过 git 安装的rlax(DeepMind)与d4rl(Berkeley,提供 Adroit 演示数据集)。

安装命令:

pip install -r requirements.txt

两个需要特别注意的前提条件:

  1. Python 版本:官方声明 AQuaDem 兼容Python 3.9,且依赖版本(如jax==0.2.28tensorflow==2.7.0)均为较早期版本,建议使用 Python 3.9 的虚拟环境安装,避免与新版依赖产生冲突;
  2. MuJoCo 物理引擎:AQuaDem 需要在 MuJoCo2.1.1下运行(Adroit 灵巧手任务依赖 MuJoCo),请按 MuJoCo 官方提供的 2.1.1 版本安装指引完成安装后,再运行上面的pip install

三、快速启动:运行 AQuaDem

安装完成后,在仓库根目录执行以下命令即可启动训练:

python -m aquadem.run_aquadqn --workdir='/tmp/aquadem' --env_name='door-human-v1'

其中workdir为日志输出目录(训练与评估指标以 CSV 形式写入),env_name为要运行的环境。默认情况下会使用 100 万步环境交互训练door-human-v1(Adroit 门开关任务的人类演示版本)。

3.1 命令行参数

入口脚本 run_aquadqn.py 通过 absl.flags 定义全部命令行参数:

参数类型默认值说明
--workdirstr/tmp/aquadqn日志输出目录,训练/评估指标以 CSV 保存
--env_namestrdoor-human-v1运行的环境名称(D4RL Adroit 任务,如door-human-v1hammer-human-v1等)
--num_demonstrationsintNone使用的专家演示条数;None表示使用完整数据集
--num_stepsint1000000训练总环境步数
--eval_everyint10000评估频率(每多少步评估一次)
--seedint0RL agent 的随机种子

注意:源码要求num_steps % eval_every == 0(run_aquadqn.py),否则程序会通过 assert 中断。

3.2 训练-评估循环

run_aquadqn.py 中采用"先评估后训练"的交替循环:

for _ in range(FLAGS.num_steps // FLAGS.eval_every): eval_loop.run(num_episodes=10) # 每次评估跑 10 个 episode train_loop.run(num_steps=FLAGS.eval_every) # 再训练 eval_every 步 eval_loop.run(num_episodes=10)

评估环境在训练环境基础上额外包裹了SuccessRewardWrapper(wrappers.py),即整个 episode 首次累积回报达到阈值即返回奖励 1,因此评估回报就是"是否成功"的 0/1 指示,便于直观衡量任务成功率。

四、源码架构总览

aquadem目录共有 9 个文件,职责划分清晰:

文件职责
run_aquadqn.py程序入口:组装环境、builder、网络、训练/评估循环
config.pyAquademConfig数据类,集中管理全部算法超参数
builder.pyAquademBuilder:实现 Acme 的ActorLearnerBuilder接口,负责组装 learner/actor/回放
learning.pyMultiBCLearnerAquademLearner:两阶段学习核心逻辑
networks.pyFlaxEncoder网络(多候选动作生成)、DQN Q 网络
actor.pyAquademActor:执行"离散选择 → 连续动作映射"的动作生成
utils.py环境创建、演示数据集加载与奖励稀疏化
wrappers.pyAdroit 任务的稀疏奖励/成功奖励包装器
requirements.txtPython 依赖清单

整个算法可以概括为"Encoder 预训练阶段 + 离散 RL 阶段"两阶段流水线,下面分别深入剖析。

五、核心机制一:多分类行为克隆(MultiBC)学习动作候选

5.1 Encoder 网络结构

动作候选生成器定义在 networks.py 的Encoder(Flaxlinen.Module)中,其结构为共享 torso +num_actions个独立 head

  • 输入观测先过一个共享 torso(默认torso_layer_sizes=(256,),即一层 256 维全连接 + ReLU);
  • 随后并行搭建num_actions个 head(默认head_layer_sizes=(256,)),每个 head 最终输出action_dim维连续动作;
  • 最后通过jnp.stack(actions, axis=-1)把所有候选动作堆叠为形状[batch, action_dim, num_actions]的张量。

也就是说,给定一个观测,Encoder 一次给出num_actions个不同的候选动作,每个候选对应一种可能的行为模式。网络在输入层与隐藏层均使用 dropout(默认input_dropout_rate=0.1hidden_dropout_rate=0.1),以鼓励多个 head 分化、避免坍缩到同一动作。

5.2 Softmin 距离损失

MultiBCLearner(learning.py)使用一种"软最小值"(softmin)损失训练 Encoder。其核心函数aqualoss(learning.py)的计算过程是:

  1. 计算每个候选动作与演示动作的平方 L2 距离:action_distances = sum((predicted - action)^2)
  2. num_actions个距离做 softmin 聚合:
softmin_action_distances = temperature * ( jax.nn.logsumexp(-action_distances / temperature) - jnp.log(num_actions)) loss = -softmin_action_distances

直觉上,损失鼓励"至少有一个候选动作接近专家动作"(因为 softmin 近似于取最小值),而非强迫所有 head 都预测同一个动作——这正是它能学到多模态演示分布的关键。参数temperature(默认0.001)控制 softmin 的锐利程度。

5.3 预训练阶段参数

预训练由AquademConfig(config.py)中的字段控制:

配置项默认值说明
num_actions10学习到的候选动作数量(量化粒度)
encoder_learning_rate3e-4MultiBC 使用的 Adam 学习率
encoder_batch_size256预训练 batch 大小
encoder_num_steps50_000MultiBC 预训练总步数
encoder_eval_every1_000预训练内部记录频率(每步包含encoder_eval_every次 SGD)
temperature0.001softmin 聚合的温度

在 builder.py 中,Encoder 使用optax.adam(encoder_learning_rate)优化;learning.py 中AquademLearner构造时即先完成整个预训练:以encoder_batch_size * encoder_eval_every条演示构造数据集,循环encoder_num_steps // encoder_eval_every次调用MultiBCLearner.step(),每次内部通过jax.jit编译执行encoder_eval_every次 SGD(见process_multiple_batches,learning.py)。

六、核心机制二:离散 RL 在候选动作上学习

6.1 从连续演示到离散标签

预训练完成后,AQuaDem 需要把回放数据"翻译"成离散 RL 可用的形式。_generate_aquadem_samples(learning.py)以概率demonstration_ratio(默认0.25)从演示数据中采样,并将连续专家动作映射为距离最近的候选动作索引:

discrete_actions = np.argmin( np.linalg.norm(continuous_actions_candidates - demonstrations.action[:, :, None], axis=1), axis=-1)

同时,若配置了min_demo_reward,还会把演示样本的奖励下限提升到该值(reward = max(min_demo_reward, reward)),从而"鼓励 agent 加入专家的支撑集"。其余情况下直接透传回放缓冲区中的交互样本。

6.2 离散 RL:Munchausen Q-Learning

在离散化后的动作空间上,AQuaDem 复用了 Acme 的 DQN 实现(run_aquadqn.py):

loss_fn = dqn.losses.MunchausenQLearning(max_abs_reward=100.) dqn_config = dqn.DQNConfig( min_replay_size=1000, n_step=3, num_sgd_steps_per_step=8, learning_rate=1e-4, samples_per_insert=256) rl_agent = dqn.DQNBuilder(config=dqn_config, loss_fn=loss_fn)

其中关键的超参数含义为:

参数默认值说明
min_replay_size1000回放缓冲区最少积累多少样本后开始学习
n_step33 步回报(n-step TD)
num_sgd_steps_per_step8每个环境步执行的 SGD 次数
learning_rate1e-4DQN 学习率
samples_per_insert256每次插入回放时采样的样本数
max_abs_reward100.Munchausen 损失的奖励裁剪上限

Q 网络由make_q_network(networks.py)构建,默认使用LayerNormMLP,隐藏层为(512, 512, 256),输出维度等于num_actions。值得注意的是,源码保留了architecture='MLP'分支(注释为 "AQuaOff architecture"),即论文后续工作 AQuaOff 的变体,可通过该参数切换网络架构。

6.3 两阶段的数据流

AquademLearner(learning.py)把两阶段串成完整闭环:

  1. 构造MultiBCLearner并完成encoder_num_steps步预训练(见 5.3);
  2. 将演示迭代器与回放迭代器喂给_generate_aquadem_samples,生成"离散标签 + 混合演示/交互"的学习数据流lfd_iterator(Learning from Demonstrations);
  3. lfd_iterator作为离散 RL learner 的数据源,此后每次step()只推进离散 RL 的学习。

因此,num_actionsdemonstration_ratiomin_demo_reward是决定"演示在离散 RL 阶段发挥多大作用"的三个核心旋钮。

七、Actor:如何把离散决策变成连续动作

训练与评估时执行动作的组件是AquademActor(actor.py),它包装了一个离散动作 actor:

  1. select_action(observation)先调用内部离散策略(DQN 的default_behavior_policy)选出候选索引discrete_action
  2. 再用 Encoder 对观测输出全部候选动作,取第discrete_action个作为最终连续动作:
def aquadem_policy(params, observation, discrete_action): predicted_actions = networks.encoder.apply(params, observation) return predicted_actions[..., discrete_action]
  1. 环境交互回放时记录的是离散动作observe()使用self._last_discrete_action),从而保证离散 RL 的训练数据一致。

训练期间离散策略使用exploration_epsilon = 0.01的 ε-greedy 探索(run_aquadqn.py),评估时 ε 设为 0 执行纯贪心策略。AquademActor中的 encoder 变量通过VariableClient从 learner 同步,且由于 Encoder 预训练完成后不再更新,update_period被设为极大的值以"永不更新"(builder.py)。

八、演示数据与环境处理

8.1 D4RL Adroit 数据集加载

演示数据通过 TFDS 加载 D4RL Adroit 数据集,utils.py 的_d4rl_dataset_name把环境名转换为 TFDS 数据集名,例如:

  • door-human-v1d4rl_adroit_door/v1-human
  • hammer-human-v1d4rl_adroit_hammer/v1-human

get_make_demonstrations_fn(utils.py)负责:加载 TFDS 数据集 →(可选)截取前num_demonstrations条 →稀疏化奖励(按任务阈值把稠密奖励转为 0/1)→ 转换为 Acme Transition 迭代器,最终返回一个按 batch 大小生成随机演示 batch 的函数。

8.2 奖励稀疏化与任务阈值

Adroit 灵巧手任务(door、hammer、pen、relocate)的稀疏奖励阈值定义在 utils.py:

SPARSE_REWARD_THRESHOLDS = {'door': 15, 'hammer': 50, 'pen': 30, 'relocate': 5}

演示数据中的奖励被稀疏化为reward > threshold的 0/1 指示。训练环境则通过AdroitSparseRewardWrapper(wrappers.py)直接用环境自身的info['goal_achieved']作为奖励,与演示数据稀疏化阈值保持一致;评估环境再叠加SuccessRewardWrapper,把"整条轨迹是否成功"作为评估指标。

8.3 环境包装链

make_environment(utils.py)构建环境的完整包装链为:

gym.make(task) → AdroitSparseRewardWrapper(goal_achieved 作为奖励) → GymWrapper(转为 dm_env 接口) → CanonicalSpecWrapper(clip=True,裁剪动作到 spec 范围) → SinglePrecisionWrapper(单精度) → (仅评估时)SuccessRewardWrapper(整条轨迹成功即 1)

九、配置参数速查与调参建议

综合 config.py 与 run_aquadqn.py,AQuaDem 的全部核心可调参数汇总如下:

算法层(AquademConfig,在代码中修改)num_actions(候选动作数,越大表达能力越强但离散 RL 难度越高)、encoder_learning_rateencoder_batch_sizeencoder_num_steps(预训练步数)、temperature(softmin 温度)、demonstration_ratio(演示样本混合比例)、min_demo_reward(演示奖励下限)。

RL 层(DQNConfigmin_replay_sizen_stepnum_sgd_steps_per_steplearning_ratesamples_per_insertmax_abs_reward

命令行层--workdir--env_name--num_demonstrations--num_steps--eval_every--seed

常见调参思路(基于源码逻辑推断):

  • 候选动作数num_actions越大,Encoder 对多模态演示的覆盖越充分,但离散 RL 的 action space 也越大,可相应增加num_sgd_steps_per_step或总步数--num_steps
  • 演示充足时可提高demonstration_ratio(默认 0.25)并设置合理的min_demo_reward,强化演示对离散 RL 的引导;
  • temperature影响 softmin 的锐利度,过小可能导致训练不稳定,过大会让损失退化为均值近似;
  • 切换环境时务必确认env_name对应的任务在SPARSE_REWARD_THRESHOLDS中已有阈值(当前仅支持 door、hammer、pen、relocate 四个 Adroit 任务)。

十、总结

AQuaDem 通过"多分类 BC 学习动作候选 + 离散 RL 学习候选选择"的两阶段设计,把连续控制问题转化为离散决策问题,同时让专家演示同时作用于 Encoder 预训练与 RL 训练两个环节。仓库源码以 Acme/JAX 生态实现,结构清晰、模块边界明确:networks.py负责多候选动作生成与 Q 网络,learning.py承载两阶段学习核心,actor.py完成离散到连续的最终映射,utils.py/wrappers.py解决 Adroit 演示数据与环境奖励的一致性问题。理解这些模块后,你可以参照 run_aquadqn.py 的组装方式,将 AQuaDem 适配到自定义的连续控制任务(需自行准备对应的演示数据集与奖励阈值)。

  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/google-research
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/20 14:24:41

GTQ-FC100T可燃气体变送器原理与工业集成实战指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/20 14:24:40

JESD47I标准解析:半导体器件可靠性评估与应力测试方案设计指南

简介:JESD47I中文版是JEDEC(电子器件工程委员会)发布的集成电路压力测试考核标准的中文编译版,主要面向半导体可靠性工程师、质量与测试人员及电子工程相关专业学习者。文档系统介绍了应力测试驱动的集成电路合格认证方法&#xf…

作者头像 李华
网站建设 2026/9/20 14:24:31

MATLAB多变量时间序列预测:Transformer-LSTM与贝叶斯优化实战

简介:面向具备MATLAB及深度学习基础的开发者、研究人员,以及智能制造、金融市场、气象预报、能源管理等领域的时序预测从业者,这份项目实例围绕BO-Transformer-LSTM多变量时间序列预测展开。资源针对Transformer-LSTM复合模型结构复杂、超参数…

作者头像 李华
网站建设 2026/9/20 14:22:20

Win10服务禁用风险与依赖关系深度解析

1. 为什么“禁用Win10服务”成了重装系统的前奏?你有没有试过——刚装好干净的Win10,兴致勃勃打开“服务”管理器(services.msc),看到密密麻麻上百个条目,心里一热:“这么多后台跑着&#xff0c…

作者头像 李华
网站建设 2026/9/20 14:21:27

CANN Runtime 对外 ACL 日志接口实战:acllog 系列 API 可运行样例全解析

CANNAscend人工智能任务调度 【免费下载链接】runtime 本项目提供CANN运行时组件和维测功能组件。 项目地址: https://gitcode.com/cann/runtime 点击查看 免费下载 本篇文章以 CANN/runtime 仓库中 example/5_performance/log 目录下的 ACL 日志样例为骨架&#x…

作者头像 李华