2026-04-02 02:24:35 +00:00
|
|
|
|
import numpy as np
|
|
|
|
|
|
import torch
|
2026-04-02 16:36:55 +08:00
|
|
|
|
import torch.nn as nn
|
2026-04-02 02:24:35 +00:00
|
|
|
|
import torch.nn.functional as F
|
2026-04-02 09:48:59 +00:00
|
|
|
|
|
2026-04-02 02:24:35 +00:00
|
|
|
|
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,
|
2026-04-02 16:36:55 +08:00
|
|
|
|
tau=0.95,
|
|
|
|
|
|
actor_lr=3e-4,
|
|
|
|
|
|
critic_lr=1e-3,
|
2026-04-02 02:24:35 +00:00
|
|
|
|
clip_eps=0.2,
|
|
|
|
|
|
k_epochs=10,
|
|
|
|
|
|
minibatch_size=64,
|
2026-04-02 16:36:55 +08:00
|
|
|
|
max_grad_norm=0.5,
|
2026-04-02 09:48:59 +00:00
|
|
|
|
device="cpu",
|
2026-04-02 02:24:35 +00:00
|
|
|
|
):
|
|
|
|
|
|
self.gamma = gamma
|
|
|
|
|
|
self.tau = tau
|
|
|
|
|
|
self.clip_eps = clip_eps # PPO-Clip 的裁剪范围 ε
|
|
|
|
|
|
self.k_epochs = k_epochs # 每次更新对同一批数据的训练轮数
|
|
|
|
|
|
self.minibatch_size = minibatch_size
|
2026-04-02 16:36:55 +08:00
|
|
|
|
self.max_grad_norm = max_grad_norm # 梯度裁剪阈值
|
2026-04-02 09:48:59 +00:00
|
|
|
|
self.device = torch.device(device)
|
2026-04-02 02:24:35 +00:00
|
|
|
|
|
2026-04-02 09:48:59 +00:00
|
|
|
|
self.actor = PolicyNetwork(state_dim, action_dim, action_bound, hidden_dim).to(self.device)
|
|
|
|
|
|
self.critic = ValueNetwork(state_dim, hidden_dim).to(self.device)
|
2026-04-02 02:24:35 +00:00
|
|
|
|
|
2026-04-02 16:36:55 +08:00
|
|
|
|
self.actor_optimizer = torch.optim.Adam(self.actor.parameters(), lr=actor_lr)
|
|
|
|
|
|
self.critic_optimizer = torch.optim.Adam(self.critic.parameters(), lr=critic_lr)
|
2026-04-02 02:24:35 +00:00
|
|
|
|
|
2026-04-02 09:48:59 +00:00
|
|
|
|
def _to_tensor(self, array_like):
|
|
|
|
|
|
return torch.as_tensor(array_like, dtype=torch.float32, device=self.device)
|
|
|
|
|
|
|
2026-04-02 02:24:35 +00:00
|
|
|
|
def get_action(self, state):
|
2026-04-02 09:48:59 +00:00
|
|
|
|
state_tensor = self._to_tensor(state)
|
|
|
|
|
|
single_state = state_tensor.ndim == 1
|
|
|
|
|
|
if single_state:
|
|
|
|
|
|
state_tensor = state_tensor.unsqueeze(0)
|
|
|
|
|
|
|
2026-04-02 02:24:35 +00:00
|
|
|
|
with torch.no_grad():
|
|
|
|
|
|
dist = self.actor.evaluate(state_tensor)
|
|
|
|
|
|
action = dist.sample()
|
2026-04-02 09:48:59 +00:00
|
|
|
|
|
|
|
|
|
|
action_np = action.detach().cpu().numpy()
|
|
|
|
|
|
return action_np[0] if single_state else action_np
|
2026-04-02 02:24:35 +00:00
|
|
|
|
|
2026-04-02 16:36:55 +08:00
|
|
|
|
def get_value(self, state):
|
2026-04-02 09:48:59 +00:00
|
|
|
|
state_tensor = self._to_tensor(state)
|
|
|
|
|
|
single_state = state_tensor.ndim == 1
|
|
|
|
|
|
if single_state:
|
|
|
|
|
|
state_tensor = state_tensor.unsqueeze(0)
|
2026-04-02 16:36:55 +08:00
|
|
|
|
|
2026-04-02 09:48:59 +00:00
|
|
|
|
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):
|
2026-04-02 02:24:35 +00:00
|
|
|
|
"""
|
|
|
|
|
|
GAE (Generalized Advantage Estimation)
|
2026-04-02 16:36:55 +08:00
|
|
|
|
Eq.(11): Â_t = δ_t + (γλ)δ_{t+1} + ... + (γλ)^{T-t+1} δ_{T-1}
|
|
|
|
|
|
Eq.(12): δ_t = r_t + γV(s_{t+1}) - V(s_t)
|
2026-04-02 02:24:35 +00:00
|
|
|
|
"""
|
2026-04-02 09:48:59 +00:00
|
|
|
|
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]
|
2026-04-02 02:24:35 +00:00
|
|
|
|
gae = delta + self.gamma * self.tau * masks[i] * gae
|
2026-04-02 09:48:59 +00:00
|
|
|
|
advantages[i] = gae
|
|
|
|
|
|
|
|
|
|
|
|
returns = advantages + values
|
|
|
|
|
|
return returns, advantages
|
2026-04-02 02:24:35 +00:00
|
|
|
|
|
|
|
|
|
|
def update(self, memory):
|
|
|
|
|
|
"""
|
2026-04-02 16:36:55 +08:00
|
|
|
|
PPO-Clip 更新 — 对齐论文 Algorithm 1 和 Eq.(9)
|
2026-04-02 02:24:35 +00:00
|
|
|
|
|
2026-04-02 16:36:55 +08:00
|
|
|
|
关键改动(相比旧版):
|
|
|
|
|
|
1. Actor 和 Critic 在同一个 minibatch 循环内联合训练
|
|
|
|
|
|
L = L^CLIP - c1 * L^VF (Eq.9,c2=0 不加熵)
|
|
|
|
|
|
2. 梯度裁剪 (max_grad_norm)
|
2026-04-02 02:24:35 +00:00
|
|
|
|
"""
|
2026-04-02 09:48:59 +00:00
|
|
|
|
if not memory:
|
|
|
|
|
|
return {}
|
|
|
|
|
|
|
2026-04-02 02:24:35 +00:00
|
|
|
|
# ---------- 1. 准备数据 ----------
|
2026-04-02 09:48:59 +00:00
|
|
|
|
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"]))
|
2026-04-02 02:24:35 +00:00
|
|
|
|
|
|
|
|
|
|
# ---------- 2. GAE 优势估计 ----------
|
|
|
|
|
|
with torch.no_grad():
|
2026-04-02 09:48:59 +00:00
|
|
|
|
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)
|
2026-04-02 02:24:35 +00:00
|
|
|
|
|
2026-04-02 09:48:59 +00:00
|
|
|
|
returns, advantages = self._compute_advantages(rewards, values, next_values, masks)
|
2026-04-02 02:24:35 +00:00
|
|
|
|
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
|
|
|
|
|
|
|
2026-04-02 09:48:59 +00:00
|
|
|
|
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)
|
|
|
|
|
|
|
2026-04-02 16:36:55 +08:00
|
|
|
|
# ---------- 3. 记录旧策略 π_θold ----------
|
2026-04-02 02:24:35 +00:00
|
|
|
|
with torch.no_grad():
|
|
|
|
|
|
old_dist = self.actor.evaluate(states)
|
|
|
|
|
|
old_log_probs = old_dist.log_prob(actions).sum(dim=1)
|
|
|
|
|
|
|
2026-04-02 16:36:55 +08:00
|
|
|
|
# ---------- 4. K epochs × minibatch 联合训练 ----------
|
2026-04-02 02:24:35 +00:00
|
|
|
|
dataset_size = states.size(0)
|
|
|
|
|
|
indices = np.arange(dataset_size)
|
2026-04-02 09:48:59 +00:00
|
|
|
|
policy_loss_value = 0.0
|
|
|
|
|
|
critic_loss_value = 0.0
|
|
|
|
|
|
entropy_value = 0.0
|
|
|
|
|
|
minibatch_updates = 0
|
2026-04-02 02:24:35 +00:00
|
|
|
|
|
|
|
|
|
|
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
|
2026-04-02 09:48:59 +00:00
|
|
|
|
mb_idx = torch.as_tensor(indices[start:end], device=self.device, dtype=torch.long)
|
2026-04-02 02:24:35 +00:00
|
|
|
|
|
|
|
|
|
|
mb_states = states[mb_idx]
|
|
|
|
|
|
mb_actions = actions[mb_idx]
|
|
|
|
|
|
mb_advantages = advantages[mb_idx]
|
2026-04-02 16:36:55 +08:00
|
|
|
|
mb_returns = returns[mb_idx]
|
2026-04-02 02:24:35 +00:00
|
|
|
|
mb_old_log_probs = old_log_probs[mb_idx]
|
|
|
|
|
|
|
2026-04-02 16:36:55 +08:00
|
|
|
|
# --- Actor: PPO-Clip 损失 ---
|
2026-04-02 02:24:35 +00:00
|
|
|
|
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()
|
|
|
|
|
|
|
2026-04-02 16:36:55 +08:00
|
|
|
|
# --- Critic: 价值函数 MSE 损失 ---
|
|
|
|
|
|
value_pred = self.critic(mb_states).squeeze()
|
|
|
|
|
|
critic_loss = F.mse_loss(value_pred, mb_returns)
|
2026-04-02 02:24:35 +00:00
|
|
|
|
|
2026-04-02 16:36:55 +08:00
|
|
|
|
# Actor 更新(带梯度裁剪,防止策略大幅跳变)
|
2026-04-02 02:24:35 +00:00
|
|
|
|
self.actor_optimizer.zero_grad()
|
2026-04-02 16:36:55 +08:00
|
|
|
|
policy_loss.backward()
|
|
|
|
|
|
nn.utils.clip_grad_norm_(self.actor.parameters(), self.max_grad_norm)
|
2026-04-02 02:24:35 +00:00
|
|
|
|
self.actor_optimizer.step()
|
2026-04-02 16:36:55 +08:00
|
|
|
|
|
|
|
|
|
|
# Critic 更新(不裁剪梯度,Adam 自适应处理大梯度即可)
|
|
|
|
|
|
self.critic_optimizer.zero_grad()
|
|
|
|
|
|
critic_loss.backward()
|
|
|
|
|
|
self.critic_optimizer.step()
|
2026-04-02 09:48:59 +00:00
|
|
|
|
|
|
|
|
|
|
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,
|
|
|
|
|
|
}
|