ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

CAQL 连续动作 Q-Learning 实践指南:面向 max-Q 问题的可插拔优化器设计与源码解析

2026/9/20 12:16:54 拓冰建站 浏览量
CAQL 连续动作 Q-Learning 实践指南:面向 max-Q 问题的可插拔优化器设计与源码解析 CAQL 连续动作 Q-Learning 实践指南面向 max-Q 问题的可插拔优化器设计与源码解析【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-researchCAQLContinuous Action Q-Learning是 google-research 仓库 caql/ 目录中开源的一类连续动作空间 Q-Learning 算法。其核心设计是将 DQN 中难以处理的对动作求最大值max-Q问题解耦为可插拔的优化器组件并配套双目标dual过滤、聚类加速等技巧。本文基于仓库源码与测试完整讲解 CAQL 的算法原理、网络架构、全部可调超参数、运行方式与实现细节读者可据此复现训练流程、理解每个配置项在代码中的实际作用并在此基础上替换或扩展自己的 max-Q 求解器。一、背景连续动作下的 max-Q 难题标准 DQN 通过argmax_a Q(s, a)选择动作当动作空间是离散且有限时这个 argmax 可以暴力枚举。但面对连续动作空间例如机器人关节力矩、自动驾驶转向角argmax_a Q(s, a)变为一个连续优化问题——即论文与代码中反复出现的max-Q 问题maxq problem。CAQL 的思路不是绕开这个问题而是把它显式地建模成一个内层优化子问题并允许研究者像插拔模块一样选用不同的优化器来求解。README 明确指出CAQL is a class of algorithms for continuous-action Q-learning that can use several plug-and-play optimizers for the max-Q problem.即CAQL 是一类算法族其关键特征是max-Q 问题可选用多种即插即用优化器。仓库中实现的优化器包括梯度上升gradient ascent、交叉熵方法cross entropy以及基于对偶松弛的解析近似dual / dual-IBP论文中的 MIP混合整数规划求解器因属于 Google 内部库、尚未开源README 已明确说明暂不包含但计划后续随开源更新补充。二、整体架构Agent、Network 与 Policy 三层设计从源码结构看CAQL 的实现分为三大模块职责边界清晰模块文件职责训练/评估入口caql/train_eval.py解析全部命令行参数编排数据采集、训练、评估、checkpoint 流程Agent 层caql/caql_agent.pyCaqlAgent管理 train/target 网络、Q 函数/动作函数/λ 函数的训练、双目标过滤与聚类加速、动态容差计算网络层caql/caql_network.pyCaqlNetQ 网络、动作函数网络、λ 函数网络的具体 TF 计算图以及内层求解器梯度上升 / 交叉熵 / dual的实现对偶方法caql/dual_method.py、caql/dual_ibp_method.py基于拉格朗日对偶duality与区间传播IBP的 max-Q 上界解析近似策略层caql/agent_policy.py、caql/epsilon_greedy_policy.py、caql/gaussian_noise_policy.py、caql/policy.py贪心策略与两类探索策略ε-greedy、高斯噪声经验回放caql/replay_memory.py线程安全、可 checkpoint/恢复的ReplayMemory工具函数caql/utils.py环境创建、经验并行采集、评估、超参保存、动作投影等启动脚本caql/run.sh环境准备与最小训练运行示例训练主循环位于 caql/train_eval.py 的main()每个 iteration 先按行为策略采集经验写入回放池然后循环train_steps_per_iteration次每次采样 minibatch依次调用agent.train_q_function_network(...)更新 Q 函数与 λ 函数、agent.train_action_function_network(...)更新动作函数并按target_update_steps/hard_update_steps同步 target 网络。三、Q 函数的两种损失L2 损失与 λ hinge 损失CaqlNet在 caql/caql_network.py 中定义了两种 Q 函数训练损失由l2_loss_flag切换L2 损失l2_loss_flagTrue默认直接对 TD 目标做均方误差回归即tf.losses.mean_squared_error(true_label, q_function_network)作为基准比较方案。λ hinge 损失simple_lambda_flagTrue且l2_loss_flagFalseQ λ * max(0, y - Q)其中 λ 是惩罚系数配对的lambda_function_loss为-λ * (y - Q)。当simple_lambda_flagTrue时λ 是一个可训练标量变量被投影到[0, lambda_max]区间内默认lambda_max5e3见 caql/caql_network.py 中lambda_proj相关代码否则 λ 由单独的 λ 函数网络结构与 Q 网络相同的 MLP逐样本输出。三个网络使用独立的 Adam 优化器但共享同一learning_rate动作函数另有learning_rate_action梯度上升求解器另有learning_rate_ga。四、核心创新max-Q 问题的可插拔求解器这是 CAQL 的灵魂所在也是 README 所述plug-and-play optimizers的落地实现。内层求解器通过solver参数选择支持的取值在 caql/train_eval.py 的 flag 定义中枚举为[dual, gradient_ascent, cross_entropy, ails, mip]。CaqlAgent会将其传给train_network由compute_best_actions()分发到具体实现caql/caql_network.py。1. 梯度上升gradient_ascent默认求解器。在 caql/caql_network.py 的_create_gradient_ascent_action_tensors()中构造可训练的action_variable_tensor以负 Q 值均值cost_now -mean(Q(s, a))为优化目标最大化 Q 等价于最小化 -Q每步计算action_gradient并做归一化grad / (eps ||grad||)用梯度上升更新动作变量迭代上限为action_maximization_iterations默认 20终止条件为前后两次 Q 值变化小于tolerance容差可选的sufficient_ascent_flag开启 Armijo-Goldstein 充分下降条件搜索源码中c_armijo0.01、c_goldstein0.25、lr_decay0.1、初始学习率lr_init100.0更新完成后通过_action_projection()将动作投影回[action_min, action_max]界内。2. 交叉熵方法cross_entropy在_create_cross_entropy_action_tensors()中实现使用tfp.distributions.MultivariateNormalDiag每轮采样num_samples200个动作保留 toptop_k_portion0.5的高 Q 样本用其均值与方差更新下一轮的采样分布循环直至收敛或达到迭代上限。与梯度上升一样采样后的动作也会投影回动作边界内。3. dual对偶松弛解析求解当solverdual时CAQL 不再做迭代优化而是直接利用对偶理论duality为 ReLU 网络的 max-Q 计算一个解析上界作为 max-Q 标签。实现位于 caql/dual_method.py 的create_dual_approx()通过前向传播逐层计算神经元上下界(l, u)、区分 active/inactive/spanning ReLU 单元get_I/get_D再反向递推构造对偶变量Nu_i与gamma_i最终得到J_tilde上界。仓库还提供了 IBP区间边界传播变体 caql/dual_ibp_method.py 的create_dual_ibp_approx()CaqlNet._create_dual_maxq_label_tensor()支持duality_based、ibp以及两者取tf.minimum的混合模式。需要特别说明dual 求解器给出的是 max-Q 的上界估计而非精确解代码注释也强调dual methods for approximating target_next_values。4. 未包含的 ails 与 mip尽管 flag 枚举中包含ails和mip但 caql/caql_network.py 的__init__和compute_best_actions()中均会抛出ValueError(AILS and MIP solvers are not supported yet.)。这与 README 的说明一致——MIP 求解器依赖尚未开源的 Google 内部库。若读者自行实现这两个求解器只需保证在compute_best_actions()分支中返回 shape 为(batch_size, action_dim)的动作张量即可接入现有框架。五、训练数据的高效利用dual filter 与聚类加速为了减少昂贵的 max-Q 求解次数caql/caql_agent.py 的train_q_function_network()实现了两个关键技巧README 论文中的 Trick 1 与 Trick 2由dual_filter_clustering_flag开关Dual filter对偶过滤先用 dual 方法快速估计 target 值再通过compute_dual_active_constraint_condition()判断哪些样本的 TD 备份大于当前 Q 预测即活跃约束只有这些样本才需要精确求解 max-Q其余样本直接使用 dual 上界。代码通过portion_active_data统计活跃样本占比并写入日志与 TensorBoard。Farthest-first 聚类对活跃的 next states 执行compute_cluster_masks()FF-traversal 聚类算法eps_approx_ratio0.01控制聚类半径只对簇心cluster centroid状态精确求解 max-Q对非簇心状态用一阶泰勒展开从最近簇心状态外推 Q 值predict_state_perturbed_q_function即Q ∇_s Q · Δs并以 dual 近似值与其取minimum来抑制过估计误差源码注释A numerical trick to reduce over-estimation error。train_q_function_network()的返回值中最后两项正是portion_active_data与portion_active_data_and_cluster分别对应过滤后与过滤聚类后的活跃数据占比用于评估这两个技巧的加速效果。六、动态容差tolerance机制max-Q 求解器的终止容差并非固定值。caql/caql_agent.py 的_compute_tolerance()实现了一种随 TD 误差自适应缩放的容差若传入tolerance_decay则tolerance tolerance_init * td_rmse * tolerance_decay并裁剪到[tolerance_min, tolerance_max]区间td_rmse由compute_td_rmse()计算否则直接返回tolerance_init缺省时取tolerance_min。对应 caql/train_eval.py 中的四个 flagtolerance_init初始值默认 None、tolerance_min默认 1e-4、tolerance_max默认 100.0、tolerance_decay衰减率默认 None。测试 caql/caql_agent_test.py 的testComputeTolerance验证了tolerance 0.1 * 0.2 * 0.9 0.018的计算过程。训练早期 TD 误差大容差被放大、求解更快随着训练收敛误差变小容差收紧、求解更精确。此外warmstart开关默认 True让求解器以上一动作函数网络的预测作为迭代初值进一步减少迭代次数。七、动作函数学习 argmax 的捷径为了避免每步决策都运行一次内层优化器CAQL 额外训练了一个动作函数网络action functiona f_θ(s)直接回归给定状态下使 Q 最大的动作。其训练损失是MSE(Q(s, f_θ(s)), best_q_label)其中标签best_q_label由dual_q_label控制dual_q_labelTrue默认用 dual 方法计算的 max-Q 上界作标签dual_q_labelFalse用 primal 的 max-Q 解即内层求解器实际求出的值作标签。推理时CaqlAgent.best_action()默认直接调用predict_action_function()输出动作use_action_functionTrue从而把内层优化器的开销完全转移到训练阶段。动作函数网络同样将输出投影到动作边界内caql/caql_network.py 的_build_action_function_net()。八、完整超参一览以下参数全部来自 caql/train_eval.py 的 absl flag 定义可在命令行直接覆盖环境与任务参数默认值说明--env_namePendulum支持Pendulum、Hopper、Walker2D、HalfCheetah、Ant、Humanoid见 caql/utils.py 的create_env()前一个为 gym其余为 tf-agents suite_mujoco 环境--discount_factor0.99折扣因子 γ--time_out200单条 episode 最大步数--action_boundsNone逗号分隔的 min,max如-.5,.5不设置则使用环境原生动作界--seed0随机种子网络与训练参数默认值说明--hidden_layers32,16各隐藏层单元数逗号分隔--batch_size64minibatch 大小--learning_rate0.001Q 函数 / λ 函数学习率--learning_rate_action0.005动作函数学习率--learning_rate_ga0.01梯度上升求解器学习率仅该求解器生效--action_maximization_iterations20内层 max-Q 求解迭代次数上限--replay_memory_capacity100000回放池容量须大于batch_size × train_steps_per_iteration--train_steps_per_iteration20每个 iteration 内的梯度步数--target_update_steps1每 N 步软更新 target 网络算法选项参数默认值说明--solvergradient_ascentdual/gradient_ascent/cross_entropy/ails/mip后两者未实现--tau_copy0.001target 网络软更新比例tau1时 polyak 插值否则直接硬拷贝--clipped_targetTrue是否启用 clipped double DQN再维护第二个 target 网络取 min--hard_update_steps5000第二个 target 网络的硬更新周期仅 clipped 模式--l2_loss_flagTrue使用 L2 损失否则使用 λ hinge 损失--simple_lambda_flagFalse使用可训练标量 λ需配合 hinge 损失--initial_lambda1.0hinge 损失的 λ 初值--dual_filter_clustering_flagFalse启用 dual filter 与聚类加速--dual_q_labelTrue动作函数用 dual max-Q 标签否则用 primal 标签--warmstartTruemax-Q 求解是否以动作函数预测为初值--tolerance_init/--tolerance_min/--tolerance_max/--tolerance_decayNone/1e-4/100.0/Nonemax-Q 求解容差及其衰减探索策略参数默认值说明--exploration_policygaussianegreedy/gaussian/none--epsilon/--epsilon_decay/--epsilon_min1.0/0.999/0.025ε-greedy 参数caql/epsilon_greedy_policy.py--sigma/--sigma_decay/--sigma_min1.0/0.999/0.025高斯噪声参数caql/gaussian_noise_policy.py噪声按sigma × action_max × N(0,1)缩放流程控制参数默认值说明--max_iterations10000最大迭代数--num_episodes_per_iteration1每 iteration 采集的 episode 数--collect_experience_parallelism1并行采集线程数--checkpoint_dir/--result_dirNone模型 checkpoint / 结果目录--checkpoint_iterations/--eval_iterations/--num_evals50/50/10定义于 caql/utils.pycheckpoint 周期、评估周期与每次评估的 episode 数九、环境准备与运行仓库提供了一键脚本 caql/run.sh内容为创建 Python 3 虚拟环境、安装依赖并以最小迭代数运行set -e set -x virtualenv -p python3 venv source ./venv/bin/activate pip install tensorflow pip install tf-agents pip install gym python -m caql.train_eval --max_iterations3注意 caql/run.sh 使用的Pendulum-v0/-v2系列 gym 版本较老当前环境运行时可能需要安装对应版本的 gym 并处理兼容性。代码本身面向 TensorFlow 1.x 计算图tensorflow.compat.v1并调用tf.disable_v2_behavior()这是运行时的前提条件。在仓库根目录下也可以直接指定更完整的训练命令例如python -m caql.train_eval \ --env_namePendulum \ --solvergradient_ascent \ --hidden_layers32,16 \ --batch_size64 \ --learning_rate0.001 \ --learning_rate_action0.005 \ --learning_rate_ga0.01 \ --action_maximization_iterations20 \ --tau_copy0.001 \ --clipped_targetTrue \ --hard_update_steps5000 \ --l2_loss_flagTrue \ --dual_q_labelTrue \ --tolerance_min1e-4 \ --tolerance_max100.0 \ --warmstartTrue \ --exploration_policygaussian \ --sigma1.0 --sigma_decay0.999 --sigma_min0.025 \ --max_iterations10000 \ --checkpoint_dir/tmp/caql_ckpt \ --result_dir/tmp/caql_result运行时日志会输出每个 iteration 的avg_q_function loss、avg_lambda_function loss、avg_action_function loss、avg portion active data、avg portion active data and cluster等指标caql/train_eval.py 主循环末尾便于监控 dual filter 与聚类的实际加速比例。十、数据采集、评估与持久化经验采集utils.collect_experience_parallel()支持多线程并行生成 episode线程数由collect_experience_parallelism控制每条经验形如[state, action, reward, next_state, done, info]写入线程安全的ReplayMemory内部为collections.deque(maxlencapacity)threading.Lock。评估utils.periodic_updates()每eval_iterations次迭代启动num_evals个并行评估任务用贪心策略跑完整 episode统计avg_score、avg_episode_len、avg_action_magnitude动作无穷范数均值并写入 TensorBoard。持久化每checkpoint_iterations次迭代保存模型model.ckpt-step与回放池replay_memory-*.pkl-n支持删除旧文件。训练中断后重新启动时CaqlAgent.initialize()会从checkpoint_dir恢复权重并返回恢复后的 global step训练从断点继续——caql/caql_agent_test.py 的testInitializeWithCheckpoint用 17 步训练后保存、再恢复并断言 step 回到 17完整验证了这一流程。超参记录指定--result_dir时训练开始前会把全部超参以 pickle 形式写入result_dir/hparam.pickle。十一、代码模块速查caql/caql_agent.pyCaqlAgent——双目标过滤、FF 聚类、容差计算、target 软/硬更新、Q 与动作函数的训练编排。caql/caql_network.pyCaqlNet——Q / 动作 / λ 网络构建、梯度上升与交叉熵求解器、dual max-Q 标签、TD 备份与 TD RMSE。caql/dual_method.py 与 caql/dual_ibp_method.py对偶松弛与 IBP 区间传播两种 max-Q 上界计算。caql/train_eval.py全部 CLI 参数与训练主循环。caql/utils.py环境工厂、并行采集、评估、checkpoint、超参保存。caql/replay_memory.py可 checkpoint 的线程安全回放池。caql/agent_policy.py、caql/epsilon_greedy_policy.py、caql/gaussian_noise_policy.py贪心与探索策略。测试caql/*_test.py覆盖 Agent 初始化/断点恢复/聚类/容差caql/caql_agent_test.py、网络前向与求解器caql/caql_network_test.py、对偶方法caql/dual_method_test.py、caql/dual_ibp_method_test.py、策略、回放池与工具函数可作为自定义扩展的回归基准。十二、使用限制与扩展建议MIP/AILS 暂不可用solver参数枚举中的ails与mip目前只会在代码中抛出NotImplementedError风格的ValueError因为 MIP 优化器依赖 Google 内部库尚未开源README 已声明。TF1 计算图全部代码基于tensorflow.compat.v1的 Session 式编程并显式tf.disable_v2_behavior()迁移到 TF2 eager 或新框架如 JAX/PyTorch需要整体改造。扩展新求解器从源码结构看接入新的 max-Q 优化器只需在CaqlNet中仿照_create_gradient_ascent_action_tensors()构造求解循环并在compute_best_actions()中增加分发分支返回投影后的最优动作即可体现了 README 所强调的 plug-and-play 设计。对偶标签是上界dual 求解器与dual_q_labelTrue给出的是 max-Q 的解析上界适用于快速过滤与标签近似追求精确解时应结合 primal 求解器并合理设置tolerance与warmstart。综上CAQL 以max-Q 求解器可插拔为核心设计将连续动作 Q-learning 的难点显式化、模块化并通过 dual filter、聚类外推、动态容差、动作函数蒸馏等手段把内层优化的开销控制在可接受范围。本文所有参数与行为均可在 caql/train_eval.py 与 caql/caql_agent.py 等源码中直接验证是深入理解与二次开发该算法的最佳起点。【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考