跳转至

base.py — 定义了 SAC(Soft Actor-Critic)算法的接口

文件路径: verl/experimental/vla/sac/base.py 模块路径: verl.experimental.vla.sac.base

文件概述

定义了 SAC(Soft Actor-Critic)算法的接口。包括 SupportSACTraining 接口(模型需要实现的 SAC 方法)和 BaseSACActor 抽象基类(SAC Actor 的基类)。

核心接口

SupportSACTraining

任何想要支持 SAC 训练的模型都需要实现这个接口:

class SupportSACTraining:
    """SAC 训练支持接口

    模型(如 PI0ForActionPrediction)需要实现这些方法
    才能被 SAC 训练器使用。
    """

    def sac_init(self):
        """初始化 SAC 相关组件(Critic、Target Network 等)"""
        ...

    def sac_forward_critic(self, state_features, actions):
        """Critic 前向:估计 \(Q(s, a)\)"""
        ...

    def sac_forward_actor(self, state_features):
        """Actor 前向:采样动作 \(a \sim \pi(\cdot|s)\)"""
        ...

    def sac_state_features(self, data):
        """提取状态特征(如视觉语言编码)"""
        ...

    def sac_update_target_network(self, tau=0.005):
        """目标网络软更新

        \(\theta_{\text{target}} = \tau \cdot \theta_{\text{current}} + (1-\tau) \cdot \theta_{\text{target}}\)
        """
        ...

BaseSACActor

class BaseSACActor(ABC):
    """SAC Actor 的抽象基类"""

    @abstractmethod
    def update_actor(self, data):
        """更新 Actor 网络"""
        ...

    @abstractmethod
    def update_critic(self, data):
        """更新 Critic 网络"""
        ...

SAC 算法简介

SAC 核心组件:

  1. Actor \(\pi(a|s)\): 输出动作的概率分布
  2. Critic \(Q(s,a)\): 估计状态-动作值
  3. Target \(Q\): \(Q\) 的慢速跟踪副本(提高稳定性)
  4. 熵系数 \(\alpha\): 自动调整探索程度

更新规则:

  • Critic: $$ \min L_Q = \left(Q(s,a) - \left(r + \gamma \left(Q_{\text{target}}(s', a') - \alpha \log \pi(a'|s')\right)\right)\right)^2 $$

  • Actor: $$ \max J = Q(s, a) - \alpha \log \pi(a|s) $$

  • \(\alpha\): 自动调整使熵 \(H[\pi]\) 接近目标值

核心类列表

名称 类型 说明
SupportSACTraining 接口 模型需要实现的 SAC 方法
BaseSACActor 抽象基类 SAC Actor 基类

与其他模块的关系

  • PI0ForActionPrediction 实现了 SupportSACTraining
  • RobDataParallelSACActor 继承了 BaseSACActor

小结

这个文件定义了 SAC 算法与 VLA 模型之间的"契约"。模型实现 SupportSACTraining 接口,训练器使用 BaseSACActor 接口,两者通过这些抽象层解耦。