File size: 4,430 Bytes
9f8cf99
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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