重构 PPO/TRPO 训练流程并添加对比绘图
- PPO: 改为 Actor/Critic 联合小批量训练,新增梯度裁剪 (max_grad_norm), 分离 actor_lr/critic_lr,添加 get_value(),GAE 部分补充论文公式注释 - TRPO: 添加 get_value(),调整 tau 从 0.97 到 0.95 - Networks: 移除 PolicyNet 输出层的 tanh,初始化 log_std=0 以增强探索 - Main: 抽取 train_agent() 通用训练函数,新增 TRPO 训练和 PPO vs TRPO 对比曲线图(原始曲线 + 滑动平均平滑曲线)
This commit is contained in:
+5
-80
@@ -6,7 +6,7 @@ from torch.nn.utils import parameters_to_vector, vector_to_parameters
|
||||
from networks import PolicyNetwork, ValueNetwork
|
||||
|
||||
class TRPOAgent:
|
||||
def __init__(self, state_dim, action_dim, action_bound, hidden_dim=128, kl_margin=0.01, gamma=0.99, tau=0.97, cg_iters=10):
|
||||
def __init__(self, state_dim, action_dim, action_bound, hidden_dim=128, kl_margin=0.01, gamma=0.99, tau=0.95, cg_iters=10):
|
||||
self.gamma = gamma
|
||||
self.tau = tau
|
||||
self.kl_margin = kl_margin
|
||||
@@ -26,6 +26,10 @@ class TRPOAgent:
|
||||
action = dist.sample()
|
||||
return action.squeeze(0).numpy()
|
||||
|
||||
def get_value(self, state):
|
||||
with torch.no_grad():
|
||||
return self.critic(torch.FloatTensor(state).unsqueeze(0)).squeeze().item()
|
||||
|
||||
def _compute_advantages(self, rewards, values, masks):
|
||||
"""
|
||||
使用广义优势估计 (GAE) 计算优势函数
|
||||
@@ -158,82 +162,3 @@ class TRPOAgent:
|
||||
|
||||
if not success:
|
||||
vector_to_parameters(old_params, self.actor.parameters())
|
||||
"""
|
||||
利用收集到的轨迹数据更新 Actor 和 Critic
|
||||
"""
|
||||
states = torch.FloatTensor(np.array([m[0] for m in memory]))
|
||||
actions = torch.FloatTensor(np.array([m[1] for m in memory]))
|
||||
rewards = [m[2] for m in memory]
|
||||
next_states = torch.FloatTensor(np.array([m[3] for m in memory]))
|
||||
masks = [m[4] for m in memory]
|
||||
|
||||
# 1. 拟合价值网络 (Critic)
|
||||
with torch.no_grad():
|
||||
values = self.critic(states).squeeze().numpy().tolist()
|
||||
next_value = self.critic(next_states[-1].unsqueeze(0)).squeeze().item()
|
||||
values.append(next_value)
|
||||
|
||||
returns = self._compute_advantages(rewards, values, masks)
|
||||
returns = torch.FloatTensor(returns)
|
||||
values = torch.FloatTensor(values[:-1])
|
||||
advantages = returns - values
|
||||
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
|
||||
|
||||
for _ in range(10):
|
||||
critic_loss = F.mse_loss(self.critic(states).squeeze(), returns)
|
||||
self.critic_optimizer.zero_grad()
|
||||
critic_loss.backward()
|
||||
self.critic_optimizer.step()
|
||||
|
||||
# --- 关键修复 3:在截断梯度的环境下生成严格的旧分布 ---
|
||||
with torch.no_grad():
|
||||
old_mean, old_std = self.actor(states)
|
||||
old_dist = Normal(old_mean, old_std)
|
||||
old_log_probs = old_dist.log_prob(actions).sum(dim=1)
|
||||
|
||||
def compute_surrogate_loss():
|
||||
# 计算替代目标函数 (Surrogate Objective)
|
||||
dist = self.actor.evaluate(states)
|
||||
log_probs = dist.log_prob(actions).sum(dim=1)
|
||||
ratio = torch.exp(log_probs - old_log_probs)
|
||||
surrogate_loss = (ratio * advantages).mean()
|
||||
return surrogate_loss, dist
|
||||
|
||||
surrogate_loss, dist = compute_surrogate_loss()
|
||||
|
||||
loss_grad = torch.autograd.grad(surrogate_loss, self.actor.parameters())
|
||||
loss_grad_flat = torch.cat([grad.view(-1) for grad in loss_grad])
|
||||
|
||||
# 传入 old_dist,确保海森矩阵计算包含准确的曲率信息
|
||||
step_dir = self._conjugate_gradient(states, old_dist, loss_grad_flat, self.cg_iters)
|
||||
|
||||
shs = 0.5 * torch.dot(step_dir, self._hessian_vector_product(states, old_dist, step_dir))
|
||||
|
||||
# 增加数值保护:防止 shs 出现负数或极小值导致报错
|
||||
if shs < 1e-8:
|
||||
return
|
||||
|
||||
lm = torch.sqrt(shs / self.kl_margin)
|
||||
fullstep = step_dir / lm
|
||||
|
||||
old_params = parameters_to_vector(self.actor.parameters())
|
||||
|
||||
# 线性搜索 (Line Search)
|
||||
success = False
|
||||
step_size = 1.0
|
||||
for _ in range(10):
|
||||
new_params = old_params + step_size * fullstep
|
||||
vector_to_parameters(new_params, self.actor.parameters())
|
||||
|
||||
with torch.no_grad():
|
||||
new_surrogate_loss, new_dist = compute_surrogate_loss()
|
||||
kl = torch.distributions.kl_divergence(old_dist, new_dist).mean()
|
||||
|
||||
if new_surrogate_loss > surrogate_loss and kl <= self.kl_margin:
|
||||
success = True
|
||||
break
|
||||
|
||||
step_size *= 0.5
|
||||
|
||||
if not success:
|
||||
vector_to_parameters(old_params, self.actor.parameters())
|
||||
|
||||
Reference in New Issue
Block a user