相比普通 SACAgent 的关键差异

发布时间:2026/7/22 6:56:03
相比普通 SACAgent 的关键差异 普通 SACAgent 的 actor 输出完整动作critic 接收完整动作。但 SACAgentHybridSingleArm 做了三层动作拆分第一层连续动作由 SAC actor 输出初始化 policy 时action_dim 被设为环境动作维度减 1policy_def Policy(…action_dimactions.shape[-1]-1, # 7 - 1 6)如果环境 action 是 7 维actor 只输出前 6 维——末端执行器的连续控制量。第二层普通 critic 也只评估连续动作Critic 初始化时只用去掉夹爪的动作critic [observations, actions[…, :-1]]训练 critic 时也只取前 6 维actions batch[“actions”][…, :-1]这意味着普通 SAC critic 学习的是Qee(s,aee)而不是Q(s,aee,agripper)。它专注于连续末端执行器动作的价值估计。第三层夹爪动作由 GraspCritic 单独学习SACAgentHybridSingleArm 在网络集合里额外注册了 “grasp_critic”并给它单独配置 optimizernetworks {“actor”: actor_def,“critic”: critic_def,“grasp_critic”: grasp_critic_def,“temperature”: temperature_def,}“grasp_critic”: make_optimizer(**grasp_critic_optimizer_kwargs)在 create_pixels 中GraspCritic 的构造方式如下grasp_critic_def partial(GraspCritic, encoderencoders[“grasp_critic”], networkgrasp_critic_backbone)(name“grasp_critic”)6.3 动作执行流程Rolloutsample_actions 是 SAC actor 和 GraspCritic 配合最直观的接口。执行流程分为四步第一步actor 采样连续末端执行器动作dist self.forward_policy(observations, rngseed, trainFalse)ee_actions dist.sample(seedseed)这里得到前 6 维动作[x, y, z, roll, pitch, yaw]。第二步GraspCritic 输出 3 个夹爪 Q 值grasp_q_values self.forward_grasp_critic(observations, rnggrasp_key, trainFalse)输出类似Q(close) 1.2Q(keep) 0.7Q(open) 0.1第三步argmax 选择夹爪动作grasp_action grasp_q_values.argmax(axis-1)grasp_action grasp_action - 1 # {0,1,2} → {-1,0,1}映射关系argmax0 → 环境动作 -1关argmax1 → 环境动作 0保持argmax2 → 环境动作 1开。第四步拼接成完整动作return jnp.concatenate([ee_actions, grasp_action[…, None]], axis-1)最终输出[continuous_0, continuous_1, continuous_2,continuous_3, continuous_4, continuous_5,discrete_gripper_action] ← 环境真正需要的完整 7 维动作连续控制分支的 Rollout 符合经典 SAC 范式——只依赖 Actor 网络Critic 不参与推理。离散控制分支则不同——由于没有独立的 Actor 网络GraspCritic 在推理时直接充当决策者通过 argmax 选出最优离散动作。6.4 训练流程SACAgentHybridSingleArm.loss_fns 返回四个独立的 lossreturn {“critic”: self.critic_loss_fn,“grasp_critic”: self.grasp_critic_loss_fn,“actor”: self.policy_loss_fn,“temperature”: self.temperature_loss_fn,}训练脚本中hybrid agent 的更新分为两个阶段Critic 训练阶段同时更新 critic 和 grasp_critictrain_critic_networks_to_update frozenset({“critic”, “grasp_critic”})完整训练阶段更新全部四个网络train_networks_to_update frozenset({“critic”, “grasp_critic”, “actor”, “temperature”})Critic Loss连续动作 SAC目标 Q 采用 Clipped Double-Qyrγ⋅miniQ¯θi(s′,a′)代码实现target_next_qs self.forward_target_critic(…)target_next_min_q target_next_qs.min(axis0)target_q rewards discount * masks * target_next_min_qpredicted_qs self.forward_critic(batch[“observations”], actions, …)critic_loss jnp.mean((predicted_qs - target_qs) ** 2)GraspCritic LossDQN 风格GraspCritic 采用 Double DQN 风格——用 online 网络选动作用 target 网络评估next_grasp_qs self.forward_grasp_critic(batch[“next_observations”], rngrng)best_next_grasp_action next_grasp_qs.argmax(axis-1)target_next_grasp_qs self.forward_target_grasp_critic(…)target_next_grasp_q target_next_grasp_qs[jnp.arange(batch_size), best_next_grasp_action]grasp_rewards batch[“rewards”] batch[“grasp_penalty”]target_grasp_q grasp_rewards discount * masks * target_next_grasp_qpredicted_grasp_q predicted_grasp_qs[jnp.arange(batch_size), grasp_action]grasp_critic_loss jnp.mean((predicted_grasp_q - target_grasp_q) ** 2)Actor Loss标准 SACActor 采样连续动作最大化objectiveQ(s,a)−αlogπ(a|s)actor_objective predicted_q - temperature * log_probsactor_loss -jnp.mean(actor_objective)Temperature Loss自动调温温度参数α自动调节保证策略熵不低于目标值——策略熵低于 target 则α增大并增加探索高于 target 则α减小并减少探索。6.5 Reward 的分工设计这是混合 agent 中最精妙的工程细节之一。普通 SAC critic 的目标使用batch[“rewards”]GraspCritic 的目标使用grasp_rewards batch[“rewards”] batch[“grasp_penalty”]设计者的意图很明确夹爪网络不仅要学习任务成功奖励还要特别学习不要做无意义夹爪动作。例如 USB pickup insertion 任务中grasp_penalty 的计算逻辑是if (action[-1] -0.5 and self.last_gripper_pos 0.9) or (action[-1] 0.5 and self.last_gripper_pos 0.9):info[“grasp_penalty”] self.penaltyelse:info[“grasp_penalty”] 0.0通俗理解夹爪已经开得很大还继续开或者已经关得很紧还继续关——这种动作没有实际意义甚至可能损坏硬件或扰乱任务。这类惩罚不适合影响机械臂连续运动的 SAC critic——因为 SAC critic 评估的是 6 维连续动作的价值加上夹爪惩罚会混淆它对末端执行器动作质量的判断。但 GraspCritic 专门学习夹爪的离散决策加上 grasp_penalty 可以更精准地训练什么时候该夹、什么时候该放的策略。6.7 配合全景图