sac_actor.py — RobDataParallelSACActor 实现了完整的 SAC 训练逻辑¶
文件路径:
verl/experimental/vla/sac/sac_actor.py模块路径:verl.experimental.vla.sac.sac_actor
文件概述¶
RobDataParallelSACActor 实现了完整的 SAC 训练逻辑,包括 Critic 更新、Actor 更新、自动熵调节和回放缓冲区集成。
核心类:RobDataParallelSACActor¶
初始化¶
class RobDataParallelSACActor(BaseSACActor):
def __init__(self, model, config):
self.model = model # PI0ForActionPrediction(实现了 SupportSACTraining)
# 优化器
self.actor_optimizer = Adam(model.actor_params(), lr=config.actor_lr)
self.critic_optimizer = Adam(model.critic_params(), lr=config.critic_lr)
# 自动熵调节
self.log_alpha = torch.zeros(1, requires_grad=True)
self.alpha_optimizer = Adam([self.log_alpha], lr=config.alpha_lr)
self.target_entropy = -config.action_dim # 目标熵
# 回放缓冲区
self.replay_pool = SACReplayPool(config.replay_pool_size)
Critic 更新¶
def _forward_critic(self, batch):
"""更新 Critic 网络
\(L = \text{MSE}\left(Q(s,a),\; r + \gamma \left(\min Q_{\text{target}}(s',a') - \alpha \log \pi(a'|s')\right)\right)\)
"""
# 当前 Q 值
q_values = self.model.sac_forward_critic(state_features, actions)
# 计算目标值(不需要梯度)
with torch.no_grad():
next_actions = self.model.sac_forward_actor(next_state_features)
next_log_probs = self.model._get_logprobs(next_state_features, next_actions)
# Double Q: 取两个 Critic 的最小值
q_targets = [head(next_features) for head in self.model.target_network_heads]
min_q_target = torch.min(*q_targets)
target = reward + gamma * (min_q_target - alpha * next_log_probs)
# Critic 损失
critic_loss = sum(F.mse_loss(q, target) for q in q_values)
critic_loss.backward()
self.critic_optimizer.step()
Actor 更新¶
def _forward_actor(self, batch):
"""更新 Actor 网络
目标: 最大化 \(Q(s, a) - \alpha \log \pi(a|s)\)
"""
state_features = self.model.sac_state_features(batch)
actions = self.model.sac_forward_actor(state_features)
log_probs = self.model._get_logprobs(state_features, actions)
q_values = self.model.sac_forward_critic(state_features, actions)
min_q = torch.min(*q_values)
actor_loss = (self.alpha * log_probs - min_q).mean()
actor_loss.backward()
self.actor_optimizer.step()
# 自动调节熵系数
alpha_loss = -(self.log_alpha * (log_probs + self.target_entropy).detach()).mean()
alpha_loss.backward()
self.alpha_optimizer.step()
回放缓冲区集成¶
def update_with_replay(self, new_data):
"""使用回放缓冲区进行训练"""
# 将新数据加入缓冲区
self.replay_pool.add_batch(new_data)
# 从缓冲区采样
batch = self.replay_pool.sample_batch(self.batch_size)
# 更新
self._forward_critic(batch)
self._forward_actor(batch)
# 软更新目标网络
self.model.sac_update_target_network(tau=0.005)
核心类/函数列表¶
| 名称 | 类型 | 说明 |
|---|---|---|
RobDataParallelSACActor |
类 | SAC Actor 完整实现 |
_forward_critic |
方法 | Critic 网络更新 |
_forward_actor |
方法 | Actor 网络更新 + 熵调节 |
update_with_replay |
方法 | 使用回放缓冲区训练 |
与其他模块的关系¶
- 使用
SupportSACTraining接口(base.py) - 使用
SACReplayPool(replay_pool.py)做经验回放 - 被
RobRaySACTrainer(sac_ray_trainer.py)调用
小结¶
这个类将 SAC 算法的所有训练逻辑集中在一起。与 PPO 的关键区别是:SAC 使用回放缓冲区(离策略),有独立的 Critic 网络,并且通过自动熵调节平衡探索和利用。