100 lines
3.7 KiB
Python
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
|