175 lines
6.0 KiB
Python
175 lines
6.0 KiB
Python
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)
|