import torch import numpy as np class RolloutBuffer: """ 经验回放池:用于收集智能体与环境交互的轨迹数据, 并在一个回合(或一个 Batch)结束后计算优势函数 GAE 和目标价值。 """ def __init__(self): # 初始化存储列表 self.states = [] self.actions = [] self.rewards = [] self.next_states = [] self.dones = [] self.log_probs = [] self.values = [] def add(self, state, action, reward, next_state, done, log_prob, value): """ 向池子中添加一步交互的数据 """ self.states.append(state) self.actions.append(action) self.rewards.append(reward) self.next_states.append(next_state) self.dones.append(done) self.log_probs.append(log_prob) self.values.append(value) def clear(self): """ 清空池子,准备收集下一批数据 """ self.states.clear() self.actions.clear() self.rewards.clear() self.next_states.clear() self.dones.clear() self.log_probs.clear() self.values.clear() def compute_returns_and_advantages(self, last_value, gamma=0.99, lam=0.95): """ 计算广义优势估计 (GAE) 和 目标价值 (Returns)。 这是 TRPO/PPO 最核心的数据处理步骤! 参数: last_value: 截断处(或回合结束时)的最后一个状态的 V 值。 gamma: 折扣因子 (Discount factor)。 lam: GAE 的平滑参数 (Lambda),用于权衡偏差和方差。 """ # 将列表转换为 NumPy 数组,方便进行向量化运算 rewards = np.array(self.rewards, dtype=np.float32) values = np.array(self.values, dtype=np.float32) dones = np.array(self.dones, dtype=np.float32) # 预分配数组空间 advantages = np.zeros_like(rewards, dtype=np.float32) last_gae_lam = 0 # 逆序遍历轨迹:从最后一步往前推算 for t in reversed(range(len(rewards))): if t == len(rewards) - 1: # 如果是最后一步,next_value 就是传入的 last_value next_non_terminal = 1.0 - dones[t] next_value = last_value else: # 否则,next_value 就是下一步的 value next_non_terminal = 1.0 - dones[t] next_value = values[t + 1] # 计算 TD 误差 (Temporal Difference Error) # delta = r_t + gamma * V(s_{t+1}) - V(s_t) delta = rewards[t] + gamma * next_value * next_non_terminal - values[t] # 递推计算 GAE # A_t = delta_t + gamma * lambda * A_{t+1} advantages[t] = last_gae_lam = delta + gamma * lam * next_non_terminal * last_gae_lam # 目标价值 = 优势函数 + 状态价值 returns = advantages + values # 优势函数标准化 (Advantage Normalization) # 这是一个极度重要的工程 Trick,能大幅提升训练稳定性 advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8) return returns, advantages def get_data(self): """ 将收集到的所有数据转换为 PyTorch Tensor,供后续网络训练使用 """ # 将 NumPy 数组转为 Tensor state_tensor = torch.tensor(np.array(self.states), dtype=torch.float32) action_tensor = torch.tensor(np.array(self.actions), dtype=torch.float32) old_log_probs_tensor = torch.tensor(np.array(self.log_probs), dtype=torch.float32) return state_tensor, action_tensor, old_log_probs_tensor