Files

175 lines
6.0 KiB
Python
Raw Permalink 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 as nn
import torch.nn.functional as F
from torch.distributions import Normal
# 确保代码可以在 GPU 上跑(如果有的话)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
class Critic(nn.Module):
def __init__(self, state_dim, action_dim):
super(Critic, self).__init__()
# Q1 网络架构
self.l1 = nn.Linear(state_dim + action_dim, 256)
self.l2 = nn.Linear(256, 256)
self.l3 = nn.Linear(256, 1)
# Q2 网络架构(和 Q1 完全一样,但参数是独立初始化的)
self.l4 = nn.Linear(state_dim + action_dim, 256)
self.l5 = nn.Linear(256, 256)
self.l6 = nn.Linear(256, 1)
def forward(self, state, action):
# 把状态和动作拼接在一起,作为 Q 网络的输入
sa = torch.cat([state, action], 1)
# Q1 的前向传播
q1 = F.relu(self.l1(sa))
q1 = F.relu(self.l2(q1))
q1 = self.l3(q1)
# Q2 的前向传播
q2 = F.relu(self.l4(sa))
q2 = F.relu(self.l5(q2))
q2 = self.l6(q2)
# 训练时,我们需要同时返回两个 Q 值来算误差
return q1, q2
# 定义标准差的上下界,防止网络输出极端值导致计算崩溃(NaN)
LOG_SIG_MAX = 2
LOG_SIG_MIN = -20
class Actor(nn.Module):
def __init__(self, state_dim, action_dim, max_action):
super(Actor, self).__init__()
# 共享特征提取层
self.l1 = nn.Linear(state_dim, 256)
self.l2 = nn.Linear(256, 256)
# 均值输出层
self.mean_linear = nn.Linear(256, action_dim)
# 对数标准差输出层(预测 log_std 比直接预测 std 更好优化)
self.log_std_linear = nn.Linear(256, action_dim)
# 动作的最大物理边界(比如 Pendulum 的力矩最大是 2.0
self.max_action = max_action
def forward(self, state):
x = F.relu(self.l1(state))
x = F.relu(self.l2(x))
mean = self.mean_linear(x)
log_std = self.log_std_linear(x)
# 限制 log_std 的范围,防止数值不稳定
log_std = torch.clamp(log_std, min=LOG_SIG_MIN, max=LOG_SIG_MAX)
return mean, log_std
def sample(self, state):
mean, log_std = self.forward(state)
std = log_std.exp()
# 构造一个高斯分布
normal = Normal(mean, std)
# normal.rsample() 内部执行的就是 a = mean + std * epsilon (其中 epsilon 是标准正态噪声)
# 这就是公式 (11) 的代码实现!用 rsample 才能让梯度传导回网络。
x_t = normal.rsample()
# 把动作压缩到 [-1, 1] 区间(这就是论文附录 C 里的 tanh 压扁函数)
y_t = torch.tanh(x_t)
# 映射到真实的物理动作区间,比如 [-2.0, 2.0]
action = y_t * self.max_action
# 计算这个动作的对数概率 log(pi(a|s)),用于后面算熵
# 这行公式对应论文附录 C 的公式 (21),是应用 tanh 后的概率修正
log_prob = normal.log_prob(x_t)
log_prob -= torch.log(self.max_action * (1 - y_t.pow(2)) + 1e-6)
log_prob = log_prob.sum(1, keepdim=True)
# mean 经过 tanh 就是测试时用的确定性动作
mean = torch.tanh(mean) * self.max_action
return action, log_prob, mean
# ===========================================================================
# PPO 专用网络
# ===========================================================================
class PPOActor(nn.Module):
"""
PPO 策略网络 — 标准高斯策略(不使用 tanh 压缩)
与 SAC Actor 的核心区别:
- SAC 需要 tanh squashing + log_prob Jacobian 修正来精确计算熵
- PPO 直接使用高斯分布的 log_prob / entropy,再 clamp 到合法范围
- log_std 是全局可学习参数(不依赖状态),更稳定
"""
def __init__(self, state_dim, action_dim, max_action):
super(PPOActor, self).__init__()
self.max_action = max_action
self.net = nn.Sequential(
nn.Linear(state_dim, 64),
nn.Tanh(),
nn.Linear(64, 64),
nn.Tanh(),
nn.Linear(64, action_dim),
)
# 全局可学习的对数标准差,初始化为 0 → std=1.0(足够的初始探索)
self.log_std = nn.Parameter(torch.zeros(action_dim))
def forward(self, state):
mean = self.net(state)
std = self.log_std.exp().expand_as(mean)
return mean, std
def get_dist(self, state):
mean, std = self.forward(state)
return Normal(mean, std)
def sample(self, state):
"""
采样动作,直接 clamp 到 [-max_action, max_action]
返回: (action, log_prob)
"""
dist = self.get_dist(state)
action_unbounded = dist.rsample()
# 直接 clamp(不做 tanh,避免 log_prob 被 Jacobian 修正污染)
action = torch.clamp(action_unbounded, -self.max_action, self.max_action)
log_prob = dist.log_prob(action_unbounded).sum(1, keepdim=True)
return action, log_prob, action_unbounded
def evaluate(self, state, action_unbounded):
"""
给定之前保存的未裁剪动作,重新计算 log_prob 和熵(用于 PPO K 轮更新)
"""
dist = self.get_dist(state)
log_prob = dist.log_prob(action_unbounded).sum(1, keepdim=True)
entropy = dist.entropy().sum(1, keepdim=True)
return log_prob, entropy
class ValueNet(nn.Module):
"""
PPO 价值网络:估计状态价值函数 V(s)
"""
def __init__(self, state_dim):
super(ValueNet, self).__init__()
self.net = nn.Sequential(
nn.Linear(state_dim, 256),
nn.Tanh(),
nn.Linear(256, 256),
nn.Tanh(),
nn.Linear(256, 1),
)
def forward(self, state):
return self.net(state)