Files
RL_TRPO/agent/ppo.py
T

192 lines
7.4 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
import torch.nn as nn
import torch.nn.functional as F
from networks import PolicyNetwork, ValueNetwork
class PPOAgent:
"""
PPO-Clip (Proximal Policy Optimization, Clipped version)
Reference: Schulman et al., "Proximal Policy Optimization Algorithms", 2017.
与 TRPO 的核心区别:
- 不再使用共轭梯度 + 线搜索求解约束优化
- 用 clip(ratio, 1-ε, 1+ε) 限制策略更新幅度
- 支持多 epoch 的小批量更新(每次从经验池中采样)
"""
def __init__(
self,
state_dim,
action_dim,
action_bound,
hidden_dim=128,
gamma=0.99,
tau=0.95,
actor_lr=3e-4,
critic_lr=1e-3,
clip_eps=0.2,
k_epochs=10,
minibatch_size=64,
max_grad_norm=0.5,
device="cpu",
):
self.gamma = gamma
self.tau = tau
self.clip_eps = clip_eps # PPO-Clip 的裁剪范围 ε
self.k_epochs = k_epochs # 每次更新对同一批数据的训练轮数
self.minibatch_size = minibatch_size
self.max_grad_norm = max_grad_norm # 梯度裁剪阈值
self.device = torch.device(device)
self.actor = PolicyNetwork(state_dim, action_dim, action_bound, hidden_dim).to(self.device)
self.critic = ValueNetwork(state_dim, hidden_dim).to(self.device)
self.actor_optimizer = torch.optim.Adam(self.actor.parameters(), lr=actor_lr)
self.critic_optimizer = torch.optim.Adam(self.critic.parameters(), lr=critic_lr)
def _to_tensor(self, array_like):
return torch.as_tensor(array_like, dtype=torch.float32, device=self.device)
def get_action(self, state):
state_tensor = self._to_tensor(state)
single_state = state_tensor.ndim == 1
if single_state:
state_tensor = state_tensor.unsqueeze(0)
with torch.no_grad():
dist = self.actor.evaluate(state_tensor)
action = dist.sample()
action_np = action.detach().cpu().numpy()
return action_np[0] if single_state else action_np
def get_value(self, state):
state_tensor = self._to_tensor(state)
single_state = state_tensor.ndim == 1
if single_state:
state_tensor = state_tensor.unsqueeze(0)
with torch.no_grad():
values = self.critic(state_tensor).squeeze(-1)
values_np = values.detach().cpu().numpy()
return float(values_np[0]) if single_state else values_np
def _compute_advantages(self, rewards, values, next_values, masks):
"""
GAE (Generalized Advantage Estimation)
Eq.(11): Â_t = δ_t + (γλ)δ_{t+1} + ... + (γλ)^{T-t+1} δ_{T-1}
Eq.(12): δ_t = r_t + γV(s_{t+1}) - V(s_t)
"""
advantages = torch.zeros_like(rewards, device=self.device)
gae = torch.zeros(rewards.size(1), dtype=torch.float32, device=self.device)
for i in reversed(range(rewards.size(0))):
delta = rewards[i] + self.gamma * next_values[i] * masks[i] - values[i]
gae = delta + self.gamma * self.tau * masks[i] * gae
advantages[i] = gae
returns = advantages + values
return returns, advantages
def update(self, memory):
"""
PPO-Clip 更新 — 对齐论文 Algorithm 1 和 Eq.(9)
关键改动(相比旧版):
1. Actor 和 Critic 在同一个 minibatch 循环内联合训练
L = L^CLIP - c1 * L^VF (Eq.9c2=0 不加熵)
2. 梯度裁剪 (max_grad_norm)
"""
if not memory:
return {}
# ---------- 1. 准备数据 ----------
states = self._to_tensor(np.asarray(memory["states"]))
actions = self._to_tensor(np.asarray(memory["actions"]))
rewards = self._to_tensor(np.asarray(memory["rewards"]))
next_states = self._to_tensor(np.asarray(memory["next_states"]))
masks = self._to_tensor(np.asarray(memory["masks"]))
# ---------- 2. GAE 优势估计 ----------
with torch.no_grad():
rollout_steps, num_envs = rewards.shape
flat_states = states.reshape(rollout_steps * num_envs, -1)
flat_next_states = next_states.reshape(rollout_steps * num_envs, -1)
values = self.critic(flat_states).squeeze(-1).reshape(rollout_steps, num_envs)
next_values = self.critic(flat_next_states).squeeze(-1).reshape(rollout_steps, num_envs)
returns, advantages = self._compute_advantages(rewards, values, next_values, masks)
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
states = flat_states
actions = actions.reshape(rollout_steps * num_envs, -1)
returns = returns.reshape(rollout_steps * num_envs)
advantages = advantages.reshape(rollout_steps * num_envs)
# ---------- 3. 记录旧策略 π_θold ----------
with torch.no_grad():
old_dist = self.actor.evaluate(states)
old_log_probs = old_dist.log_prob(actions).sum(dim=1)
# ---------- 4. K epochs × minibatch 联合训练 ----------
dataset_size = states.size(0)
indices = np.arange(dataset_size)
policy_loss_value = 0.0
critic_loss_value = 0.0
entropy_value = 0.0
minibatch_updates = 0
for _ in range(self.k_epochs):
np.random.shuffle(indices)
for start in range(0, dataset_size, self.minibatch_size):
end = start + self.minibatch_size
mb_idx = torch.as_tensor(indices[start:end], device=self.device, dtype=torch.long)
mb_states = states[mb_idx]
mb_actions = actions[mb_idx]
mb_advantages = advantages[mb_idx]
mb_returns = returns[mb_idx]
mb_old_log_probs = old_log_probs[mb_idx]
# --- Actor: PPO-Clip 损失 ---
dist = self.actor.evaluate(mb_states)
log_probs = dist.log_prob(mb_actions).sum(dim=1)
ratio = torch.exp(log_probs - mb_old_log_probs)
surr1 = ratio * mb_advantages
surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * mb_advantages
policy_loss = -torch.min(surr1, surr2).mean()
# --- Critic: 价值函数 MSE 损失 ---
value_pred = self.critic(mb_states).squeeze()
critic_loss = F.mse_loss(value_pred, mb_returns)
# Actor 更新(带梯度裁剪,防止策略大幅跳变)
self.actor_optimizer.zero_grad()
policy_loss.backward()
nn.utils.clip_grad_norm_(self.actor.parameters(), self.max_grad_norm)
self.actor_optimizer.step()
# Critic 更新(不裁剪梯度,Adam 自适应处理大梯度即可)
self.critic_optimizer.zero_grad()
critic_loss.backward()
self.critic_optimizer.step()
policy_loss_value += float(policy_loss.detach().item())
critic_loss_value += float(critic_loss.detach().item())
entropy_value += float(dist.entropy().sum(dim=1).mean().detach().item())
minibatch_updates += 1
divisor = max(1, minibatch_updates)
return {
"policy_loss": policy_loss_value / divisor,
"critic_loss": critic_loss_value / divisor,
"entropy": entropy_value / divisor,
}