Skip to content

Commit 9c48802

Browse files
committed
feat!: move CMA-ES distribution into CmaEsState (B8)
CMA-ES was the one solver parking its iterate — the search distribution (m, σ, C, B/D, evolution paths) — on the solver struct, which forced its canonical TolX test to be a hardcoded `terminate` hook that ignored its state argument (routing around tenet 3) and contradicted the crate's own articulated principle that working state lives in the state, not the solver (see LbfgsState). Introduce `CmaEsState<V, M, F>` (core/state/cma_es.rs): one shared state for both CmaEs and BoundedCmaEs, holding the distribution plus the population, with the bounded variant's adaptive-penalty bookkeeping riding along as `Option<BoundPenalty>` (mirroring LbfgsState::work). No bespoke trait — the new `CmaEsTolerance` criterion binds the concrete state and fires `TerminationReason::CmaEsTolerance`. Both solvers become configuration-only: `CmaEs::new(seed)` / `BoundedCmaEs::new(seed)`, with mean/σ/stds on `CmaEsState::new(mean, σ).with_stds(...)` and the old `with_tol_x` replaced by the criterion. Derived constants + RNG are cached on the solver. Result semantics follow canonical CMA-ES: `param()`/`cost()` return the distribution mean (xfavorite; the solver evaluates f(m) once per generation), while `best_param()`/`best_cost()` return the best evaluated sample (xbest) — an OptimizationResult surfaces both. The PopulationState contract is relaxed so `param()` need not equal candidates[0]; MaLsChCma's chain write-back reads best_param(). Updates the injection/chain consumers, lib re-exports, wasm bindings, benches, the CMA test suite across all backends, and docs (AGENTS.md observer-KV note, audit B8 marked resolved, solver-composition rule). Verified: cargo test --workspace --all-features (699 pass), clippy (all targets/features), doc build, wasm target, web build. BREAKING CHANGE: CMA-ES public API changed. `CmaEs::new` / `BoundedCmaEs::new` now take only `(seed)`; the initial mean and σ move to `CmaEsState::new(mean, sigma)`, and the Executor state for both solvers is now `CmaEsState<V, M>` instead of `BasicPopulationState<V>`. `CmaEs::with_stds` / `with_tol_x` are removed — use `CmaEsState::with_stds(...)` and the `CmaEsTolerance` termination criterion. `OptimizationResult::param()` / `cost()` now return the distribution mean (xfavorite), not the best sample; use `best_param()` / `best_cost()` for the best evaluated point. Standalone CMA-ES no longer self-terminates on TolX unless a `CmaEsTolerance` criterion is registered.
1 parent 0c01ff7 commit 9c48802

31 files changed

Lines changed: 1335 additions & 1384 deletions

‎.claude/rules/solver-composition.md‎

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -86,10 +86,10 @@ Each shipped state has its own mirror rule:
8686
`gradient_evals = gradient + jacobian + hessian` (the NLLS convention
8787
preserved — residual rolls into cost, Jacobian/Hessian into gradient).
8888
- Derivative-free states (`BasicSimplexState`, `BasicPopulationState`,
89-
`MaLsChState`): `cost_evals = total_work()` — every kind of work folded
90-
in. This is what makes a CMA-ES outer with an L-BFGS inner just *work*:
91-
the inner's gradient evals show up in the outer's `cost_evals` honestly,
92-
with no manual cross-type fold.
89+
`CmaEsState`, `MaLsChState`): `cost_evals = total_work()` — every kind of
90+
work folded in. This is what makes a CMA-ES outer (which drives a
91+
`CmaEsState`) with an L-BFGS inner just *work*: the inner's gradient evals
92+
show up in the outer's `cost_evals` honestly, with no manual cross-type fold.
9393

9494
User-defined state types plugging into `Executor` must impl `CountsMirror`;
9595
it is `pub` for exactly that reason.
@@ -107,8 +107,8 @@ inner's state, not just drive it — and inners carry different state shapes
107107
`seed_scaled(x, σ)` (defaults to `seed`; only Nelder-Mead's σ-scaled simplex
108108
overrides it). No per-trait eval-aggregation hook — same-problem composition
109109
shares the `Problem<P>` wrapper, and the `total_work()` fold in
110-
`BasicPopulationState`'s `CountsMirror` rolls every kind of inner work into
111-
the outer's single `cost_evals` automatically.
110+
`CmaEsState`'s `CountsMirror` (same rule as `BasicPopulationState`) rolls
111+
every kind of inner work into the outer's single `cost_evals` automatically.
112112

