poincare-hyper / src /ppo.py
DHDRL's picture
Rename ppo.py to src/ppo.py
f4d5e92 verified
Raw
History Blame Contribute Delete
6.36 kB
"""
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