File size: 6,364 Bytes
ae73c7f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 | """
Proximal Policy Optimization (PPO) for the multi-step Poincaré environment.
Uses the existing HyperbolicCritic and HierarchicalHyperbolicPredictor.
Compatible with Optuna best HPs and the continual-learning module.
"""
from __future__ import annotations
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.distributions import Normal
import numpy as np
from typing import List, Tuple, Optional
from src.model import HierarchicalHyperbolicPredictor, HyperbolicCritic, MultiScaleEncoder
from src.poincare import PoincareBall8D
class PoincareActor(nn.Module):
def __init__(self, obs_dim: int = 8, action_dim: int = 8, hidden: int = 64):
super().__init__()
self.net = nn.Sequential(
nn.Linear(obs_dim, hidden),
nn.GELU(),
nn.Linear(hidden, hidden),
nn.GELU(),
)
self.mean = nn.Linear(hidden, action_dim)
self.log_std = nn.Parameter(torch.zeros(action_dim) - 0.5)
def forward(self, obs: torch.Tensor):
h = self.net(obs)
mu = self.mean(h)
std = self.log_std.exp().expand_as(mu)
return mu, std
def dist(self, obs: torch.Tensor):
mu, std = self.forward(obs)
return Normal(mu, std)
def act(self, obs: torch.Tensor, deterministic: bool = False):
dist = self.dist(obs)
if deterministic:
action = dist.mean
else:
action = dist.sample()
logp = dist.log_prob(action).sum(-1)
return action, logp
class PPOTrainer:
def __init__(
self,
actor: PoincareActor,
critic: HyperbolicCritic,
poincare: PoincareBall8D,
lr: float = 3e-4,
clip_eps: float = 0.2,
gamma: float = 0.95,
gae_lambda: float = 0.95,
entropy_coef: float = 0.01,
value_coef: float = 0.5,
max_grad_norm: float = 1.0,
device: str = "cpu",
):
self.actor = actor.to(device)
self.critic = critic.to(device)
self.poincare = poincare
self.clip_eps = clip_eps
self.gamma = gamma
self.gae_lambda = gae_lambda
self.entropy_coef = entropy_coef
self.value_coef = value_coef
self.max_grad_norm = max_grad_norm
self.device = device
self.opt = torch.optim.Adam(
list(self.actor.parameters()) + list(self.critic.parameters()), lr=lr
)
def collect_rollout(self, env, n_steps: int = 64):
obs_list, act_list, logp_list, rew_list, val_list, done_list = [], [], [], [], [], []
obs, _ = env.reset()
for _ in range(n_steps):
obs_t = torch.as_tensor(obs, device=self.device, dtype=torch.float32).unsqueeze(0)
with torch.no_grad():
action, logp = self.actor.act(obs_t)
# critic expects ball points
z_ball = self.poincare.expmap0(obs_t)
value = self.critic(z_ball)
next_obs, reward, term, trunc, info = env.step(action.squeeze(0).cpu().numpy())
done = term or trunc
obs_list.append(obs)
act_list.append(action.squeeze(0).cpu().numpy())
logp_list.append(logp.item())
rew_list.append(reward)
val_list.append(value.item())
done_list.append(float(done))
obs = next_obs
if done:
obs, _ = env.reset()
return {
"obs": np.array(obs_list, dtype=np.float32),
"actions": np.array(act_list, dtype=np.float32),
"logp": np.array(logp_list, dtype=np.float32),
"rewards": np.array(rew_list, dtype=np.float32),
"values": np.array(val_list, dtype=np.float32),
"dones": np.array(done_list, dtype=np.float32),
}
def compute_gae(self, rewards, values, dones):
advantages = np.zeros_like(rewards)
last_gae = 0.0
for t in reversed(range(len(rewards))):
next_val = 0.0 if t == len(rewards) - 1 else values[t + 1]
next_nonterminal = 1.0 - dones[t]
delta = rewards[t] + self.gamma * next_val * next_nonterminal - values[t]
last_gae = delta + self.gamma * self.gae_lambda * next_nonterminal * last_gae
advantages[t] = last_gae
returns = advantages + values
return advantages, returns
def update(self, rollout, n_epochs: int = 4, batch_size: int = 32):
adv, ret = self.compute_gae(rollout["rewards"], rollout["values"], rollout["dones"])
adv = (adv - adv.mean()) / (adv.std() + 1e-8)
obs = torch.as_tensor(rollout["obs"], device=self.device)
actions = torch.as_tensor(rollout["actions"], device=self.device)
old_logp = torch.as_tensor(rollout["logp"], device=self.device)
adv_t = torch.as_tensor(adv, device=self.device)
ret_t = torch.as_tensor(ret, device=self.device)
n = len(obs)
indices = np.arange(n)
losses = []
for _ in range(n_epochs):
np.random.shuffle(indices)
for start in range(0, n, batch_size):
idx = indices[start : start + batch_size]
o = obs[idx]
a = actions[idx]
olp = old_logp[idx]
ad = adv_t[idx]
rt = ret_t[idx]
dist = self.actor.dist(o)
new_logp = dist.log_prob(a).sum(-1)
entropy = dist.entropy().sum(-1).mean()
ratio = (new_logp - olp).exp()
surr1 = ratio * ad
surr2 = torch.clamp(ratio, 1 - self.clip_eps, 1 + self.clip_eps) * ad
policy_loss = -torch.min(surr1, surr2).mean()
z_ball = self.poincare.expmap0(o)
values = self.critic(z_ball)
value_loss = F.mse_loss(values, rt)
loss = policy_loss + self.value_coef * value_loss - self.entropy_coef * entropy
self.opt.zero_grad()
loss.backward()
nn.utils.clip_grad_norm_(
list(self.actor.parameters()) + list(self.critic.parameters()),
self.max_grad_norm,
)
self.opt.step()
losses.append(loss.item())
return float(np.mean(losses)) if losses else 0.0
|