添加 PPO 算法实现及相关配置,更新训练入口以支持 SAC 和 PPO 模式

This commit is contained in:
2026-04-04 17:50:02 +08:00
parent f3d2d1a85f
commit a2ce5073c5
10 changed files with 675 additions and 73 deletions
+76
View File
@@ -96,3 +96,79 @@ class Actor(nn.Module):
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)