Skip to main content

fcmaes_core/
pgpe.rs

1//! PGPE — Rust port of the C++ `pgpe.cpp`.
2//!
3//! Parameter-exploring Policy Gradients with an ADAM center/baseline update and
4//! symmetric ("mirrored") sampling
5//! (<http://mediatum.ub.tum.de/doc/1099128/631352.pdf>, derived from EvoJax).
6//! C++-only in the original (no pure-Python twin); parity is validated by
7//! convergence rather than against a reference distribution.
8//!
9//! Note: the C++ free-function driver left `popX` unpopulated, so its reported
10//! best-x read uninitialized memory. This port decodes the population every
11//! generation (as the ask/tell path did), so the best-x is always a real point.
12
13use nalgebra::DVector;
14
15use crate::fitness::Fitness;
16use crate::rng::Rng;
17
18/// Outcome of a PGPE run (mirrors the C++ `PgpeResult`).
19#[derive(Clone, Debug)]
20pub struct PgpeResult {
21    pub x: Vec<f64>,
22    pub y: f64,
23    pub evaluations: u64,
24    pub iterations: i32,
25    pub stop: i32,
26}
27
28/// Tunable inputs for [`Pgpe::new`].
29#[derive(Clone, Debug)]
30pub struct PgpeParams {
31    pub popsize: i32,
32    pub max_evaluations: u64,
33    pub stop_fitness: f64,
34    pub lr_decay_steps: i32,
35    pub use_ranking: bool,
36    pub center_learning_rate: f64,
37    pub stdev_learning_rate: f64,
38    pub stdev_max_change: f64,
39    pub b1: f64,
40    pub b2: f64,
41    pub eps: f64,
42    pub decay_coef: f64,
43    pub seed: u64,
44    pub runid: i64,
45}
46
47impl Default for PgpeParams {
48    fn default() -> Self {
49        Self {
50            popsize: 32,
51            max_evaluations: 100_000,
52            stop_fitness: f64::NEG_INFINITY,
53            lr_decay_steps: 1000,
54            use_ranking: true,
55            center_learning_rate: 0.15,
56            stdev_learning_rate: 0.1,
57            stdev_max_change: 0.2,
58            b1: 0.9,
59            b2: 0.999,
60            eps: 1e-8,
61            decay_coef: 1.0,
62            seed: 0,
63            runid: 0,
64        }
65    }
66}
67
68fn sort_index(v: &[f64]) -> Vec<usize> {
69    let mut idx: Vec<usize> = (0..v.len()).collect();
70    idx.sort_by(|&a, &b| v[a].partial_cmp(&v[b]).unwrap_or(std::cmp::Ordering::Equal));
71    idx
72}
73
74/// ADAM optimizer for the distribution center (the C++ `ADAM`).
75struct Adam {
76    x: DVector<f64>,
77    m: DVector<f64>,
78    v: DVector<f64>,
79    b1: f64,
80    b2: f64,
81    eps: f64,
82    center_lr: f64,
83    decay_coef: f64,
84}
85
86impl Adam {
87    fn new(x0: &DVector<f64>, b1: f64, b2: f64, eps: f64, center_lr: f64, decay_coef: f64) -> Self {
88        let dim = x0.len();
89        Adam {
90            x: x0.clone(),
91            m: DVector::zeros(dim),
92            v: DVector::zeros(dim),
93            b1,
94            b2,
95            eps,
96            center_lr,
97            decay_coef,
98        }
99    }
100
101    fn step_size(&self, i: i32) -> f64 {
102        self.center_lr * self.decay_coef.powi(i)
103    }
104
105    fn update(&mut self, i: i32, g: &DVector<f64>) {
106        self.m = g * (1.0 - self.b1) + &self.m * self.b1;
107        self.v = g.map(|v| v * v) * (1.0 - self.b2) + &self.v * self.b2;
108        let bc1 = 1.0 / (1.0 - self.b1.powi(i + 1));
109        let bc2 = 1.0 / (1.0 - self.b2.powi(i + 1));
110        let mhat = &self.m * bc1;
111        let vhat = &self.v * bc2;
112        let step = self.step_size(i);
113        let delta = DVector::from_iterator(
114            self.x.len(),
115            (0..self.x.len()).map(|k| step * mhat[k] / (vhat[k].sqrt() + self.eps)),
116        );
117        self.x -= delta;
118    }
119}
120
121pub struct Pgpe {
122    fitfun: Fitness,
123    rng: Rng,
124    dim: usize,
125    popsize: usize,
126    max_evaluations: u64,
127    stopfitness: f64,
128    lr_decay_steps: i32,
129    use_ranking: bool,
130    stdev_learning_rate: f64,
131    stdev_max_change: f64,
132
133    adam: Adam,
134    center: DVector<f64>,
135    stdev: DVector<f64>,
136    scaled_noises: Vec<DVector<f64>>, // n columns
137    pop_x: Vec<DVector<f64>>,         // decoded population (popsize)
138
139    best_x: DVector<f64>,
140    best_y: f64,
141    iterations: i32,
142    stop: i32,
143    external_evaluations: u64,
144}
145
146impl Pgpe {
147    pub fn new(mut fitfun: Fitness, guess: &[f64], input_sigma: &[f64], p: &PgpeParams) -> Self {
148        let dim = fitfun.dim();
149        fitfun.reset_evaluations();
150        let mut popsize = if p.popsize > 0 {
151            p.popsize as usize
152        } else {
153            4 * dim
154        };
155        if popsize % 2 == 1 {
156            popsize += 1;
157        }
158        let center = DVector::from_vec(fitfun.encode(guess));
159        let stdev = if input_sigma.len() == 1 {
160            DVector::from_element(dim, input_sigma[0])
161        } else {
162            DVector::from_row_slice(input_sigma)
163        };
164        // ADAM optimizes the center in encoded space. (The C++ seeded ADAM with
165        // the raw guess while `center` was encoded, so after the first tell the
166        // center jumped to raw coordinates in normalized space — a latent
167        // inconsistency; seeding ADAM with the encoded center fixes it.)
168        let adam = Adam::new(
169            &center,
170            p.b1,
171            p.b2,
172            p.eps,
173            p.center_learning_rate,
174            p.decay_coef,
175        );
176        Pgpe {
177            dim,
178            popsize,
179            max_evaluations: if p.max_evaluations > 0 {
180                p.max_evaluations
181            } else {
182                50_000
183            },
184            stopfitness: p.stop_fitness,
185            lr_decay_steps: p.lr_decay_steps.max(1),
186            use_ranking: p.use_ranking,
187            stdev_learning_rate: p.stdev_learning_rate.abs(),
188            stdev_max_change: p.stdev_max_change.abs(),
189            adam,
190            center,
191            stdev,
192            scaled_noises: vec![],
193            pop_x: vec![DVector::zeros(dim); popsize],
194            best_x: DVector::zeros(dim),
195            best_y: f64::MAX,
196            iterations: 0,
197            stop: 0,
198            external_evaluations: 0,
199            rng: Rng::new(p.seed.wrapping_add(p.runid as u64)),
200            fitfun,
201        }
202    }
203
204    pub fn dim(&self) -> usize {
205        self.dim
206    }
207    pub fn popsize(&self) -> usize {
208        self.popsize
209    }
210    pub fn stop(&self) -> i32 {
211        self.stop
212    }
213
214    /// Symmetric sampling: returns `popsize` *encoded* candidates, interleaved
215    /// `[center+n0, center-n0, center+n1, center-n1, ...]`, storing the noises.
216    fn ask_encoded(&mut self) -> Vec<DVector<f64>> {
217        let n = self.popsize / 2;
218        self.scaled_noises = (0..n)
219            .map(|_| {
220                let noise =
221                    DVector::from_iterator(self.dim, (0..self.dim).map(|_| self.rng.gaussian()));
222                noise.component_mul(&self.stdev)
223            })
224            .collect();
225        let mut xs = Vec::with_capacity(self.popsize);
226        for p in 0..n {
227            xs.push(&self.center + &self.scaled_noises[p]);
228            xs.push(&self.center - &self.scaled_noises[p]);
229        }
230        xs
231    }
232
233    /// Decoded, in-bounds population (rows), stored for the reinforce update.
234    fn ask_pop_internal(&mut self) -> Vec<Vec<f64>> {
235        let xs = self.ask_encoded();
236        self.pop_x = xs
237            .iter()
238            .map(|c| {
239                let feasible = self.fitfun.closest_feasible_normed(c.as_slice());
240                DVector::from_vec(self.fitfun.decode(&feasible))
241            })
242            .collect();
243        self.pop_x.iter().map(|c| c.as_slice().to_vec()).collect()
244    }
245
246    fn process_scores(&self, ys: &[f64]) -> DVector<f64> {
247        if self.use_ranking {
248            let n = ys.len();
249            let order = sort_index(ys);
250            let mut ranks = DVector::zeros(n);
251            for (i, &idx) in order.iter().enumerate() {
252                ranks[idx] = i as f64 / n as f64 - 0.5;
253            }
254            ranks
255        } else {
256            DVector::from_row_slice(ys)
257        }
258    }
259
260    /// grad_center, grad_stdev from the REINFORCE estimator (the C++
261    /// `compute_reinforce_update`).
262    fn reinforce(&self, pop_y: &DVector<f64>) -> (DVector<f64>, DVector<f64>) {
263        let n = self.popsize / 2;
264        let mean_all = pop_y.mean();
265        let mut grad_center = DVector::zeros(self.dim);
266        let mut grad_stdev = DVector::zeros(self.dim);
267        for i in 0..self.dim {
268            let mut gc = 0.0;
269            let mut gs = 0.0;
270            for p in 0..n {
271                let fit1 = pop_y[2 * p];
272                let fit2 = pop_y[2 * p + 1];
273                let score = fit1 - fit2;
274                let avg = 0.5 * (fit1 + fit2);
275                let sn = self.scaled_noises[p][i];
276                gc += sn * score * 0.5;
277                gs += (avg - mean_all) * (sn * sn - self.stdev[i] * self.stdev[i]) / self.stdev[i];
278            }
279            grad_center[i] = gc / n as f64;
280            grad_stdev[i] = gs / n as f64;
281        }
282        (grad_center, grad_stdev)
283    }
284
285    fn update_stdev(&self, grad: &DVector<f64>) -> DVector<f64> {
286        DVector::from_iterator(
287            self.dim,
288            (0..self.dim).map(|i| {
289                let allowed = self.stdev[i].abs() * self.stdev_max_change;
290                let lo = self.stdev[i] - allowed;
291                let hi = self.stdev[i] + allowed;
292                (self.stdev[i] + self.stdev_learning_rate * grad[i]).clamp(lo, hi)
293            }),
294        )
295    }
296
297    fn tell(&mut self, ys: &[f64]) -> i32 {
298        let neg: Vec<f64> = ys.iter().map(|y| -y).collect();
299        let pop_y = self.process_scores(&neg);
300        // Track the *true* best fitness/point. (The C++ tracked `-max(process_
301        // scores(-ys))`, which under ranking is a rank value, not a fitness, so
302        // its reported best was unusable; using the raw ys is correct and keeps
303        // the result meaningful for retry/comparison.)
304        let mut best_p = 0;
305        for p in 1..self.popsize {
306            if ys[p] < ys[best_p] {
307                best_p = p;
308            }
309        }
310        if ys[best_p] < self.best_y {
311            self.best_y = ys[best_p];
312            self.best_x = self.pop_x[best_p].clone();
313            if self.best_y < self.stopfitness {
314                self.stop = 1;
315            }
316        }
317        let (grad_center, grad_stdev) = self.reinforce(&pop_y);
318        self.adam
319            .update(self.iterations / self.lr_decay_steps, &(-grad_center));
320        self.iterations += 1;
321        self.center = self.adam.x.clone();
322        self.stdev = self.update_stdev(&grad_stdev);
323        self.stop
324    }
325
326    fn make_result(&self, evaluations: u64) -> PgpeResult {
327        PgpeResult {
328            x: self.best_x.as_slice().to_vec(),
329            y: self.best_y,
330            evaluations,
331            iterations: self.iterations,
332            stop: self.stop,
333        }
334    }
335
336    /// Generational loop evaluating each population through a batch closure.
337    pub fn optimize_batch<F>(&mut self, mut eval_batch: F) -> PgpeResult
338    where
339        F: FnMut(&[Vec<f64>]) -> Vec<f64>,
340    {
341        self.iterations = 0;
342        self.fitfun.reset_evaluations();
343        while self.fitfun.evaluations() < self.max_evaluations
344            && !self.fitfun.terminate()
345            && self.stop == 0
346        {
347            let rows = self.ask_pop_internal();
348            let mut ys = eval_batch(&rows);
349            for v in ys.iter_mut() {
350                if !v.is_finite() {
351                    *v = crate::fitness::NAN_REPLACEMENT;
352                }
353            }
354            self.fitfun.incr_evaluations(self.popsize as u64);
355            self.tell(&ys);
356        }
357        self.make_result(self.fitfun.evaluations())
358    }
359
360    // ---- ask/tell interface (mirrors PgpeState::Impl) ----
361
362    pub fn ask_pop(&mut self) -> Vec<Vec<f64>> {
363        self.ask_pop_internal()
364    }
365
366    pub fn tell_pop(&mut self, ys: &[f64]) -> i32 {
367        let stop = self.tell(ys);
368        self.external_evaluations += ys.len() as u64;
369        stop
370    }
371
372    pub fn population(&self) -> Vec<Vec<f64>> {
373        self.pop_x.iter().map(|c| c.as_slice().to_vec()).collect()
374    }
375
376    pub fn result(&self) -> PgpeResult {
377        self.make_result(self.external_evaluations)
378    }
379}
380
381#[cfg(test)]
382mod tests {
383    use super::*;
384
385    fn sphere(x: &[f64]) -> f64 {
386        x.iter().map(|v| v * v).sum()
387    }
388
389    #[test]
390    fn converges_on_sphere() {
391        // use_ranking=false so best_y is the true fitness.
392        let mut fit = Fitness::bounded(5, 1, &[-5.0; 5], &[5.0; 5]);
393        fit.set_normalize(true);
394        let params = PgpeParams {
395            popsize: 40,
396            max_evaluations: 20000,
397            use_ranking: false,
398            seed: 1,
399            ..Default::default()
400        };
401        let mut opt = Pgpe::new(fit, &[3.0; 5], &[0.5; 5], &params);
402        let r = opt.optimize_batch(|rows| rows.iter().map(|x| sphere(x)).collect());
403        assert!(r.y < 1e-2, "pgpe did not converge: {}", r.y);
404    }
405
406    #[test]
407    fn ask_tell_best_x_is_real_point() {
408        let mut fit = Fitness::bounded(4, 1, &[-5.0; 4], &[5.0; 4]);
409        fit.set_normalize(true);
410        let params = PgpeParams {
411            popsize: 32,
412            use_ranking: false,
413            seed: 2,
414            ..Default::default()
415        };
416        let mut opt = Pgpe::new(fit, &[2.0; 4], &[0.5; 4], &params);
417        for _ in 0..400 {
418            let pop = opt.ask_pop();
419            let ys: Vec<f64> = pop.iter().map(|x| sphere(x)).collect();
420            opt.tell_pop(&ys);
421        }
422        let r = opt.result();
423        // best_x must evaluate close to the reported best value.
424        assert!((sphere(&r.x) - r.y).abs() < 1e-6 || r.y < 1e-2);
425        assert!(sphere(&r.x) < 1e-1, "best_x not good: {}", sphere(&r.x));
426    }
427
428    #[test]
429    fn ranking_defaults_odd_population_getters_and_nonfinite_scores() {
430        let mut fit = Fitness::bounded(2, 1, &[-1.0; 2], &[1.0; 2]);
431        fit.set_normalize(true);
432        let params = PgpeParams {
433            popsize: 3,
434            max_evaluations: 0,
435            stop_fitness: 1.0e100,
436            lr_decay_steps: 0,
437            use_ranking: true,
438            stdev_learning_rate: -0.1,
439            stdev_max_change: -0.2,
440            seed: 14,
441            ..Default::default()
442        };
443        let mut optimizer = Pgpe::new(fit, &[0.0; 2], &[0.25], &params);
444        assert_eq!(optimizer.dim(), 2);
445        assert_eq!(optimizer.popsize(), 4);
446        assert_eq!(optimizer.stop(), 0);
447        let result = optimizer.optimize_batch(|rows| vec![f64::NAN; rows.len()]);
448        assert_eq!(result.evaluations, 4);
449        assert_eq!(result.stop, 1);
450        assert_eq!(optimizer.population().len(), 4);
451        assert_eq!(optimizer.stop(), 1);
452    }
453}