sac_ray_trainer.py — RobRaySACTrainer 是机器人 VLA 模型的 SAC 训练器¶
文件路径:
verl/experimental/vla/sac/sac_ray_trainer.py模块路径:verl.experimental.vla.sac.sac_ray_trainer
文件概述¶
RobRaySACTrainer 是机器人 VLA 模型的 SAC 训练器,结构类似于 RobRayPPOTrainer,但使用 SAC 算法。主要区别是使用回放缓冲区和独立的 Actor/Critic 更新。
关键代码¶
训练循环¶
class RobRaySACTrainer:
def fit(self):
for step in range(self.total_steps):
# 1. 环境重置
self._reset_envs(step)
# 2. 收集轨迹
trajectories = self.env_loop.rollout()
# 3. 添加到回放缓冲区
transitions = self.add_transition_prefixes(trajectories)
self.replay_pool.add_batch(transitions)
# 4. 从缓冲区采样并更新
for _ in range(self.gradient_steps_per_env_step):
batch = self.replay_pool.sample_batch(self.batch_size)
self.actor_wg.update_critic(batch)
self.actor_wg.update_actor(batch)
self.actor_wg.update_target_network()
转移数据预处理¶
def add_transition_prefixes(self, trajectories):
"""为 SAC 准备转移数据
将轨迹拆分为 (s, a, r, s', done) 元组,
并计算 response_mask。
"""
transitions = {
"state": states,
"action": actions,
"reward": rewards,
"next_state": next_states,
"done": dones,
"response_mask": self.compute_response_mask(data),
}
return transitions
SAC vs PPO 训练对比¶
PPO 训练循环: SAC 训练循环:
1. 收集轨迹 1. 收集轨迹
2. 计算优势函数 2. 添加到回放缓冲区
3. 多轮 PPO 更新 3. 多次采样+更新:
- Actor (策略梯度) - Critic (TD 损失)
- Critic (价值函数) - Actor (策略损失)
4. 丢弃旧数据 - Target (软更新)
4. 保留旧数据
核心类/函数列表¶
| 名称 | 类型 | 说明 |
|---|---|---|
RobRaySACTrainer |
类 | SAC 训练器 |
fit |
方法 | 主训练循环 |
add_transition_prefixes |
方法 | 转移数据预处理 |
compute_response_mask |
方法 | 计算动作 token mask |
与其他模块的关系¶
- 使用
RobDataParallelSACActor(sac_actor.py)做 Actor/Critic 更新 - 使用
SACReplayPool(replay_pool.py)做经验回放 - 使用
EnvLoop(env_loop.py)做环境交互 - 被
main_sac.py创建和调用
小结¶
RobRaySACTrainer 将 SAC 算法适配到了 VLA 机器人训练场景。与 PPO 训练器相比,它的核心区别是使用回放缓冲区实现数据复用,以及 Actor 和 Critic 的独立更新步骤。