Files

100 lines
3.7 KiB
Python

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