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