多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

Ray RLlib 中的 Soft Actor-Critic(SAC):最大熵离策略强化学习的实现与实战指南

Ray RLlib 中的 Soft Actor-Critic(SAC):最大熵离策略强化学习的实现与实战指南 人工智能分布式训练强化学习任务调度模型推理服务【免费下载链接】rayRay is an AI compute engine. Ray consists of a core distributed runtime and a set of AI Libraries for accelerating ML workloads.项目地址https://gitcode.com/gh_mirrors/ra/ray点击查看免费下载导读SACSoft Actor-Critic是 Ray RLlib 内置的一种基于最大熵框架的 model-free、off-policy 强化学习算法在连续控制任务上表现优异。本文以 rllib/algorithms/sac/README.md 为核心结合 sac.py、sac_learner.py、torch/sac_torch_learner.py 等源码系统讲解 SAC 的核心原理、配置参数、源码实现与实战训练方法并覆盖 SAC-Discrete 离散动作变体的使用。读完后你将掌握如何在 RLlib 中用SACConfig训练连续控制与离散控制任务理解最大熵、自动温度调节alpha、twin Q 等关键机制并知道如何修改参数与调试。SAC 算法概览最大熵强化学习框架SAC论文是一种 SOTA 的 model-free、off-policy 强化学习算法在连续控制领域表现突出。它采用 actor-critic 架构通过基于最大熵框架的学习解决高样本复杂度和训练稳定性的问题。与标准 RL 目标最大化未来累积奖励之和不同SAC 的目标是在最大化累积奖励的同时最大化当前策略的期望熵。即目标函数为J(π) Σ_t E[(s_t, a_t) ~ ρ_π][ r(s_t, a_t) α * H(π(·|s_t)) ]其中αalpha为熵权重系数。除了使用基于熵目标的 actor 和 critic 进行优化外SAC 还会对熵系数本身进行优化实现自动温度调节automatic temperature tuning。SAC-Discrete 变体SAC-Discrete论文 中可以看到对于离散动作空间Q 网络输出每个动作的 Q 值输出维度为action_space.n对于连续动作空间Q 网络输出单个值。在 torch/sac_torch_learner.py 中compute_loss_for_module会根据动作空间类型Discrete或Box自动选择离散或连续的损失函数。源码中的核心实现SAC 在 RLlib 中的实现文件如下文件职责sac.py定义SACConfig配置类与SAC算法类sac_learner.py定义SACLearner负责 alpha 与 target entropy 的初始化/管理torch/sac_torch_learner.pyTorch 专属的 SAC 损失计算离散/连续与优化器配置default_sac_rl_module.py默认 RLModule 结构pi 编码器/头、Q 网络与目标网络torch/default_sac_torch_rl_module.pyTorch 前向传播实现sac_catalog.py模型构建目录encoder/head 组装sac_tf_policy.py / sac_torch_policy.py旧 API stack 的 Policy 实现init.py导出SAC、SACConfig、SACTFPolicy、SACTorchPolicySAC 类继承关系SAC类继承自DQNsac.py#L561-L570其默认配置由SACConfig()提供默认策略类根据 framework 选择SACTorchPolicy或SACTFPolicy。RLModule 结构default_sac_rl_module.py 说明了默认 RLModule 的网络结构策略actorpi_encoder状态编码器pi_head输出动作分布参数如 squashed Gaussian 的位置和对数尺度参数Q 网络criticqf_encoderqf_head输入为[obs, action]输出软 Q 值目标网络qf_target_encoderqf_target_head若twin_qTrue则额外有qf_twin_*与qf_target_twin_*。其前向传播链路可概括为[obs] - [pi_encoder] - [pi_head] - [action_dist_inputs] [obs, action] - [qf_encoder] - [qf_head] - [q-value] [obs, action] - [qf_target_encoder] - [qf_target_head] - [q-target-value]SACConfig 配置详解SACConfig继承自AlgorithmConfigsac.py#L31-L50通过链式调用.training()、.env_runners()、.environment()等方法来配置。以下是核心配置参数及其默认值训练相关参数.training()参数默认值说明twin_qTrue是否使用两个 Q 网络Clipped Double-Q每个 Q 网络拥有自己的目标网络tau5e-3目标网络软更新系数target tau * policy (1 - tau) * target_policyinitial_alpha1.0熵权重 alpha 的初始值target_entropyauto目标熵下界为auto时自动设为-|A|如 Discrete(2) 为 -2.0Box(shape(3,)) 为 -3.0n_step1N 步目标更新可为整数或(min, max)元组每次采样时从闭区间均匀抽取clip_actionsFalse是否裁剪动作若动作已归一化应设为Falsegrad_clipNone梯度裁剪值须大于 0.0actor_lr3e-5actor 学习率或[[timestep, lr], ...]形式的调度默认比 critic 低一个数量级critic_lr3e-4critic 学习率默认比 actor 高一个数量级alpha_lr3e-4alpha 温度参数的学习率默认与 critic 相同target_network_update_freq0目标网络更新频率每 N 步更新一次store_buffer_in_checkpointsFalse是否将回放缓冲区内容存入 checkpointtraining_intensityNone模型更新相对采样执行的强度None时使用自然值train_batch_size / (rollout_fragment_length * num_env_runners * num_envs_per_env_runner)num_steps_sampled_before_learning_starts1500开始学习前从 runner 收集的 timestep 数模型配置q_model_config与policy_model_config默认均为{ fcnet_hiddens: [256, 256], fcnet_activation: relu, post_fcnet_hiddens: [], post_fcnet_activation: None, custom_model: None, # 自定义 Q 模型 / 策略模型 custom_model_config: {}, }其中q_model_config用于 Q 网络obsBox(1D)时Tuple(Box(1D) Action) - concat - post_fcnetobsBox(3D)时先经 vision-net 再与 action concatobsTuple(...)时同理。policy_model_config与 Q 配置的区别在于post_fcnet 之前不进行 action 拼接。回放缓冲区配置self.replay_buffer_config { type: PrioritizedEpisodeReplayBuffer, # 优先级回放缓冲区 capacity: int(1e6), # 容量若 async_updates 开启每个 worker 独立拥有该容量 alpha: 0.6, # 优先级指数TD 误差越高被采样的概率越大0.0 为均匀采样 beta: 0.4, # 重要性采样系数抑制高采样概率样本的梯度影响 }在replay_buffer_config中还可设置prioritized_replay_eps基线采样概率保证 TD 误差为 0 时仍有采样机会与replay_batch_size等。注意当使用新 API stackenable_env_runner_and_connector_v2True时回放缓冲区必须是EpisodeReplayBuffer系列类型sac.py#L431-L460使用旧 API stack 时则必须使用MultiAgentPrioritizedReplayBuffer等旧类型sac.py#L461-L480。其他关键配置rollout_fragment_length默认auto自动设为n_step的值或n_step[1]若为元组sac.py#L118-L120、sac.py#L500-L508train_batch_size_per_learner默认256新 API stacktrain_batch_size默认256旧 API stackoptimization{actor_learning_rate: 3e-4, critic_learning_rate: 3e-4, entropy_learning_rate: 3e-4}可通过optimization_config覆盖exploration_config默认{type: StochasticSampling}SAC 通过随机采样进行探索lr必须为NoneSAC 使用actor_lr/critic_lr/alpha_lr三个独立学习率sac.py#L482-L489。完整训练示例最小可用配置以下是 sac.py#L34-L49 中给出的官方最小示例在Pendulum-v1上训练 1 个迭代from ray.rllib.algorithms.sac import SACConfig config ( SACConfig() .environment(Pendulum-v1) .env_runners(num_env_runners1) .training( gamma0.9, actor_lr0.001, critic_lr0.002, train_batch_size_per_learner32, ) ) # 根据 config 构建 SAC 算法对象并运行 1 次训练迭代 algo config.build() algo.train()测试用例中的完整配置test_sac.py 展示了更完整的配置方式使用n_step3、twin_qTrue、store_buffer_in_checkpointsTrue等import ray from ray import tune from ray.rllib.algorithms import sac from ray.rllib.connectors.env_to_module.flatten_observations import FlattenObservations config ( sac.SACConfig() .training( n_step3, twin_qTrue, replay_buffer_config{ capacity: 40000, }, num_steps_sampled_before_learning_starts0, store_buffer_in_checkpointsTrue, train_batch_size10, ) .env_runners( env_to_module_connector( lambda env, spaces, device: FlattenObservations() ), num_env_runners0, rollout_fragment_length10, ) ) algo config.build() results algo.train() algo.stop()该测试还验证了 SAC 支持 Dict / Tuple 混合观测空间含图片、离散与连续分量的环境如random_dict_envDict({a: Box(3), b: Discrete(2), c: Box(84,84,3)})与random_tuple_envtest_sac.py#L84-L124。自定义环境的处理SAC 对自定义 Gymnasium 环境无特殊要求只需遵循标准reset/step接口。测试中使用的SimpleEnvtest_sac.py#L20-L43展示了最小环境写法action_spaceBox(0,1,(1,))、observation_spaceBox(0,1,(1,))奖励定义为1.0 - |max(action) - state|。可通过tune.register_env(name, lambda config: env)注册环境后传入config.environment(name)。源码级原理剖析连续动作空间的损失函数torch/sac_torch_learner.py 中的_compute_loss_for_module_continuous实现了连续控制场景的 SAC 损失Critic 损失软贝尔曼误差对应论文 eq.(7-8)q_target_next fwd_out[q_target_next] - alpha.detach() * fwd_out[logp_next_resampled] q_next_masked (1.0 - batch[Columns.TERMINATEDS].float()) * q_target_next q_selected_target ( batch[Columns.REWARDS] (config.gamma ** batch[n_step]) * q_next_masked ).detach() critic_loss torch.mean( batch[weights] * torch.nn.HuberLoss(reductionnone, delta1.0)( q_selected, q_selected_target ) )值得注意的是实现中使用Huber loss 替代 MSE以提升训练性能源码注释明确说明。Actor 损失最大化 Q 值与熵实现中取负号做最小化actor_loss torch.mean(alpha.detach() * fwd_out[logp_resampled] - fwd_out[q_curr])Alpha 损失自动温度调节alpha_loss -torch.mean( self.curr_log_alpha[module_id] * (fwd_out[logp_resampled].detach() self.target_entropy[module_id]) )其中 alpha 在对数域存储与优化curr_log_alpha见 sac_learner.py#L33-L46使用时通过alpha torch.exp(self.curr_log_alpha[module_id])还原sac_torch_learner.py#L279。重参数化技巧前向传播中使用rsample()而非sample()采样动作保证梯度可通过采样节点回传torch/default_sac_torch_rl_module.py#L153-L159。离散动作空间的损失函数_compute_loss_for_module_discretesac_torch_learner.py#L148-L268针对离散动作空间通过 softmax 得到动作概率与对数概率计算下一状态的价值next_v (action_probs_next * (next_q - alpha.detach() * action_log_probs_next)).sum(-1)critic 损失通过gather提取已选动作的 Q 值计算 TD 误差actor 损失为(action_probs * (alpha.detach() * action_log_probs - qf)).sum(-1).mean()。目标网络更新与优化器目标网络软更新tau控制软更新比例target_network_update_freq可控制硬更新的频率。四组优化器configure_optimizers_for_modulesac_torch_learner.py#L54-L113为每个模块注册qf、qf_twin若启用、policy、alpha四组 Adam 优化器分别使用critic_lr、actor_lr、alpha_lr。TD 误差与优先级回放critic 计算的 TD 误差含 twin Q 时取平均会以item_series方式记录sac_torch_learner.py#L229-L233供PrioritizedEpisodeReplayBuffer计算采样优先级。双 Q 网络Clipped Double-Qtwin_qTrue时actor 损失中使用两个 Q 网络的最小值qf torch.min(fwd_out[QF_PREDS], fwd_out[QF_TWIN_PREDS]).detach()连续场景见 sac_torch_learner.py#L339-L341。同时目标网络也会为 twin Q 建立对应副本default_sac_rl_module.py#L77-L111。调参建议与注意事项双时间尺度学习率源码注释明确建议 actor 学习率比 critic 低一个数量级默认actor_lr3e-5、critic_lr3e-4以保证 critic 能为策略提供可靠的价值估计。可参考 sac.py#L263-L298 的详细说明。rollout_fragment_length与n_step的约束当n_step为元组时rollout_fragment_length不得小于n_step[1]否则validate()会抛出ValueErrorsac.py#L392-L410。框架支持新 API stack 下 SAC 仅支持torch框架get_default_rl_module_spec与get_default_learner_class在非 torch 时抛错sac.py#L511-L532若在 TF 下运行需安装tensorflow_probabilitysac.py#L422-L429。探索策略SAC 默认使用StochasticSampling探索sac.py#L53-L61即直接从策略的随机分布采样动作无需 epsilon-greedy。checkpoint 与回放缓冲区如需恢复训练并保留回放数据设置store_buffer_in_checkpointsTrue恢复时若数据缺失会给出警告sac.py#L198-L204。调试选项_deterministic_lossTrue可跳过随机动作采样进行确定性损失计算仅对连续动作、用于调试_use_beta_distributionTrue可用 Beta 分布替代 SquashedGaussian不推荐仅调试用。如何运行与验证在仓库根目录下运行 SAC 的单元测试即可验证实现正确性python -m pytest rllib/algorithms/sac/tests/test_sac.py -v该测试覆盖了SAC 在混合观测空间环境上的构建与训练test_sac_compilation、Dict 观测空间键序稳定性test_sac_dict_obs_order。运行前需安装 Ray 及其 RLlib 依赖Torch 等。总结RLlib 中的 SAC 实现了完整的最大熵离策略训练框架两大目标最大化累积奖励 最大化策略熵三个可训练组件actor策略、criticQ 网络可启用 twin Q、温度系数 alpha自动调节两种动作空间连续SquashedGaussian 重参数化采样与离散SAC-Discrete 变体两类 API stack新 API stackSACLearnerDefaultSACTorchRLModule默认与旧 API stackSACTorchPolicy/SACTFPolicy。通过SACConfig的链式配置你可以在数行代码内完成从环境注册、回放缓冲区设置到训练启动的完整流程并根据任务特性微调twin_q、tau、initial_alpha、n_step与三组学习率等关键超参数。赞分享人工智能分布式训练强化学习任务调度模型推理服务【免费下载链接】rayRay is an AI compute engine. Ray consists of a core distributed runtime and a set of AI Libraries for accelerating ML workloads.项目地址https://gitcode.com/gh_mirrors/ra/ray点击查看免费下载相关推荐Soft Actor-Critic深度强化学习的未来Soft Actor Critic深度强化学习的未来 项目介绍 Soft Actor CriticSAC是一个用于连续域中训练最大熵策略的深度强化学习框架Stable Baselines3 中的 SAC 算法详解随机策略下的最大熵离线强化学习Stable Baselines3 中的 SAC 算法详解随机策略下的最大熵离线强化学习 本文基于 docs/modules/sac.md https://l人工智能强化学习机器学习深入解析garage中的Soft Actor-Critic(SAC)算法深入解析garage中的Soft Actor Critic SAC 算法 算法概述 Soft Actor CriticSAC是一种在强化学习领域中表现优异的上一篇GitHub Readme Streak Stats性能优化如何减少加载时间下一篇如何突破网络边界ZeroTierOne的OpenTelemetry分布式监控完整指南创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表