修复并行采样并完善训练文档
This commit is contained in:
+72
-27
@@ -2,7 +2,7 @@ import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch.distributions import Normal
|
||||
|
||||
from networks import PolicyNetwork, ValueNetwork
|
||||
|
||||
|
||||
@@ -32,6 +32,7 @@ class PPOAgent:
|
||||
k_epochs=10,
|
||||
minibatch_size=64,
|
||||
max_grad_norm=0.5,
|
||||
device="cpu",
|
||||
):
|
||||
self.gamma = gamma
|
||||
self.tau = tau
|
||||
@@ -39,37 +40,58 @@ class PPOAgent:
|
||||
self.k_epochs = k_epochs # 每次更新对同一批数据的训练轮数
|
||||
self.minibatch_size = minibatch_size
|
||||
self.max_grad_norm = max_grad_norm # 梯度裁剪阈值
|
||||
self.device = torch.device(device)
|
||||
|
||||
self.actor = PolicyNetwork(state_dim, action_dim, action_bound, hidden_dim)
|
||||
self.critic = ValueNetwork(state_dim, hidden_dim)
|
||||
self.actor = PolicyNetwork(state_dim, action_dim, action_bound, hidden_dim).to(self.device)
|
||||
self.critic = ValueNetwork(state_dim, hidden_dim).to(self.device)
|
||||
|
||||
self.actor_optimizer = torch.optim.Adam(self.actor.parameters(), lr=actor_lr)
|
||||
self.critic_optimizer = torch.optim.Adam(self.critic.parameters(), lr=critic_lr)
|
||||
|
||||
def _to_tensor(self, array_like):
|
||||
return torch.as_tensor(array_like, dtype=torch.float32, device=self.device)
|
||||
|
||||
def get_action(self, state):
|
||||
state_tensor = torch.FloatTensor(state).unsqueeze(0)
|
||||
state_tensor = self._to_tensor(state)
|
||||
single_state = state_tensor.ndim == 1
|
||||
if single_state:
|
||||
state_tensor = state_tensor.unsqueeze(0)
|
||||
|
||||
with torch.no_grad():
|
||||
dist = self.actor.evaluate(state_tensor)
|
||||
action = dist.sample()
|
||||
return action.squeeze(0).numpy()
|
||||
|
||||
action_np = action.detach().cpu().numpy()
|
||||
return action_np[0] if single_state else action_np
|
||||
|
||||
def get_value(self, state):
|
||||
with torch.no_grad():
|
||||
return self.critic(torch.FloatTensor(state).unsqueeze(0)).squeeze().item()
|
||||
state_tensor = self._to_tensor(state)
|
||||
single_state = state_tensor.ndim == 1
|
||||
if single_state:
|
||||
state_tensor = state_tensor.unsqueeze(0)
|
||||
|
||||
def _compute_advantages(self, rewards, values, masks):
|
||||
with torch.no_grad():
|
||||
values = self.critic(state_tensor).squeeze(-1)
|
||||
|
||||
values_np = values.detach().cpu().numpy()
|
||||
return float(values_np[0]) if single_state else values_np
|
||||
|
||||
def _compute_advantages(self, rewards, values, next_values, masks):
|
||||
"""
|
||||
GAE (Generalized Advantage Estimation)
|
||||
Eq.(11): Â_t = δ_t + (γλ)δ_{t+1} + ... + (γλ)^{T-t+1} δ_{T-1}
|
||||
Eq.(12): δ_t = r_t + γV(s_{t+1}) - V(s_t)
|
||||
"""
|
||||
returns = []
|
||||
gae = 0
|
||||
for i in reversed(range(len(rewards))):
|
||||
delta = rewards[i] + self.gamma * values[i + 1] * masks[i] - values[i]
|
||||
advantages = torch.zeros_like(rewards, device=self.device)
|
||||
gae = torch.zeros(rewards.size(1), dtype=torch.float32, device=self.device)
|
||||
|
||||
for i in reversed(range(rewards.size(0))):
|
||||
delta = rewards[i] + self.gamma * next_values[i] * masks[i] - values[i]
|
||||
gae = delta + self.gamma * self.tau * masks[i] * gae
|
||||
returns.insert(0, gae + values[i])
|
||||
return returns
|
||||
advantages[i] = gae
|
||||
|
||||
returns = advantages + values
|
||||
return returns, advantages
|
||||
|
||||
def update(self, memory):
|
||||
"""
|
||||
@@ -80,25 +102,32 @@ class PPOAgent:
|
||||
L = L^CLIP - c1 * L^VF (Eq.9,c2=0 不加熵)
|
||||
2. 梯度裁剪 (max_grad_norm)
|
||||
"""
|
||||
if not memory:
|
||||
return {}
|
||||
|
||||
# ---------- 1. 准备数据 ----------
|
||||
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]
|
||||
states = self._to_tensor(np.asarray(memory["states"]))
|
||||
actions = self._to_tensor(np.asarray(memory["actions"]))
|
||||
rewards = self._to_tensor(np.asarray(memory["rewards"]))
|
||||
next_states = self._to_tensor(np.asarray(memory["next_states"]))
|
||||
masks = self._to_tensor(np.asarray(memory["masks"]))
|
||||
|
||||
# ---------- 2. GAE 优势估计 ----------
|
||||
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)
|
||||
rollout_steps, num_envs = rewards.shape
|
||||
flat_states = states.reshape(rollout_steps * num_envs, -1)
|
||||
flat_next_states = next_states.reshape(rollout_steps * num_envs, -1)
|
||||
values = self.critic(flat_states).squeeze(-1).reshape(rollout_steps, num_envs)
|
||||
next_values = self.critic(flat_next_states).squeeze(-1).reshape(rollout_steps, num_envs)
|
||||
|
||||
returns = self._compute_advantages(rewards, values, masks)
|
||||
returns = torch.FloatTensor(returns)
|
||||
values = torch.FloatTensor(values[:-1])
|
||||
advantages = returns - values
|
||||
returns, advantages = self._compute_advantages(rewards, values, next_values, masks)
|
||||
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
|
||||
|
||||
states = flat_states
|
||||
actions = actions.reshape(rollout_steps * num_envs, -1)
|
||||
returns = returns.reshape(rollout_steps * num_envs)
|
||||
advantages = advantages.reshape(rollout_steps * num_envs)
|
||||
|
||||
# ---------- 3. 记录旧策略 π_θold ----------
|
||||
with torch.no_grad():
|
||||
old_dist = self.actor.evaluate(states)
|
||||
@@ -107,13 +136,17 @@ class PPOAgent:
|
||||
# ---------- 4. K epochs × minibatch 联合训练 ----------
|
||||
dataset_size = states.size(0)
|
||||
indices = np.arange(dataset_size)
|
||||
policy_loss_value = 0.0
|
||||
critic_loss_value = 0.0
|
||||
entropy_value = 0.0
|
||||
minibatch_updates = 0
|
||||
|
||||
for _ in range(self.k_epochs):
|
||||
np.random.shuffle(indices)
|
||||
|
||||
for start in range(0, dataset_size, self.minibatch_size):
|
||||
end = start + self.minibatch_size
|
||||
mb_idx = indices[start:end]
|
||||
mb_idx = torch.as_tensor(indices[start:end], device=self.device, dtype=torch.long)
|
||||
|
||||
mb_states = states[mb_idx]
|
||||
mb_actions = actions[mb_idx]
|
||||
@@ -144,3 +177,15 @@ class PPOAgent:
|
||||
self.critic_optimizer.zero_grad()
|
||||
critic_loss.backward()
|
||||
self.critic_optimizer.step()
|
||||
|
||||
policy_loss_value += float(policy_loss.detach().item())
|
||||
critic_loss_value += float(critic_loss.detach().item())
|
||||
entropy_value += float(dist.entropy().sum(dim=1).mean().detach().item())
|
||||
minibatch_updates += 1
|
||||
|
||||
divisor = max(1, minibatch_updates)
|
||||
return {
|
||||
"policy_loss": policy_loss_value / divisor,
|
||||
"critic_loss": critic_loss_value / divisor,
|
||||
"entropy": entropy_value / divisor,
|
||||
}
|
||||
|
||||
+92
-35
@@ -3,44 +3,79 @@ import torch
|
||||
import torch.nn.functional as F
|
||||
from torch.distributions import Normal
|
||||
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.95, 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,
|
||||
critic_epochs=40,
|
||||
device="cpu",
|
||||
):
|
||||
self.gamma = gamma
|
||||
self.tau = tau
|
||||
self.kl_margin = kl_margin
|
||||
self.kl_margin = kl_margin
|
||||
self.cg_iters = cg_iters
|
||||
self.critic_epochs = critic_epochs
|
||||
self.device = torch.device(device)
|
||||
|
||||
self.actor = PolicyNetwork(state_dim, action_dim, action_bound, hidden_dim)
|
||||
self.critic = ValueNetwork(state_dim, hidden_dim)
|
||||
self.actor = PolicyNetwork(state_dim, action_dim, action_bound, hidden_dim).to(self.device)
|
||||
self.critic = ValueNetwork(state_dim, hidden_dim).to(self.device)
|
||||
self.critic_optimizer = torch.optim.Adam(self.critic.parameters(), lr=1e-3)
|
||||
|
||||
def _to_tensor(self, array_like):
|
||||
return torch.as_tensor(array_like, dtype=torch.float32, device=self.device)
|
||||
|
||||
def get_action(self, state):
|
||||
"""
|
||||
根据当前状态采样动作
|
||||
"""
|
||||
state_tensor = torch.FloatTensor(state).unsqueeze(0)
|
||||
state_tensor = self._to_tensor(state)
|
||||
single_state = state_tensor.ndim == 1
|
||||
if single_state:
|
||||
state_tensor = state_tensor.unsqueeze(0)
|
||||
|
||||
with torch.no_grad():
|
||||
dist = self.actor.evaluate(state_tensor)
|
||||
action = dist.sample()
|
||||
return action.squeeze(0).numpy()
|
||||
|
||||
action_np = action.detach().cpu().numpy()
|
||||
return action_np[0] if single_state else action_np
|
||||
|
||||
def get_value(self, state):
|
||||
with torch.no_grad():
|
||||
return self.critic(torch.FloatTensor(state).unsqueeze(0)).squeeze().item()
|
||||
state_tensor = self._to_tensor(state)
|
||||
single_state = state_tensor.ndim == 1
|
||||
if single_state:
|
||||
state_tensor = state_tensor.unsqueeze(0)
|
||||
|
||||
def _compute_advantages(self, rewards, values, masks):
|
||||
with torch.no_grad():
|
||||
values = self.critic(state_tensor).squeeze(-1)
|
||||
|
||||
values_np = values.detach().cpu().numpy()
|
||||
return float(values_np[0]) if single_state else values_np
|
||||
|
||||
def _compute_advantages(self, rewards, values, next_values, masks):
|
||||
"""
|
||||
使用广义优势估计 (GAE) 计算优势函数
|
||||
"""
|
||||
returns = []
|
||||
gae = 0
|
||||
for i in reversed(range(len(rewards))):
|
||||
delta = rewards[i] + self.gamma * values[i + 1] * masks[i] - values[i]
|
||||
advantages = torch.zeros_like(rewards, device=self.device)
|
||||
gae = torch.zeros(rewards.size(1), dtype=torch.float32, device=self.device)
|
||||
|
||||
for i in reversed(range(rewards.size(0))):
|
||||
delta = rewards[i] + self.gamma * next_values[i] * masks[i] - values[i]
|
||||
gae = delta + self.gamma * self.tau * masks[i] * gae
|
||||
returns.insert(0, gae + values[i])
|
||||
return returns
|
||||
advantages[i] = gae
|
||||
|
||||
returns = advantages + values
|
||||
return returns, advantages
|
||||
|
||||
# --- 关键修复 1:将固定的 old_dist 作为参数传入 ---
|
||||
def _hessian_vector_product(self, states, old_dist, vector, damping=0.1):
|
||||
@@ -58,7 +93,7 @@ class TRPOAgent:
|
||||
|
||||
# 与给定向量点乘
|
||||
kl_v = (flat_grad_kl * vector).sum()
|
||||
|
||||
|
||||
# 二阶导数
|
||||
grads = torch.autograd.grad(kl_v, self.actor.parameters())
|
||||
flat_grad_grad_kl = torch.cat([grad.contiguous().view(-1) for grad in grads])
|
||||
@@ -88,30 +123,39 @@ class TRPOAgent:
|
||||
return x
|
||||
|
||||
def update(self, memory):
|
||||
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]
|
||||
if not memory:
|
||||
return {}
|
||||
|
||||
states = self._to_tensor(np.asarray(memory["states"]))
|
||||
actions = self._to_tensor(np.asarray(memory["actions"]))
|
||||
rewards = self._to_tensor(np.asarray(memory["rewards"]))
|
||||
next_states = self._to_tensor(np.asarray(memory["next_states"]))
|
||||
masks = self._to_tensor(np.asarray(memory["masks"]))
|
||||
|
||||
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
|
||||
rollout_steps, num_envs = rewards.shape
|
||||
flat_states = states.reshape(rollout_steps * num_envs, -1)
|
||||
flat_next_states = next_states.reshape(rollout_steps * num_envs, -1)
|
||||
values = self.critic(flat_states).squeeze(-1).reshape(rollout_steps, num_envs)
|
||||
next_values = self.critic(flat_next_states).squeeze(-1).reshape(rollout_steps, num_envs)
|
||||
|
||||
returns, advantages = self._compute_advantages(rewards, values, next_values, masks)
|
||||
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
|
||||
|
||||
states = flat_states
|
||||
actions = actions.reshape(rollout_steps * num_envs, -1)
|
||||
returns = returns.reshape(rollout_steps * num_envs)
|
||||
advantages = advantages.reshape(rollout_steps * num_envs)
|
||||
|
||||
# 【修复】:加大 Critic 的训练力度,从 10 提升到 40 Epochs
|
||||
# 确保裁判的眼光足够准确
|
||||
for _ in range(40):
|
||||
critic_loss_value = 0.0
|
||||
for _ in range(self.critic_epochs):
|
||||
critic_loss = F.mse_loss(self.critic(states).squeeze(), returns)
|
||||
self.critic_optimizer.zero_grad()
|
||||
critic_loss.backward()
|
||||
self.critic_optimizer.step()
|
||||
critic_loss_value += float(critic_loss.detach().item())
|
||||
|
||||
with torch.no_grad():
|
||||
old_mean, old_std = self.actor(states)
|
||||
@@ -139,26 +183,39 @@ class TRPOAgent:
|
||||
|
||||
lm = torch.sqrt(shs / self.kl_margin)
|
||||
fullstep = step_dir / lm
|
||||
|
||||
|
||||
old_params = parameters_to_vector(self.actor.parameters())
|
||||
|
||||
# 线性搜索
|
||||
success = False
|
||||
step_size = 1.0
|
||||
kl_value = 0.0
|
||||
surrogate_value = float(surrogate_loss.detach().item())
|
||||
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 - 1e-8 and kl <= self.kl_margin * 1.5:
|
||||
success = True
|
||||
surrogate_value = float(new_surrogate_loss.item())
|
||||
kl_value = float(kl.item())
|
||||
break
|
||||
|
||||
|
||||
step_size *= 0.5
|
||||
|
||||
kl_value = float(kl.item())
|
||||
|
||||
if not success:
|
||||
vector_to_parameters(old_params, self.actor.parameters())
|
||||
|
||||
return {
|
||||
"critic_loss": critic_loss_value / max(1, self.critic_epochs),
|
||||
"surrogate_loss": surrogate_value,
|
||||
"kl": kl_value,
|
||||
"line_search_success": float(success),
|
||||
"step_scale": float(step_size if success else 0.0),
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user