Files
RL_SAC/utils/rollout_buffer.py
T

105 lines
4.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import numpy as np
import torch
class RolloutBuffer:
"""
On-policy 滚动缓冲区,用于 PPO 算法。
每次 collect() 收集 T 步数据后,调用 compute_returns_and_advantages()
计算 GAE 优势估计,然后通过 get_batches() 将数据切分为 mini-batch
供 K 轮 epoch 使用,最后 clear() 清空等待下一轮收集。
"""
def __init__(self, state_dim, action_dim, steps_per_update, gamma, gae_lambda, device):
self.steps = steps_per_update
self.gamma = gamma
self.gae_lambda = gae_lambda
self.device = device
# 预分配存储空间
self.states = np.zeros((steps_per_update, state_dim), dtype=np.float32)
self.actions_unbounded = np.zeros((steps_per_update, action_dim), dtype=np.float32) # clamp 之前的高斯采样值
self.log_probs = np.zeros((steps_per_update, 1), dtype=np.float32)
self.rewards = np.zeros((steps_per_update, 1), dtype=np.float32)
self.dones = np.zeros((steps_per_update, 1), dtype=np.float32)
self.values = np.zeros((steps_per_update, 1), dtype=np.float32)
# 计算后填充
self.returns = np.zeros((steps_per_update, 1), dtype=np.float32)
self.advantages = np.zeros((steps_per_update, 1), dtype=np.float32)
self.ptr = 0
self.full = False
def add(self, state, action_unbounded, log_prob, reward, done, value):
"""
向缓冲区写入一步数据。
action_unbounded: 未裁剪的高斯采样值(形状 [action_dim]
log_prob: 该步的 log π(a|s)(标量)
value: V(s) 的估计值(标量)
"""
idx = self.ptr
self.states[idx] = state
self.actions_unbounded[idx] = action_unbounded
self.log_probs[idx] = log_prob
self.rewards[idx] = reward
self.dones[idx] = done
self.values[idx] = value
self.ptr += 1
if self.ptr >= self.steps:
self.full = True
def compute_returns_and_advantages(self, last_value):
"""
反向遍历轨迹,计算 GAE 优势估计(论文公式11/12)。
只对实际填充的 self.ptr 步数据计算,避免无效数据参与。
Args:
last_value: V(s_{T+1}),下一个状态的价值估计(若 episode 结束则为 0)
"""
n = self.ptr # 实际有效数据条数
last_gae = 0.0
for t in reversed(range(n)):
if t == n - 1:
next_non_terminal = 1.0 - self.dones[t]
next_value = last_value
else:
next_non_terminal = 1.0 - self.dones[t]
next_value = self.values[t + 1]
delta = self.rewards[t] + self.gamma * next_value * next_non_terminal - self.values[t]
last_gae = delta + self.gamma * self.gae_lambda * next_non_terminal * last_gae
self.advantages[t] = last_gae
# 回报 G_t = Â_t + V(s_t)
self.returns[:n] = self.advantages[:n] + self.values[:n]
# 优势归一化
adv = self.advantages[:n]
self.advantages[:n] = (adv - adv.mean()) / (adv.std() + 1e-8)
def get_batches(self, batch_size):
"""
将实际填充的数据随机打乱后按 batch_size 切片,生成 mini-batch。
Yields:
(states, actions_unbounded, log_probs, returns, advantages) — 均为 Tensor
"""
n = self.ptr
indices = np.random.permutation(n)
for start in range(0, n, batch_size):
idx = indices[start: start + batch_size]
yield (
torch.FloatTensor(self.states[idx]).to(self.device),
torch.FloatTensor(self.actions_unbounded[idx]).to(self.device),
torch.FloatTensor(self.log_probs[idx]).to(self.device),
torch.FloatTensor(self.returns[idx]).to(self.device),
torch.FloatTensor(self.advantages[idx]).to(self.device),
)
def clear(self):
"""清空缓冲区,为下一轮收集做准备。"""
self.ptr = 0
self.full = False