| """ |
| 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) |
| |
| 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 |
|
|