跳转至

replay_pool.py — SACReplayPool 实现了 SAC 算法所需的经验回放缓冲区(Replay Bu...

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

文件概述

SACReplayPool 实现了 SAC 算法所需的经验回放缓冲区(Replay Buffer)。它使用循环缓冲区存储转移元组 (s, a, r, s', done),支持批量添加、均匀采样和持久化保存/加载。

关键代码

循环缓冲区

class SACReplayPool:
    """SAC 回放缓冲区"""

    def __init__(self, capacity: int):
        self.capacity = capacity
        self.buffer = []
        self.position = 0  # 循环写入位置

    def add_batch(self, transitions: dict):
        """批量添加转移数据

        当缓冲区满时,新数据覆盖最旧的数据(循环写入)。
        """
        batch_size = len(transitions["reward"])
        for i in range(batch_size):
            transition = {k: v[i] for k, v in transitions.items()}
            if len(self.buffer) < self.capacity:
                self.buffer.append(transition)
            else:
                self.buffer[self.position] = transition
            self.position = (self.position + 1) % self.capacity

均匀采样

    def sample_batch(self, batch_size: int) -> dict:
        """从缓冲区中均匀随机采样一个 batch"""
        indices = np.random.choice(len(self.buffer), batch_size, replace=False)
        batch = {
            key: torch.stack([self.buffer[i][key] for i in indices])
            for key in self.buffer[0].keys()
        }
        return batch

插入并重采样

    def insert_and_resample(self, new_data, batch_size):
        """原子操作:插入新数据并立即采样

        在分布式训练中,这确保了新数据有机会被采样到。
        """
        self.add_batch(new_data)
        return self.sample_batch(batch_size)

持久化

    def save(self, path: str):
        """保存缓冲区到文件"""
        torch.save({
            "buffer": self.buffer,
            "position": self.position,
            "capacity": self.capacity,
        }, path)

    def load(self, path: str):
        """从文件加载缓冲区"""
        data = torch.load(path)
        self.buffer = data["buffer"]
        self.position = data["position"]

核心类/函数列表

名称 类型 说明
SACReplayPool 类 循环回放缓冲区
add_batch 方法 批量添加转移
sample_batch 方法 均匀随机采样
insert_and_resample 方法 插入+采样原子操作
save / load 方法 持久化保存/加载

与其他模块的关系

  • 被 RobDataParallelSACActor(sac_actor.py)使用
  • 被 RobRaySACTrainer(sac_ray_trainer.py)管理

小结

经验回放是 SAC 等离策略算法的核心组件。它打破了数据的时间相关性(通过随机采样),提高了样本效率(数据可以被多次使用)。循环缓冲区设计确保了内存使用是有界的。