跳转至

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 网络,并且通过自动熵调节平衡探索和利用。