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 核心组件:
- Actor \(\pi(a|s)\): 输出动作的概率分布
- Critic \(Q(s,a)\): 估计状态-动作值
- Target \(Q\): \(Q\) 的慢速跟踪副本(提高稳定性)
- 熵系数 \(\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实现了SupportSACTrainingRobDataParallelSACActor继承了BaseSACActor
小结¶
这个文件定义了 SAC 算法与 VLA 模型之间的"契约"。模型实现 SupportSACTraining 接口,训练器使用 BaseSACActor 接口,两者通过这些抽象层解耦。