SAC算法:连续动作空间强化学习的原理与实现 1. SAC算法核心原理与实现思路在强化学习领域连续动作空间的控制一直是个颇具挑战性的问题。传统的DQN等算法只能处理离散动作而像Pendulum-v1这样的环境需要输出连续扭矩值范围[-2,2]。Soft Actor-CriticSAC算法通过引入熵正则化机制在最大化累积奖励的同时鼓励动作探索成为解决这类问题的利器。1.1 最大熵强化学习框架SAC的核心创新在于其优化目标 $$ \pi^* \arg\max_\pi \mathbb{E}_{\tau\sim\pi}\left[\sum_t r(s_t,a_t) \alpha H(\pi(\cdot|s_t))\right] $$ 其中α是温度系数H(π(·|s))是策略熵。这个公式意味着算法不仅要追求高奖励还要保持策略的随机性。在实际实现中我发现α的取值非常关键——太小会导致探索不足太大又会影响策略收敛。经过多次实验最终采用自动调整α的机制将其初始值设为0.2目标熵设为动作维度的负数Pendulum-v1中为-1。1.2 关键技术实现要点针对连续动作空间SAC有几个精妙设计重参数化技巧策略网络输出高斯分布的μ和σ通过$\epsilon \sim \mathcal{N}(0,1)$采样计算$a \tanh(\mu \sigma \odot \epsilon)$。这种参数化方式使得采样过程可导便于梯度回传。双Q网络结构使用两个独立的Q网络取较小值作为目标有效缓解Q值高估问题。在代码中可以看到critic_1和critic_2的并行结构。目标网络软更新通过参数τ控制更新幅度代码中设为0.005使目标网络缓慢跟踪当前网络提升训练稳定性。提示在实际编码时注意tanh变换后的概率密度修正。由于tanh是非线性变换需要对应调整对数概率log_prob - torch.log(1 - torch.tanh(action).pow(2) 1e-7)2. 代码架构与核心模块实现2.1 网络结构设计2.1.1 策略网络(PolicyNetContinuous)class PolicyNetContinuous(torch.nn.Module): def __init__(self, state_dim, hidden_dim, action_dim, action_bound): super().__init__() self.fc1 nn.Linear(state_dim, hidden_dim) self.fc_mu nn.Linear(hidden_dim, action_dim) self.fc_std nn.Linear(hidden_dim, action_dim) self.action_bound action_bound def forward(self, x): x F.relu(self.fc1(x)) mu self.fc_mu(x) std F.softplus(self.fc_std(x)) # 保证标准差为正 dist Normal(mu, std) normal_sample dist.rsample() log_prob dist.log_prob(normal_sample) action torch.tanh(normal_sample) # 概率密度修正 log_prob - torch.log(1 - torch.tanh(action).pow(2) 1e-7) return action * self.action_bound, log_prob2.1.2 Q值网络(QValueNetContinuous)采用两层隐藏层的MLP结构输入为状态和动作的拼接class QValueNetContinuous(torch.nn.Module): def __init__(self, state_dim, hidden_dim, action_dim): super().__init__() self.fc1 nn.Linear(state_dim action_dim, hidden_dim) self.fc2 nn.Linear(hidden_dim, hidden_dim) self.fc_out nn.Linear(hidden_dim, 1) def forward(self, x, a): x F.relu(self.fc1(torch.cat([x, a], dim1))) x F.relu(self.fc2(x)) return self.fc_out(x)2.2 经验回放机制实现了一个循环缓冲区的经验回放池class ReplayBuffer: def __init__(self, capacity): self.buffer collections.deque(maxlencapacity) # 使用deque实现循环缓冲区 def push(self, state, action, reward, next_state, done): self.buffer.append((state, action, reward, next_state, done)) def sample(self, batch_size): transitions random.sample(self.buffer, batch_size) return map(np.array, zip(*transitions)) def __len__(self): return len(self.buffer)在实际使用中发现当环境奖励尺度变化较大时如Pendulum-v1的原始奖励范围是[-16.2,0]对奖励进行归一化可以显著提升训练稳定性。代码中采用了(rewards 8.0)/8.0将奖励映射到[0,1]附近。3. 完整训练流程与调优技巧3.1 训练循环实现训练过程采用经典的离线策略(off-policy)模式智能体与环境交互收集经验从经验池随机采样batch数据计算TD目标并更新网络参数关键训练代码如下def train_off_policy_agent(env, agent, num_episodes, replay_buffer, minimal_size, batch_size): return_list [] for i_episode in range(num_episodes): state, _ env.reset() episode_return 0 done False while not done: action agent.take_action(state) next_state, reward, done, _ env.step(action) replay_buffer.push(state, action, reward, next_state, done) state next_state episode_return reward if len(replay_buffer) minimal_size: b_s, b_a, b_r, b_ns, b_d replay_buffer.sample(batch_size) transition_dict { states: b_s, actions: b_a, next_states: b_ns, rewards: b_r, dones: b_d } agent.update(transition_dict) return_list.append(episode_return) return return_list3.2 关键参数设置经过多次实验验证以下参数组合在Pendulum-v1上表现良好参数推荐值作用说明actor_lr3e-4策略网络学习率critic_lr3e-3Q网络学习率alpha_lr3e-4温度系数学习率gamma0.99折扣因子tau0.005软更新系数buffer_size100000经验池容量batch_size64训练batch大小hidden_dim128网络隐藏层维度注意学习率的设置非常关键。实践中发现策略网络的学习率应该小于Q网络因为策略更新对Q值的估计误差更敏感。如果策略学习太快容易导致训练不稳定。3.3 可视化与调试技巧实时渲染分离创建独立的环境实例用于渲染避免拖慢训练速度env gym.make(Pendulum-v1) # 训练用 env_render gym.make(Pendulum-v1, render_modehuman) # 渲染用滑动平均曲线使用窗口大小为9的滑动平均处理训练曲线更清晰观察趋势def moving_average(a, window_size): cumulative_sum np.cumsum(np.insert(a, 0, 0)) middle (cumulative_sum[window_size:] - cumulative_sum[:-window_size]) / window_size return middle训练过程监控通过tqdm进度条实时显示平均回报方便判断收敛情况4. 常见问题与解决方案4.1 训练不收敛问题排查奖励尺度异常现象Q值爆炸式增长或变为NaN解决检查环境奖励范围必要时进行缩放如Pendulum-v1的奖励重塑策略熵失控现象温度系数α持续增大或减小解决调整目标熵值检查策略网络输出是否合理梯度爆炸现象网络参数突然变为NaN解决添加梯度裁剪减小学习率4.2 性能优化建议并行数据收集使用多个环境实例并行采样加快经验收集速度自动α调整实现动态调整温度系数避免手动调参优先经验回放对重要的transition赋予更高采样概率定期保存模型保存训练过程中的检查点防止意外中断4.3 迁移到其他环境当将本实现迁移到其他连续控制环境如MuJoCo系列时需要注意调整动作边界action_bound匹配新环境的动作空间范围可能需要增大网络容量如hidden_dim设为256或512对于高维状态输入如图像需要考虑使用CNN提取特征更复杂的环境通常需要更大的经验池和更长的训练时间我在实际项目中发现这套SAC实现稍作修改就能在HalfCheetah-v4上取得不错的效果关键调整包括增大batch_size到256延长训练到5000回合以及使用更深的Q网络3层隐藏层。

本月热点