相比普通 SACAgent 的关键差异

📅 2026/7/20 18:09:41 👁️ 阅读次数 📝 编程学习
相比普通 SACAgent 的关键差异

普通 SACAgent 的 actor 输出完整动作,critic 接收完整动作。但 SACAgentHybridSingleArm 做了三层动作拆分:

第一层:连续动作由 SAC actor 输出

初始化 policy 时,action_dim 被设为环境动作维度减 1:

policy_def = Policy(

action_dim=actions.shape[-1]-1, # 7 - 1 = 6
)

如果环境 action 是 7 维,actor 只输出前 6 维——末端执行器的连续控制量。

第二层:普通 critic 也只评估连续动作

Critic 初始化时只用去掉夹爪的动作:

critic = [observations, actions[…, :-1]]

训练 critic 时也只取前 6 维:

actions = batch[“actions”][…, :-1]

这意味着普通 SAC critic 学习的是
Q
ee
(
s
,
a
ee
)
,而不是
Q
(
s
,
a
ee
,
a
gripper
)
。它专注于连续末端执行器动作的价值估计。

第三层:夹爪动作由 GraspCritic 单独学习

SACAgentHybridSingleArm 在网络集合里额外注册了 “grasp_critic”,并给它单独配置 optimizer:

networks = {
“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, encoder=encoders[“grasp_critic”], network=grasp_critic_backbone
)(name=“grasp_critic”)

6.3 动作执行流程(Rollout)
sample_actions 是 SAC actor 和 GraspCritic 配合最直观的接口。执行流程分为四步:

第一步:actor 采样连续末端执行器动作

dist = self.forward_policy(observations, rng=seed, train=False)
ee_actions = dist.sample(seed=seed)

这里得到前 6 维动作:[x, y, z, roll, pitch, yaw]。

第二步:GraspCritic 输出 3 个夹爪 Q 值

grasp_q_values = self.forward_grasp_critic(observations, rng=grasp_key, train=False)

输出类似:

Q(close) = 1.2
Q(keep) = 0.7
Q(open) = 0.1

第三步:argmax 选择夹爪动作

grasp_action = grasp_q_values.argmax(axis=-1)
grasp_action = grasp_action - 1 # {0,1,2} → {-1,0,1}

映射关系:argmax=0 → 环境动作 -1(关),argmax=1 → 环境动作 0(保持),argmax=2 → 环境动作 +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 返回四个独立的 loss:

return {
“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_critic:

train_critic_networks_to_update = frozenset({“critic”, “grasp_critic”})

完整训练阶段更新全部四个网络:

train_networks_to_update = frozenset({“critic”, “grasp_critic”, “actor”, “temperature”})

Critic Loss(连续动作 SAC)
目标 Q 采用 Clipped Double-Q:

y

r
+
γ

min
i

Q
¯
θ
i
(
s

,
a

)
代码实现:

target_next_qs = self.forward_target_critic(…)
target_next_min_q = target_next_qs.min(axis=0)
target_q = rewards + discount * masks * target_next_min_q

predicted_qs = self.forward_critic(batch[“observations”], actions, …)
critic_loss = jnp.mean((predicted_qs - target_qs) ** 2)

GraspCritic Loss(DQN 风格)
GraspCritic 采用 Double DQN 风格——用 online 网络选动作,用 target 网络评估:

next_grasp_qs = self.forward_grasp_critic(batch[“next_observations”], rng=rng)
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_q

predicted_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(标准 SAC)
Actor 采样连续动作,最大化:

objective

Q
(
s
,
a
)

α
log
π
(
a
|
s
)
actor_objective = predicted_q - temperature * log_probs
actor_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.penalty
else:
info[“grasp_penalty”] = 0.0

通俗理解:夹爪已经开得很大还继续开,或者已经关得很紧还继续关——这种动作没有实际意义,甚至可能损坏硬件或扰乱任务。

这类惩罚不适合影响机械臂连续运动的 SAC critic——因为 SAC critic 评估的是 6 维连续动作的价值,加上夹爪惩罚会混淆它对末端执行器动作质量的判断。但 GraspCritic 专门学习夹爪的离散决策,加上 grasp_penalty 可以更精准地训练"什么时候该夹、什么时候该放"的策略。

6.7 配合全景图