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