添加 PPO 算法实现及相关配置,更新训练入口以支持 SAC 和 PPO 模式
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user