#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Changeover {
Control,
Coordinate,
Fixture,
}
#[derive(Debug, Clone, PartialEq)]
pub enum AxisValue {
Num(f64),
Label(String),
Bool(bool),
}
impl AxisValue {
pub fn as_num(&self) -> f64 {
match self {
AxisValue::Num(f) => *f,
AxisValue::Bool(b) => {
if *b {
1.0
} else {
0.0
}
}
AxisValue::Label(_) => 0.0,
}
}
}
impl std::fmt::Display for AxisValue {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
AxisValue::Num(n) => write!(f, "{n}"),
AxisValue::Label(s) => write!(f, "{s}"),
AxisValue::Bool(b) => write!(f, "{b}"),
}
}
}
impl From<f64> for AxisValue {
fn from(v: f64) -> Self {
AxisValue::Num(v)
}
}
pub type Coord = Vec<AxisValue>;
#[derive(Debug, Clone)]
pub enum AxisKind {
Continuous { lo: f64, hi: f64, min_step: f64 },
Discrete { detents: Vec<AxisValue> },
Categorical { options: Vec<AxisValue> },
}
#[derive(Debug, Clone)]
pub struct Axis {
pub name: String,
pub kind: AxisKind,
pub changeover: Changeover,
}
impl Axis {
pub fn center(&self) -> AxisValue {
match &self.kind {
AxisKind::Continuous { lo, hi, .. } => AxisValue::Num(0.5 * (lo + hi)),
AxisKind::Discrete { detents } => {
let mut d = detents.clone();
d.sort_by(|a, b| {
a.as_num()
.partial_cmp(&b.as_num())
.unwrap_or(std::cmp::Ordering::Equal)
});
d.get(d.len() / 2).cloned().unwrap_or(AxisValue::Num(0.0))
}
AxisKind::Categorical { options } => {
options.first().cloned().unwrap_or(AxisValue::Num(0.0))
}
}
}
}
#[derive(Debug, Clone)]
pub struct SearchSpace {
pub axes: Vec<Axis>,
}
impl SearchSpace {
pub fn new(axes: Vec<Axis>) -> Self {
Self { axes }
}
pub fn dims(&self) -> usize {
self.axes.len()
}
pub fn center(&self) -> Coord {
self.axes.iter().map(Axis::center).collect()
}
}
#[derive(Debug, Clone)]
pub struct Observation {
pub value: f64,
pub feasible: bool,
pub cost: f64,
pub metrics: Vec<(String, f64)>,
}
impl Observation {
pub fn value(v: f64) -> Self {
Self {
value: v,
feasible: true,
cost: 0.0,
metrics: Vec::new(),
}
}
pub fn infeasible() -> Self {
Self {
value: f64::NEG_INFINITY,
feasible: false,
cost: 0.0,
metrics: Vec::new(),
}
}
}
pub trait Objective {
fn query(&mut self, x: &[AxisValue]) -> Observation;
fn query_fidelity(&mut self, x: &[AxisValue], _fidelity: f64) -> Observation {
self.query(x)
}
}
#[derive(Debug, Clone)]
pub struct Budget {
pub max_evals: usize,
pub max_seconds: Option<f64>,
pub seed: u64,
}
impl Budget {
pub fn seeded(max_evals: usize, seed: u64) -> Self {
Self {
max_evals,
max_seconds: None,
seed,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StopReason {
Converged,
BudgetExhausted,
NoFeasiblePoint,
Aborted,
}
#[derive(Debug, Clone)]
pub struct AxisImpact {
pub name: String,
pub main_effect: f64,
pub curvature: f64,
}
#[derive(Debug, Clone)]
pub struct Report {
pub best: Coord,
pub best_value: f64,
pub evals: usize,
pub stop: StopReason,
pub ranked_axes: Vec<AxisImpact>,
pub history: Vec<(Coord, f64)>,
}
#[derive(Debug, Clone, Default)]
pub struct OptimizerParams {
pub overrides: Vec<(String, f64)>,
}
impl OptimizerParams {
pub fn new() -> Self {
Self::default()
}
pub fn with(mut self, key: impl Into<String>, value: f64) -> Self {
self.overrides.push((key.into(), value));
self
}
pub fn get(&self, key: &str, default: f64) -> f64 {
self.overrides
.iter()
.rev()
.find(|(k, _)| k == key)
.map(|(_, v)| *v)
.unwrap_or(default)
}
}
pub trait CoordinateSource: Send {
fn as_feedback(&mut self) -> Option<&mut dyn FeedbackSource> {
None
}
fn as_pull(&mut self) -> Option<&mut dyn PullSource> {
None
}
}
pub trait PullSource: Send {
fn pull(&mut self) -> Option<Vec<Coord>>;
}
pub trait FeedbackSource {
fn step(&mut self, evaluated: &[(Coord, f64)]) -> Option<Vec<Coord>>;
}
pub struct LexSource {
lists: Vec<Vec<AxisValue>>,
idx: Vec<usize>,
done: bool,
}
impl LexSource {
pub fn new(space: &SearchSpace) -> Self {
let lists: Vec<Vec<AxisValue>> = space
.axes
.iter()
.map(|a| match &a.kind {
AxisKind::Discrete { detents } => detents.clone(),
AxisKind::Categorical { options } => options.clone(),
AxisKind::Continuous { lo, hi, .. } => {
vec![AxisValue::Num(*lo), AxisValue::Num(*hi)]
}
})
.collect();
let n = lists.len();
Self {
lists,
idx: vec![0; n],
done: false,
}
}
}
impl PullSource for LexSource {
fn pull(&mut self) -> Option<Vec<Coord>> {
if self.done {
return None;
}
let n = self.lists.len();
let point: Vec<AxisValue> = (0..n).map(|d| self.lists[d][self.idx[d]].clone()).collect();
let mut d = n;
loop {
if d == 0 {
self.done = true;
break;
}
d -= 1;
self.idx[d] += 1;
if self.idx[d] < self.lists[d].len() {
break;
}
self.idx[d] = 0;
}
Some(vec![point])
}
}
pub struct PullOnly(pub Box<dyn PullSource>);
impl CoordinateSource for PullOnly {
fn as_pull(&mut self) -> Option<&mut dyn PullSource> {
Some(self.0.as_mut())
}
}
pub trait Optimizer: Send {
fn name(&self) -> &str;
fn doc_md(&self) -> &str;
fn coordinate_source(
&self,
space: &SearchSpace,
budget: &Budget,
lex: Box<dyn PullSource>,
) -> Box<dyn CoordinateSource>;
fn optimize(&self, space: &SearchSpace, obj: &mut dyn Objective, budget: &Budget) -> Report {
drive_source(self, space, obj, budget)
}
}
fn drive_source(
opt: &(impl Optimizer + ?Sized),
space: &SearchSpace,
obj: &mut dyn Objective,
budget: &Budget,
) -> Report {
let lex: Box<dyn PullSource> = Box::new(LexSource::new(space));
let mut src = opt.coordinate_source(space, budget, lex);
let mut best = space.center();
let mut best_value = f64::NEG_INFINITY;
let mut history: Vec<(Coord, f64)> = Vec::new();
let mut any_feasible = false;
let mut evals = 0usize;
let mut budget_hit = false;
fn next(src: &mut Box<dyn CoordinateSource>, evaluated: &[(Coord, f64)]) -> Option<Vec<Coord>> {
if let Some(f) = src.as_feedback() {
f.step(evaluated)
} else if let Some(p) = src.as_pull() {
p.pull()
} else {
None
}
}
let mut batch = next(&mut src, &[]);
while let Some(coords) = batch.take() {
let mut evaluated: Vec<(Coord, f64)> = Vec::with_capacity(coords.len());
for c in coords {
if evals >= budget.max_evals {
budget_hit = true;
break;
}
let obs = obj.query(&c);
evals += 1;
let v = if obs.feasible && obs.value.is_finite() {
any_feasible = true;
obs.value
} else {
-1.0e18
};
if v > best_value {
best_value = v;
best = c.clone();
}
history.push((c.clone(), v));
evaluated.push((c, v));
}
if budget_hit {
break;
}
batch = next(&mut src, &evaluated);
}
let stop = if !any_feasible {
StopReason::NoFeasiblePoint
} else if budget_hit {
StopReason::BudgetExhausted
} else {
StopReason::Converged
};
Report {
best,
best_value,
evals,
stop,
ranked_axes: Vec::new(),
history,
}
}
pub struct SweepOptimizer;
const SWEEP_DOC: &str = "# sweep\n\nThe identity optimizer — the **default** when no `method:` is \
set. Returns the default lexicographic coordinate stream unchanged — the full Cartesian product of \
the discrete-axis detents (and the `{lo, hi}` corners of continuous axes), in lex order. So `sweep` \
evaluates EVERY coordinate exhaustively and reports the best by the objective; it reproduces the \
engine's ordinary parameter sweep, now with best-selection. (Use an adaptive method — `nelder_mead`, \
`cmaes`, … — to search a continuous space without enumerating it.)\n";
impl Optimizer for SweepOptimizer {
fn name(&self) -> &str {
"sweep"
}
fn doc_md(&self) -> &str {
SWEEP_DOC
}
fn coordinate_source(
&self,
_space: &SearchSpace,
_budget: &Budget,
lex: Box<dyn PullSource>,
) -> Box<dyn CoordinateSource> {
Box::new(PullOnly(lex)) }
}
pub struct OptimizerRegistration {
pub name: fn() -> &'static str,
pub doc_md: fn() -> &'static str,
pub make: fn(&OptimizerParams) -> Box<dyn Optimizer>,
}
inventory::collect!(OptimizerRegistration);
#[derive(Debug, Clone)]
pub struct OptimizerInfo {
pub name: &'static str,
pub doc_md: &'static str,
}
pub fn by_name(name: &str, params: &OptimizerParams) -> Option<Box<dyn Optimizer>> {
if name == "sweep" {
return Some(Box::new(SweepOptimizer));
}
inventory::iter::<OptimizerRegistration>
.into_iter()
.find(|r| (r.name)() == name)
.map(|r| (r.make)(params))
}
pub fn describe() -> Vec<OptimizerInfo> {
let mut out = vec![OptimizerInfo {
name: "sweep",
doc_md: SWEEP_DOC,
}];
for r in inventory::iter::<OptimizerRegistration> {
out.push(OptimizerInfo {
name: (r.name)(),
doc_md: (r.doc_md)(),
});
}
out.sort_by(|a, b| a.name.cmp(b.name));
out.dedup_by(|a, b| a.name == b.name);
out
}
pub fn registered_names() -> Vec<&'static str> {
describe().into_iter().map(|i| i.name).collect()
}
#[cfg(test)]
mod tests {
use super::*;
struct Paraboloid {
target: Vec<f64>,
}
impl Objective for Paraboloid {
fn query(&mut self, x: &[AxisValue]) -> Observation {
let v: f64 = x
.iter()
.map(AxisValue::as_num)
.zip(&self.target)
.map(|(a, t)| -(a - t) * (a - t))
.sum();
Observation::value(v)
}
}
#[test]
fn sweep_is_always_available_and_sweeps_the_grid() {
let space = SearchSpace::new(vec![
Axis {
name: "x".into(),
kind: AxisKind::Discrete {
detents: vec![
AxisValue::Num(0.0),
AxisValue::Num(1.0),
AxisValue::Num(2.0),
],
},
changeover: Changeover::Coordinate,
},
Axis {
name: "y".into(),
kind: AxisKind::Discrete {
detents: vec![
AxisValue::Num(0.0),
AxisValue::Num(1.0),
AxisValue::Num(2.0),
],
},
changeover: Changeover::Coordinate,
},
]);
let opt = by_name("sweep", &OptimizerParams::new()).expect("sweep is built in");
let mut obj = Paraboloid {
target: vec![1.0, 2.0],
};
let r = opt.optimize(&space, &mut obj, &Budget::seeded(100, 0));
assert_eq!(r.evals, 9);
assert_eq!(r.best, vec![AxisValue::Num(1.0), AxisValue::Num(2.0)]);
assert!(r.best_value.abs() < 1e-9);
assert_eq!(r.stop, StopReason::Converged);
}
#[test]
fn describe_includes_sweep_and_unknown_is_none() {
let infos = describe();
assert!(infos.iter().any(|i| i.name == "sweep"));
assert!(infos.iter().any(|i| i.doc_md.contains("# sweep")));
assert!(by_name("definitely_not_registered", &OptimizerParams::new()).is_none());
}
}