113113
Two consumer families validate the split: the barrier / AL methods bound `So:
114114
WarmStart<V>` with `So::State: GradientState + CountsMirror` (gradient inners

‎AGENTS.md‎

Lines changed: 7 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -50,11 +50,13 @@ the relevant files. Don't duplicate it here:
5050
to it, and existing `Observe` impls keep compiling via the forwarding default
5151
— a concrete `Kv` type keeps the trait object-safe for `Box<dyn Observe>`. So
5252
this is *deferred*, not foreclosed. The genuine future motivation, if it
53-
comes: solver-internal working state (CMA-ES σ / covariance / evolution
54-
paths, LM μ / ν / diag) lives in the *solver* struct, not the state, so the
55-
"expose it on a richer state trait" answer does not cover those scalars.
56-
Don't "fix" the absence of a KV channel by reflex — but don't treat it as
57-
permanently closed either.
53+
comes: some solver-internal working state (LM μ / ν / diag) lives in the
54+
*solver* struct, not the state, so the "expose it on a richer state trait"
55+
answer does not cover those scalars. (CMA-ES is *no longer* an example: its
56+
σ / covariance / evolution paths moved onto `CmaEsState` so its TolX test
57+
could become the composable `CmaEsTolerance` criterion — see that state's
58+
rustdoc.) Don't "fix" the absence of a KV channel by reflex — but don't
59+
treat it as permanently closed either.
5860

5961
- **No `Solver::name()` introspection.** The `Solver` trait (`core/solver.rs`)
6062
has `type Error` + `init` / `next_iter` / `terminate` and deliberately no

‎crates/basin-wasm/src/lib.rs‎

Lines changed: 8 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -25,9 +25,9 @@ use basin::problems::{styblinski_tang, styblinski_tang_gradient};
2525
use basin::solver::lbfgs::{Lbfgs, Unbounded as LbfgsUnbounded};
2626
use basin::{
2727
Backtracking, BasicPopulationState, BasicSimplexState, BasicState, BoxConstraints, CmaEs,
28-
Constant, CostFunction, De, DenseMatrix, Executor, FiniteDiff, Gradient, GradientDescent,
29-
LbfgsState, MoreThuente, NelderMead, PopulationState, RandomSearch, Ssga, State, StepOutcome,
30-
Stepper, TerminationReason,
28+
CmaEsState, Constant, CostFunction, De, DenseMatrix, Executor, FiniteDiff, Gradient,
29+
GradientDescent, LbfgsState, MoreThuente, NelderMead, PopulationState, RandomSearch, Ssga,
30+
State, StepOutcome, Stepper, TerminationReason,
3131
};
3232
use serde::{Deserialize, Serialize};
3333
use wasm_bindgen::prelude::*;
@@ -298,7 +298,7 @@ type LbfgsStepper = Stepper<Problem2D, LbfgsState<Vec<f64>>, Lbfgs<LbfgsUnbounde
298298
/// Concrete population-solver stepper aliases. Same motivation as
299299
/// [`LbfgsStepper`] — keep the [`Inner`] variants readable.
300300
type CmaEsStepper =
301-
Stepper<Problem2D, BasicPopulationState<Vec<f64>>, CmaEs<Vec<f64>, DenseMatrix>>;
301+
Stepper<Problem2D, CmaEsState<Vec<f64>, DenseMatrix>, CmaEs<Vec<f64>, DenseMatrix>>;
302302
type DeStepper = Stepper<Problem2DBounded, BasicPopulationState<Vec<f64>>, De>;
303303
type RandomSearchStepper = Stepper<Problem2DBounded, BasicPopulationState<Vec<f64>>, RandomSearch>;
304304
type SsgaStepper = Stepper<Problem2DBounded, BasicPopulationState<Vec<f64>>, Ssga>;
@@ -602,20 +602,16 @@ impl Run {
602602
} else {
603603
0.25 * 0.5 * ((opts.xmax - opts.xmin) + (opts.ymax - opts.ymin))
604604
};
605-
let mut solver =
606-
CmaEs::<Vec<f64>, DenseMatrix>::new(initial.clone(), sigma, opts.seed);
605+
let mut solver = CmaEs::<Vec<f64>, DenseMatrix>::new(opts.seed);
607606
// λ < 4 is invalid for CMA-ES recombination weights; treat
608607
// small overrides as "auto" and let the solver pick.
609-
let lambda = if opts.cma_lambda >= 4 {
608+
if opts.cma_lambda >= 4 {
610609
solver = solver.with_lambda(opts.cma_lambda);
611-
opts.cma_lambda
612-
} else {
613-
CmaEs::<Vec<f64>, DenseMatrix>::default_lambda(2)
614-
};
610+
}
615611
let stepper = Executor::new(
616612
p,
617613
solver,
618-
BasicPopulationState::<Vec<f64>>::with_size(lambda),
614+
CmaEsState::<Vec<f64>, DenseMatrix>::new(initial.clone(), sigma),
619615
)
620616
.max_iter(max_iter as u64)
621617
.into_stepper()

‎crates/basin/benches/solver_backends.rs‎

Lines changed: 10 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -38,8 +38,8 @@ use std::time::Duration;
3838

3939
use basin::problems::{Ackley, Levy, Rastrigin, Rosenbrock, SparseLeastSquares, StyblinskiTang};
4040
use basin::{
41-
BasicPopulationState, BasicSimplexState, BasicState, Bfgs, CmaEs, DenseMatrix, Executor,
42-
GaussNewton, GradientDescent, LbfgsState, Lbfgsb, LevenbergMarquardt, MoreThuente, NelderMead,
41+
BasicSimplexState, BasicState, Bfgs, CmaEs, CmaEsState, DenseMatrix, Executor, GaussNewton,
42+
GradientDescent, LbfgsState, Lbfgsb, LevenbergMarquardt, MoreThuente, NelderMead,
4343
QuasiNewtonState,
4444
};
4545
use criterion::{BatchSize, BenchmarkId, Criterion, criterion_group, criterion_main};
@@ -280,17 +280,16 @@ fn bench_cmaes(c: &mut Criterion) {
280280
let mut g = c.benchmark_group(format!("cmaes_rastrigin_n{n}"));
281281
// In-domain start away from the global optimum at the origin.
282282
let m0 = vec![3.0; n];
283-
// λ is backend-independent; match the solver's internal default in the
284-
// population state so the contract holds.
285-
let lambda = CmaEs::<Vec<f64>, DenseMatrix>::default_lambda(n);
283+
// The mean / σ live on `CmaEsState`; the state builder seeds it from the
284+
// per-batch start vector and the solver derives λ internally.
286285

287286
contestant!(
288287
g,
289288
"vec",
290289
|| m0.clone(),
291290
Rastrigin<Vec<f64>>,
292-
CmaEs::<Vec<f64>, DenseMatrix>::new(m0.clone(), 0.3, 42),
293-
|_x0| BasicPopulationState::<Vec<f64>>::with_size(lambda),
291+
CmaEs::<Vec<f64>, DenseMatrix>::new(42),
292+
|x0| CmaEsState::<Vec<f64>, DenseMatrix>::new(x0, 0.3),
294293
);
295294

296295
let m0n = DVector::from_vec(m0.clone());
@@ -299,8 +298,8 @@ fn bench_cmaes(c: &mut Criterion) {
299298
"nalgebra",
300299
|| m0n.clone(),
301300
Rastrigin<DVector<f64>>,
302-
CmaEs::<DVector<f64>, DMatrix<f64>>::new(m0n.clone(), 0.3, 42),
303-
|_x0| BasicPopulationState::<DVector<f64>>::with_size(lambda),
301+
CmaEs::<DVector<f64>, DMatrix<f64>>::new(42),
302+
|x0| CmaEsState::<DVector<f64>, DMatrix<f64>>::new(x0, 0.3),
304303
);
305304

306305
let m0f = Col::<f64>::from_fn(n, |i| m0[i]);
@@ -309,8 +308,8 @@ fn bench_cmaes(c: &mut Criterion) {
309308
"faer",
310309
|| m0f.clone(),
311310
Rastrigin<Col<f64>>,
312-
CmaEs::<Col<f64>, Mat<f64>>::new(m0f.clone(), 0.3, 42),
313-
|_x0| BasicPopulationState::<Col<f64>>::with_size(lambda),
311+
CmaEs::<Col<f64>, Mat<f64>>::new(42),
312+
|x0| CmaEsState::<Col<f64>, Mat<f64>>::new(x0, 0.3),
314313
);
315314
g.finish();
316315
}

‎crates/basin/src/core/state.rs‎

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -18,9 +18,12 @@
1818
//! pipeline composes at `F = f32`, and the *Provisional choices* section of
1919
//! `CONTRIBUTING.md`.
2020
21+
/// CMA-ES distribution state (`CmaEsState`).
22+
pub mod cma_es;
2123
/// Limited-memory BFGS / L-BFGS-B state (`LbfgsState`).
2224
pub mod lbfgs;
2325

26+
pub use cma_es::CmaEsState;
2427
pub use lbfgs::LbfgsState;
2528

2629
use crate::core::math::{MatrixIdentity, Scalar, VectorLen};
@@ -298,14 +301,20 @@ pub trait SimplexState: State {
298301
/// [`costs`](Self::costs) sorted by **ascending cost** at the start
299302
/// and end of every
300303
/// [`Solver::next_iter`](crate::core::solver::Solver::next_iter)
301-
/// call (and at the end of [`Solver::init`](crate::core::solver::Solver::init)).
302-
/// So [`State::param`] / [`State::cost`] always return the current
303-
/// best candidate (`candidates[0]` / `costs[0]`).
304+
/// call (and at the end of [`Solver::init`](crate::core::solver::Solver::init)),
305+
/// so `candidates[0]` / `costs[0]` are always the best sampled
306+
/// candidate.
304307
/// - **Implementor must:** sort `NaN` costs *last*, so a single bad
305308
/// evaluation can't drag itself to the front and become the
306309
/// "best" candidate.
307310
/// - **Implementor must:** keep the two slices the same length and in
308311
/// parallel order — `costs[i]` is the cost at `candidates[i]`.
312+
/// - What [`State::param`] / [`State::cost`] return is the [`State`]
313+
/// impl's responsibility and need *not* equal `candidates[0]`. Most
314+
/// population states (e.g. [`BasicPopulationState`]) return the best
315+
/// candidate; distribution-based states like
316+
/// [`CmaEsState`] return the distribution mean
317+
/// (`xfavorite`) while the population stays the sampled candidates.
309318
pub trait PopulationState: State {
310319
/// All `λ` candidates, sorted by ascending cost.
311320
fn candidates(&self) -> &[Self::Param];

0 commit comments

Comments
 (0)