Files
RL_SAC/algorithms/ppo.py
T

140 lines
5.5 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 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),
}