相比普通 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 配合全景图