Gyanateet Dutta
Fix Space loading: direct Streamlit, lazy imports, ReNova page, fix deps
9f8cf99
Raw
History Blame Contribute Delete
4.43 kB
"""PEPG-style optimizer for policy search over control parameters."""
from __future__ import annotations
from dataclasses import asdict, dataclass
from pathlib import Path
import json
import numpy as np
@dataclass
class PEPGState:
"""Serializable optimizer state."""
mean: list[float]
sigma: list[float]
learning_rate: float
sigma_learning_rate: float
iteration: int
seed: int
rng_state: dict
class PEPGOptimizer:
"""Parameter-Exploring Policy Gradients (PEPG) with diagonal Gaussian."""
def __init__(
self,
parameter_dim: int,
*,
seed: int = 0,
init_sigma: float = 0.1,
learning_rate: float = 0.05,
sigma_learning_rate: float = 0.02,
sigma_min: float = 1e-4,
):
if parameter_dim <= 0:
raise ValueError("parameter_dim must be positive")
self.parameter_dim = parameter_dim
self.mean = np.zeros(parameter_dim, dtype=np.float64)
self.sigma = np.full(parameter_dim, init_sigma, dtype=np.float64)
self.learning_rate = float(learning_rate)
self.sigma_learning_rate = float(sigma_learning_rate)
self.sigma_min = float(sigma_min)
self.iteration = 0
self.seed = int(seed)
self._rng = np.random.default_rng(self.seed)
def ask(self, population_size: int) -> tuple[np.ndarray, np.ndarray]:
"""Sample antithetic population and return (candidates, perturbations)."""
if population_size <= 0 or population_size % 2 != 0:
raise ValueError("population_size must be a positive even integer")
half = population_size // 2
eps = self._rng.standard_normal((half, self.parameter_dim))
perturb = np.vstack([eps, -eps]) * self.sigma
candidates = self.mean + perturb
return candidates, perturb
def tell(self, perturbations: np.ndarray, rewards: np.ndarray) -> None:
"""Update mean and sigma from population rewards."""
perturbations = np.asarray(perturbations, dtype=np.float64)
rewards = np.asarray(rewards, dtype=np.float64)
if perturbations.shape[0] != rewards.shape[0]:
raise ValueError("perturbations and rewards length mismatch")
centered = rewards - np.mean(rewards)
scale = np.std(rewards)
if scale > 0:
centered = centered / (scale + 1e-8)
grad_mean = perturbations.T @ centered / len(rewards)
self.mean += self.learning_rate * grad_mean
sigma_grad = (
((perturbations**2 - self.sigma**2) / np.maximum(self.sigma, 1e-8)).T
@ centered
/ len(rewards)
)
self.sigma = np.maximum(self.sigma + self.sigma_learning_rate * sigma_grad, self.sigma_min)
self.iteration += 1
def state_dict(self) -> PEPGState:
"""Return serializable state."""
return PEPGState(
mean=self.mean.tolist(),
sigma=self.sigma.tolist(),
learning_rate=self.learning_rate,
sigma_learning_rate=self.sigma_learning_rate,
iteration=self.iteration,
seed=self.seed,
rng_state=self._rng.bit_generator.state,
)
def load_state_dict(self, state: PEPGState) -> None:
"""Load optimizer from state."""
self.mean = np.asarray(state.mean, dtype=np.float64)
self.sigma = np.asarray(state.sigma, dtype=np.float64)
self.learning_rate = state.learning_rate
self.sigma_learning_rate = state.sigma_learning_rate
self.iteration = state.iteration
self.seed = state.seed
self._rng = np.random.default_rng()
self._rng.bit_generator.state = state.rng_state
def save_checkpoint(self, path: str | Path) -> None:
"""Persist optimizer checkpoint as JSON."""
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8") as f:
json.dump(asdict(self.state_dict()), f, indent=2)
@classmethod
def load_checkpoint(cls, path: str | Path) -> "PEPGOptimizer":
"""Restore optimizer from checkpoint."""
path = Path(path)
with path.open("r", encoding="utf-8") as f:
payload = json.load(f)
state = PEPGState(**payload)
opt = cls(parameter_dim=len(state.mean), seed=state.seed)
opt.load_state_dict(state)
return opt