- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
本指南以 Google Research 的 JRL(Jax Reinforcement Learning)代码库中的 CQL(Conservative Q-Learning)智能体为主线,系统讲解其架构组成、全部 gin 可配置参数、num_sgd_steps_per_step加速机制,以及基于jrl.localized.runner在 D4RL 数据集上运行完整训练实验的命令行流程。读完本文,你将掌握如何在 JRL 代码库中复现 CQL 基准实验、如何针对 D4RL gym 与 antmaze 场景调整超参数,并能深入理解 BC 预训练(behavior cloning warm-up)与 Q 值过滤评估(q-filter evaluation)等实现细节。
一、CQL 在 JRL 代码库中的定位与模块结构
JRL 是面向离线强化学习研究的 Jax 代码库,基于 Acme RL 库实现(参见 jrl/README.md)。CQL(Conservative Q-Learning)是其中的一个离线 RL 智能体(agent),其完整实现位于 jrl/agents/cql 目录,包含五个文件:
- README.md:使用说明与运行示例,即本文主体来源;
- config.py:
CQLConfiggin 可配置 dataclass,集中定义全部算法参数; - builder.py:Acme
ActorLearnerBuilder实现,负责组装 learner 与 actor; - learning.py:
CQLLearner核心训练逻辑(critic 损失、CQL 正则项、actor 损失与 BC 迭代); - networks.py:策略网络与多 critic 网络的 Haiku 定义。
按照 jrl/agents/README.md 描述的代码组织约定,每个 agent 均由 RL 组件(注册于 jrl/agents/init.py)、Builder、Config、Learner、Networks 五部分组成,CQL 完全遵循这一范式。与同目录下的bc(behavior cloning)、msg、batch_ensemble_msg、snr等智能体相比,CQL 的核心特点是:完全离线训练(不与环境交互)、以 critic 的 CQL 正则项约束 Q 值估计、并可选地通过 BC 迭代初始化策略。
二、全局 flag 与 gin 配置:理解两类参数的分工
CQL 的运行参数分为两类,理解这一分工是正确配置实验的前提:
- 全局 flag:定义在 jrl/localized/runner_flags.py,通过命令行
--flag 值直接传入,用于控制运行框架本身; - 算法参数:定义在 jrl/agents/cql/config.py 的
CQLConfigdataclass 中,通过--gin_bindings='cql.config.CQLConfig.xxx=...'传入。
2.1 全局 flag 速查表
jrl/localized/runner_flags.py 定义的常用全局 flag 如下:
| Flag | 默认值 | 说明 |
|---|---|---|
--algorithm | 'msg' | 要运行的算法名,运行 CQL 时需设为'cql' |
--task_class | 'd4rl' | 任务类别,如d4rl |
--task_name | 'halfcheetah-medium-replay-v0' | 具体任务名,如'antmaze-large-diverse-v0' |
--num_steps | 1_000_000 | 训练步数(注意与num_sgd_steps_per_step的换算,见下文) |
--eval_every_steps | 10_000 | 每隔多少 learner 步做一次评估 |
--episodes_per_eval | 10 | 每次评估运行的回合数 |
--batch_size | 256 | 训练 batch 大小(注意需乘以num_sgd_steps_per_step) |
--seed | 1 | 实验随机种子 |
--root_dir | '/tmp/test_msg' | 实验输出目录 |
--create_saved_model_actor | False | 是否生成 actor 的 saved model |
--eval_with_q_filter | False | 评估时是否用 Q 值过滤动作(非所有算法支持) |
--debug_nans | False | 排查 NaN 问题时开启 |
--spoof_multi_device | False | 本地调试时模拟多设备 |
--disable_jit | False | 关闭 jit/pmap,便于调试 |
--single_precision_env | False | 是否将环境改为单精度 |
--gin_configs/--gin_bindings | 空 | gin 配置文件路径 / 参数绑定 |
2.2 为什么推荐使用 gin 而非 flag
算法参数统一收敛到CQLConfig(jrl/agents/cql/config.py),该 dataclass 用@gin.configurable修饰,因此可以用 gin 绑定覆盖任意字段。这带来两个好处:
- 一个实验的全部算法超参数可以集中、可复现地表达在命令行或 gin 文件中;
- 每个参数都有明确的默认值与含义,便于横向对比(下文给出完整参数表)。
三、CQLConfig全参数详解
jrl/agents/cql/config.py 中CQLConfig的全部字段、默认值与作用如下:
| 参数 | 默认值 | 作用 |
|---|---|---|
policy_lr | 3e-5 | 策略网络学习率 |
q_lr | 3e-4 | Q(critic)网络学习率 |
num_bc_iters | 50_000 | 初始 behavior cloning 迭代次数(按真实步数计) |
cql_alpha | 5.0 | CQL 正则项权重 |
num_importance_acts | 10 | CQL 重要性采样动作数 |
target_entropy | 0.0 | 自适应熵系数使用的目标熵 |
num_sgd_steps_per_step | 1 | 每次 learner step 内执行的 SGD 步数 |
actor_network_hidden_sizes | (256, 256) | actor 隐藏层尺寸 |
critic_network_hidden_sizes | (256, 256, 256) | critic 隐藏层尺寸 |
num_critics | 2 | critic 个数(取 min 合并) |
tau | 0.005 | target 网络软更新系数 |
eval_with_q_filter | False | 评估时是否启用 Q 值过滤 |
num_eval_samples | 10 | Q 值过滤评估时采样的动作数 |
snr_alpha | 0.0 | SNR 正则项权重(0 表示关闭) |
snr_kwargs | SNRKwargs() | SNR 子参数对象 |
其中snr_alpha、snr_kwargs与 jrl/agents/snr 模块相关:SNR(spectral norm regularization)源自独立的研究方向,与 CQL 本身无关(README 明确说明 "The SNR hyperparameters are unrelated to CQL and are from orthogonal research ideas")。snr_alpha=0时 learning.py 中self._use_snr = snr_alpha > 0.为 False,SNR 完全不参与计算。
3.1num_sgd_steps_per_step加速机制
README 特别强调:所有 agent 都有num_sgd_steps_per_step配置,它决定每次调用 learner 的 step 函数时执行多少次训练步骤。调大该值能让 Jax 执行跨 batch 的编译优化从而显著提速,但必须同步调整batch_size与num_steps:
batch_size应设为num_sgd_steps_per_step × 单步期望 batch 大小;num_steps应设为期望总训练步数 / num_sgd_steps_per_step。
在 learning.py 中可以看到实现细节:learner 通过utils.process_multiple_batches(..., num_sgd_steps_per_step)包装更新函数,step()内每步消费的 transition 数量即按该系数放大,且计数器按steps=self._num_sgd_steps_per_step递增。例如 README 中的调试命令:
--num_sgd_steps_per_step 1 \ --batch_size 64 \ --num_steps 1000 \ --episodes_per_eval 10 \ --gin_bindings='cql.config.CQLConfig.num_sgd_steps_per_step=1'即单步 batch 64、共 1000 个 learner step,适合本地快速验证。
3.2num_bc_iters与pretrain_iters按真实步数计
README 提醒:cql.config.CQLConfig.num_bc_iters(BC 预热迭代数)以"真实步数"计,无需按num_sgd_steps_per_step换算;同理pretrain_iters。在 learning.py 的step()中,通过cur_step < self._num_bc_iters判断当前是否处于 BC 阶段(in_initial_bc_iters),这里的cur_step来自 counter 的learner_steps,与num_sgd_steps_per_step无关。
四、CQL 训练流程与 BC 预热阶段
CQL 的训练逻辑集中在 jrl/agents/cql/learning.py 的CQLLearner中,整体是 SAC 风格的双 critic + 自适应熵训练,外加 CQL 正则项与可选 BC 预热。
4.1 两个训练阶段
- BC 阶段(前
num_bc_iters步):actor_loss = -log_prob,即纯行为克隆,最大化数据集中动作的对数似然;此时 critic 的 CQL 项依然生效(total = critic_loss_term + cql_alpha * cql_term),但 actor 不接收 Q 值信号,也不应用 SNR(源码中# No SNR in bc iters)。 - 正常阶段(BC 之后):
actor_loss = alpha * log_prob - min_q,即标准的 SAC 风格策略更新:最大化熵正则的 Q 值;若snr_alpha > 0,还会叠加snr_alpha * sn的 SNR 正则项。
4.2 CQL 正则项的实现
cql_loss(learning.py)按如下方式构造 CQL 正则:
- 用策略采样
num_importance_acts个动作(pi_acts)并计算其对数概率; - 在
[-1, 1]均匀分布中采样同等数量的动作(unif_acts),对数概率为-log(2) * act_dim; - 对两类动作分别计算 Q 值,拼接后取
logsumexp,再减去数据集中真实动作的 Q 值data_q,得到cql_term = mean(logsumexp - data_q); - 总 critic 损失 = 标准 TD critic 损失 +
cql_alpha × cql_term。
这里cql_alpha与num_importance_acts是调节保守程度与计算量的两个关键参数:cql_alpha越大,对分布外动作的 Q 值惩罚越强;num_importance_acts越大,重要性采样估计越准但计算开销越大。
4.3 多 critic 与 target 软更新
num_critics控制 critic 个数(默认 2),多个 critic 的输出沿最后一维拼接(见 networks.py 中_all_critic_stuff的jnp.concatenate(critic_preds, axis=-1)),在 TD 目标与 actor 损失中均取jnp.min(..., axis=-1)融合。target 网络按tau做指数滑动平均:target_q_params = (1 - tau) * target_q_params + tau * q_params。
4.4 关于复现结果的说明
README 给出重要的复现经验(引用自 MSG 论文的实验报告):
- 本实现能很好地复现 D4RL gym 的 CQL 结果(论文中同时采用了本实现与另一份非本仓库实现);
- 未能复现 CQL 的 antmaze 结果——作者与同事用其他公开实现也无法复现 antmaze 结果,说明该问题并非本实现独有;
- 仅在
halfcheetah、hopper、walker实验中使用num_critics=2,该选择差异不显著,作者未回退重跑全套实验; - 对于
antmaze域,num_critics取 1 或 2 都尝试过。
这些说明对后续研究者具有直接参考价值:在 D4RL gym 上以本文命令为基线,antmaze 上则需自行斟酌。
五、D4RL 上的完整运行命令
README 给出的 CQL 标准运行命令(antmaze-large-diverse 任务):
python3 -m jrl.localized.runner \ --pdb_post_mortem \ --debug_nans=False \ --create_saved_model_actor=False \ --num_steps 11000 \ --eval_every_steps 500 \ --episodes_per_eval 100 \ --batch_size 51200 \ --root_dir '/tmp/test_cql' \ --seed 42 \ --algorithm 'cql' \ --task_class 'd4rl' \ --task_name 'antmaze-large-diverse-v0' \ --gin_bindings='cql.config.CQLConfig.num_sgd_steps_per_step=200' \ --gin_bindings='cql.config.CQLConfig.num_bc_iters=50000' \ --gin_bindings='cql.config.CQLConfig.cql_alpha=0.05' \ --gin_bindings='cql.config.CQLConfig.num_importance_acts=10' \ --gin_bindings='cql.config.CQLConfig.actor_network_hidden_sizes=(256, 256, 256)' \ --gin_bindings='cql.config.CQLConfig.critic_network_hidden_sizes=(256, 256, 256)' \ --gin_bindings='cql.config.CQLConfig.num_critics=2' \ --gin_bindings='cql.config.CQLConfig.tau=0.005' \ --gin_bindings='cql.config.CQLConfig.eval_with_q_filter=False' \ --gin_bindings='cql.config.CQLConfig.num_eval_samples=32' \ --gin_bindings='cql.config.CQLConfig.snr_kwargs=@snr.config.SNRKwargs()' \ --gin_bindings='cql.config.CQLConfig.snr_alpha=0' \ --gin_bindings='snr.config.SNRKwargs.snr_mode="params_kernel"' \ --gin_bindings='snr.config.SNRKwargs.snr_loss_type="svd_kamyar_v1"' \ --gin_bindings='snr.config.SNRKwargs.use_log_space_matrix=False'逐项解读该命令:
- 框架参数:
--num_steps 11000(learner step 数)、--eval_every_steps 500(每 500 步评估一次)、--episodes_per_eval 100(每次评估 100 个回合)、--batch_size 51200、--seed 42、--root_dir '/tmp/test_cql'。 - batch 换算示例:
num_sgd_steps_per_step=200时,batch_size=51200相当于单步真实 batch 为 256(51200 ÷ 200),即"每步 200 个 batch × 每 batch 256 条 transition"。若按单步 256 计算,11000 learner steps × 200 × 256 ≈ 5.6 亿条 transition 的等效训练量;相应地,num_steps已被除以 200,因此不要再额外换算。 - 算法参数:
num_bc_iters=50000(按真实步数计,BC 预热 5 万步)、cql_alpha=0.05(antmaze 常用较小值)、num_importance_acts=10、actor/critic 均为三隐层 256、num_critics=2、tau=0.005。 - 评估相关:
eval_with_q_filter=False关闭 Q 值过滤(此时num_eval_samples=32不生效);snr_alpha=0关闭 SNR,但仍显式绑定snr_kwargs与三个 SNR 子参数——这是为了让cql.config.CQLConfig.snr_kwargs的 gin 绑定生效(见 config.py 的注释:绑定snr.config.SNRConfig.snr_kwargs的方式对 CQL 同样适用)。
5.1 为什么batch_size要取这么大
结合 builder.py 可以看到,CQL 是纯离线训练:make_replay_tables返回空表、make_dataset_iterator与make_adder返回None,训练数据全部来自make_demonstrations提供的演示数据迭代器。因此batch_size与num_sgd_steps_per_step的乘积直接决定每个 learner step 的吞吐量,调大num_sgd_steps_per_step配合大 batch,能让 Jax 在一次编译后的多 batch 处理中摊销开销、显著提速(learning.py 中utils.process_multiple_batches即为此设计)。
5.2 本地调试推荐配置
README 给出的最小化调试配置(回到本文 3.1 节的命令):num_sgd_steps_per_step=1、batch_size=64、num_steps=1000、episodes_per_eval=10。此外可组合 runner_flags.py 中的--spoof_multi_device(单机模拟多设备)与--disable_jit(关闭 jit/pmap)来快速定位 NaN 或形状错误。
六、Q 值过滤评估(可选)
当eval_with_q_filter=True时,builder.py 的make_actor会同时向 actor 传入['policy', 'q']两组参数,并使用 networks.py 中build_q_filtered_actor构建评估 actor:从策略分布采样num_eval_samples个动作(可叠加[-1,1]均匀采样),用所有 critic 的 min-Q 打分并选取最优动作执行。该机制可在评估时利用 Q 值对动作做二次筛选,代价是推理开销增大,README 的标准命令默认关闭它。
七、运行环境与注意事项
- 环境准备遵循 jrl/README.md 的建议:
conda env create -f requirements.yml,依赖定义见 jrl/requirements.yml。 - 运行入口为 jrl/localized/runner.py,以
python3 -m jrl.localized.runner方式启动,--pdb_post_mortem表示出错时进入事后调试器。 - 修改算法或新增 agent 时,参照 jrl/agents/README.md 的五层结构约定:RL 组件注册到 jrl/agents/init.py,配置统一走 gin dataclass。
- CQL 的
pretrain_iters与num_bc_iters均按真实步数计,无需除以num_sgd_steps_per_step;反之batch_size、num_steps必须换算。
结语
本文从 jrl/agents/cql/README.md 出发,结合 config.py、builder.py、learning.py、networks.py 与 runner_flags.py 等源码,完整覆盖了 JRL 中 CQL 智能体的参数体系、BC 预热机制、CQL 正则项原理与 D4RL 训练实操。关键结论可归纳为:全局框架参数走 flag、算法参数走 gin;调大num_sgd_steps_per_step可提速但需同步换算 batch 与步数;num_bc_iters按真实步数计;D4RL gym 结果可复现而 antmaze 存在已知的复现困难。基于以上内容,你可以直接在 D4RL gym 任务上展开 CQL 实验,并以此为基线继续研究离线 RL 的保守性与不确定性估计问题。
- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
相关推荐
TensorLayer离线强化学习算法对比:BCQ、CQL与TD3+BC
TensorLayer离线强化学习算法对比:BCQ、CQL与TD3+BC 你是否在训练强化学习智能体时遇到过数据稀缺或分布偏移的问题?离线强化学习(Offlin
人工智能深度学习机器学习强化学习MarkItDown 快速上手:把办公文档转成 Markdown 的保姆级指南
MarkItDown 快速上手:把办公文档转成 Markdown 的保姆级指南 你有没有被这种场景折腾过:手里一堆 PDF、Word、Excel 和 PPT,想
人工智能深度学习NLP计算机视觉强化学习从源码到部署:LibVNCServer全平台编译与安装教程
从源码到部署:LibVNCServer全平台编译与安装教程 LibVNCServer是一款功能强大的跨平台C语言库,能够帮助开发者轻松实现VNC服务器或客户端功
网络通信
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考