JRL 中的 CQL 离线强化学习实现:配置参数、BC 预热与 D4RL 训练实战指南
2026/9/21 15:53:21 网站建设 项目流程
  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/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:AcmeActorLearnerBuilder实现,负责组装 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)、msgbatch_ensemble_msgsnr等智能体相比,CQL 的核心特点是:完全离线训练(不与环境交互)、以 critic 的 CQL 正则项约束 Q 值估计、并可选地通过 BC 迭代初始化策略。

二、全局 flag 与 gin 配置:理解两类参数的分工

CQL 的运行参数分为两类,理解这一分工是正确配置实验的前提:

  1. 全局 flag:定义在 jrl/localized/runner_flags.py,通过命令行--flag 值直接传入,用于控制运行框架本身;
  2. 算法参数:定义在 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_steps1_000_000训练步数(注意与num_sgd_steps_per_step的换算,见下文)
--eval_every_steps10_000每隔多少 learner 步做一次评估
--episodes_per_eval10每次评估运行的回合数
--batch_size256训练 batch 大小(注意需乘以num_sgd_steps_per_step
--seed1实验随机种子
--root_dir'/tmp/test_msg'实验输出目录
--create_saved_model_actorFalse是否生成 actor 的 saved model
--eval_with_q_filterFalse评估时是否用 Q 值过滤动作(非所有算法支持)
--debug_nansFalse排查 NaN 问题时开启
--spoof_multi_deviceFalse本地调试时模拟多设备
--disable_jitFalse关闭 jit/pmap,便于调试
--single_precision_envFalse是否将环境改为单精度
--gin_configs/--gin_bindingsgin 配置文件路径 / 参数绑定

2.2 为什么推荐使用 gin 而非 flag

算法参数统一收敛到CQLConfig(jrl/agents/cql/config.py),该 dataclass 用@gin.configurable修饰,因此可以用 gin 绑定覆盖任意字段。这带来两个好处:

  • 一个实验的全部算法超参数可以集中、可复现地表达在命令行或 gin 文件中;
  • 每个参数都有明确的默认值与含义,便于横向对比(下文给出完整参数表)。

三、CQLConfig全参数详解

jrl/agents/cql/config.py 中CQLConfig的全部字段、默认值与作用如下:

参数默认值作用
policy_lr3e-5策略网络学习率
q_lr3e-4Q(critic)网络学习率
num_bc_iters50_000初始 behavior cloning 迭代次数(按真实步数计)
cql_alpha5.0CQL 正则项权重
num_importance_acts10CQL 重要性采样动作数
target_entropy0.0自适应熵系数使用的目标熵
num_sgd_steps_per_step1每次 learner step 内执行的 SGD 步数
actor_network_hidden_sizes(256, 256)actor 隐藏层尺寸
critic_network_hidden_sizes(256, 256, 256)critic 隐藏层尺寸
num_critics2critic 个数(取 min 合并)
tau0.005target 网络软更新系数
eval_with_q_filterFalse评估时是否启用 Q 值过滤
num_eval_samples10Q 值过滤评估时采样的动作数
snr_alpha0.0SNR 正则项权重(0 表示关闭)
snr_kwargsSNRKwargs()SNR 子参数对象

其中snr_alphasnr_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_sizenum_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_iterspretrain_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 正则:

  1. 用策略采样num_importance_acts个动作(pi_acts)并计算其对数概率;
  2. [-1, 1]均匀分布中采样同等数量的动作(unif_acts),对数概率为-log(2) * act_dim
  3. 对两类动作分别计算 Q 值,拼接后取logsumexp,再减去数据集中真实动作的 Q 值data_q,得到cql_term = mean(logsumexp - data_q)
  4. 总 critic 损失 = 标准 TD critic 损失 +cql_alpha × cql_term

这里cql_alphanum_importance_acts是调节保守程度与计算量的两个关键参数:cql_alpha越大,对分布外动作的 Q 值惩罚越强;num_importance_acts越大,重要性采样估计越准但计算开销越大。

4.3 多 critic 与 target 软更新

num_critics控制 critic 个数(默认 2),多个 critic 的输出沿最后一维拼接(见 networks.py 中_all_critic_stuffjnp.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 结果,说明该问题并非本实现独有;
  • 仅在halfcheetahhopperwalker实验中使用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=2tau=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_iteratormake_adder返回None,训练数据全部来自make_demonstrations提供的演示数据迭代器。因此batch_sizenum_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=1batch_size=64num_steps=1000episodes_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_itersnum_bc_iters均按真实步数计,无需除以num_sgd_steps_per_step;反之batch_sizenum_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

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

相关推荐

上一篇:FanControl 风扇控制设置指南:3 步让机箱风扇不再狂叫
下一篇:Prometheus Dashboard定制指南:从零开始设计可视化监控面板

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询