140 lines
5.5 KiB
Python
140 lines
5.5 KiB
Python
import torch
|
||
import torch.nn.functional as F
|
||
import torch.optim as optim
|
||
import numpy as np
|
||
from models.networks import PPOActor, ValueNet
|
||
from utils.rollout_buffer import RolloutBuffer
|
||
|
||
|
||
class PPO:
|
||
"""
|
||
Proximal Policy Optimization (PPO-Clip)
|
||
参考论文:Schulman et al., 2017 (arXiv:1707.06347)
|
||
|
||
核心目标函数(论文公式9):
|
||
L^{CLIP+VF+S} = E[ L^CLIP - c1 * L^VF + c2 * S[π](s) ]
|
||
|
||
使用独立的 Actor 和 Value 网络,GAE 优势估计(公式11/12)。
|
||
"""
|
||
|
||
def __init__(self, state_dim, action_dim, max_action, config):
|
||
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||
|
||
# 超参数
|
||
self.gamma = config.get('gamma', 0.99)
|
||
self.gae_lambda = config.get('gae_lambda', 0.95)
|
||
self.clip_epsilon = config.get('clip_epsilon', 0.2)
|
||
self.n_epochs = config.get('n_epochs', 10)
|
||
self.batch_size = config.get('batch_size', 64)
|
||
self.vf_coef = config.get('vf_coef', 0.5)
|
||
self.entropy_coef = config.get('entropy_coef', 0.01)
|
||
self.max_grad_norm = config.get('max_grad_norm', 0.5)
|
||
steps_per_update = config.get('steps_per_update', 2048)
|
||
|
||
# 策略网络 (Actor)
|
||
self.actor = PPOActor(state_dim, action_dim, max_action).to(self.device)
|
||
self.actor_optimizer = optim.Adam(self.actor.parameters(), lr=config.get('lr', 3e-4))
|
||
|
||
# 价值网络 (V-net)
|
||
self.value_net = ValueNet(state_dim).to(self.device)
|
||
self.value_optimizer = optim.Adam(self.value_net.parameters(), lr=config.get('lr', 3e-4))
|
||
|
||
# On-policy 滚动缓冲区
|
||
self.rollout = RolloutBuffer(
|
||
state_dim, action_dim, steps_per_update,
|
||
self.gamma, self.gae_lambda, self.device
|
||
)
|
||
|
||
# ------------------------------------------------------------------
|
||
# 与环境交互
|
||
# ------------------------------------------------------------------
|
||
|
||
@torch.no_grad()
|
||
def select_action(self, state):
|
||
"""
|
||
采样动作并返回:
|
||
action (np.ndarray): clamp 后的实际动作
|
||
action_unbounded (np.ndarray): 未裁剪的高斯采样值(存入 RolloutBuffer 用于 evaluate)
|
||
log_prob (float)
|
||
value (float): V(s)
|
||
"""
|
||
state_t = torch.FloatTensor(state).unsqueeze(0).to(self.device)
|
||
|
||
# Actor:采样 → clamp
|
||
action, log_prob, action_unbounded = self.actor.sample(state_t)
|
||
|
||
# Critic:价值估计
|
||
value = self.value_net(state_t)
|
||
|
||
return (
|
||
action.cpu().numpy().flatten(),
|
||
action_unbounded.cpu().numpy().flatten(),
|
||
log_prob.cpu().item(),
|
||
value.cpu().item(),
|
||
)
|
||
|
||
@torch.no_grad()
|
||
def get_value(self, state):
|
||
"""获取当前状态的 V(s),用于 GAE 计算的 last_value。"""
|
||
state_t = torch.FloatTensor(state).unsqueeze(0).to(self.device)
|
||
return self.value_net(state_t).cpu().item()
|
||
|
||
# ------------------------------------------------------------------
|
||
# 核心更新
|
||
# ------------------------------------------------------------------
|
||
|
||
def update(self):
|
||
"""
|
||
使用 RolloutBuffer 中收集的数据,进行 K 轮 mini-batch 更新。
|
||
对应论文 Algorithm 1。
|
||
注意:调用前需已执行 rollout.compute_returns_and_advantages()
|
||
"""
|
||
actor_losses, value_losses, entropy_bonuses = [], [], []
|
||
|
||
for _ in range(self.n_epochs):
|
||
for states, actions_unbounded, old_log_probs, returns, advantages in \
|
||
self.rollout.get_batches(self.batch_size):
|
||
|
||
# ---- 重新计算当前策略的 log_prob 和熵 ----
|
||
new_log_probs, entropy = self.actor.evaluate(states, actions_unbounded)
|
||
|
||
# 概率比率 r_t(θ) = π_θ(a|s) / π_θ_old(a|s)
|
||
ratio = torch.exp(new_log_probs - old_log_probs)
|
||
|
||
# ---- L^CLIP 目标(论文公式7)----
|
||
surrogate1 = ratio * advantages
|
||
surrogate2 = torch.clamp(ratio, 1 - self.clip_epsilon, 1 + self.clip_epsilon) * advantages
|
||
actor_loss = -torch.min(surrogate1, surrogate2).mean()
|
||
|
||
# ---- L^VF 价值损失 ----
|
||
current_values = self.value_net(states)
|
||
value_loss = F.mse_loss(current_values, returns)
|
||
|
||
# ---- 熵奖励 ----
|
||
entropy_bonus = entropy.mean()
|
||
|
||
# ---- 总损失(公式9)----
|
||
loss = actor_loss + self.vf_coef * value_loss - self.entropy_coef * entropy_bonus
|
||
|
||
# ---- 梯度更新 ----
|
||
self.actor_optimizer.zero_grad()
|
||
self.value_optimizer.zero_grad()
|
||
loss.backward()
|
||
torch.nn.utils.clip_grad_norm_(self.actor.parameters(), self.max_grad_norm)
|
||
torch.nn.utils.clip_grad_norm_(self.value_net.parameters(), self.max_grad_norm)
|
||
self.actor_optimizer.step()
|
||
self.value_optimizer.step()
|
||
|
||
actor_losses.append(actor_loss.item())
|
||
value_losses.append(value_loss.item())
|
||
entropy_bonuses.append(entropy_bonus.item())
|
||
|
||
# 清空缓冲区,准备下一轮收集
|
||
self.rollout.clear()
|
||
|
||
return {
|
||
'actor_loss': np.mean(actor_losses),
|
||
'value_loss': np.mean(value_losses),
|
||
'entropy': np.mean(entropy_bonuses),
|
||
}
|