use crate::{
core::{Callbacks, LinearAlgebra, MinimizationSummary, NalgebraProvider, RealScalar},
traits::{Algorithm, Status},
};
use serde::Serialize;
use std::{
convert::Infallible,
fmt::{Debug, Display},
};
use tabled::{
builder::Builder,
settings::{
object::Row, style::HorizontalLine, themes::BorderCorrection, Alignment, Padding, Span,
Style, Theme,
},
};
#[derive(Debug, Clone, Default)]
pub struct MultiStartState<T: RealScalar = f64, B: LinearAlgebra<T> = NalgebraProvider> {
runs: Vec<MinimizationSummary<T, B>>,
}
impl<T: RealScalar, B: LinearAlgebra<T>> MultiStartState<T, B> {
pub const fn new() -> Self {
Self { runs: Vec::new() }
}
pub fn runs(&self) -> &[MinimizationSummary<T, B>] {
&self.runs
}
pub fn completed_runs(&self) -> usize {
self.runs.len()
}
pub fn restart_count(&self) -> usize {
self.runs.len().saturating_sub(1)
}
pub fn best(&self) -> Option<&MinimizationSummary<T, B>> {
self.best_index().map(|index| &self.runs[index])
}
pub fn best_index(&self) -> Option<usize> {
self.runs
.iter()
.enumerate()
.min_by(|(_, a), (_, b)| a.fx.total_cmp(&b.fx))
.map(|(index, _)| index)
}
pub(crate) fn push(&mut self, summary: MinimizationSummary<T, B>) {
self.runs.push(summary);
}
}
#[derive(Clone, Serialize)]
#[serde(bound(
serialize = "T: Serialize, B::VectorStorage: Serialize, B::MatrixStorage: Serialize"
))]
pub struct MultiStartSummary<T: RealScalar = f64, B: LinearAlgebra<T> = NalgebraProvider> {
pub runs: Vec<MinimizationSummary<T, B>>,
pub best_run_index: Option<usize>,
pub restart_count: usize,
}
impl<T: RealScalar, B: LinearAlgebra<T>> Display for MultiStartSummary<T, B> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let mut builder = Builder::default();
builder.push_record(["MULTISTART SUMMARY", "", "", "", ""]);
builder.push_record(["Completed runs", "Restarts", "Best run", "", ""]);
builder.push_record([
self.completed_runs().to_string(),
self.restart_count.to_string(),
self.best_run_index
.map_or_else(|| "—".to_string(), |index| index.to_string()),
String::new(),
String::new(),
]);
builder.push_record(["Run", "Best?", "Status", "f(x)", "# f(x)"]);
for (index, run) in self.runs.iter().enumerate() {
builder.push_record([
index.to_string(),
if self.best_run_index == Some(index) {
"Yes".to_string()
} else {
String::new()
},
if run.message.success() {
"Converged".to_string()
} else {
"Not converged".to_string()
},
format!("{:.5}", run.fx),
run.evals.f().to_string(),
]);
}
let mut table = builder.build();
let mut theme = Theme::from_style(Style::rounded().remove_horizontals());
for row in 1..=3 {
theme.insert_horizontal_line(row, HorizontalLine::inherit(Style::modern()));
}
table
.with(theme)
.modify(
Row::from(0),
(Alignment::center(), Padding::new(1, 1, 0, 0)),
)
.modify((0, 0), Span::column(5))
.modify((1, 2), Span::column(3))
.modify((2, 2), Span::column(3))
.with(BorderCorrection::span());
f.write_str(&table.to_string())
}
}
impl<T: RealScalar, B: LinearAlgebra<T>> Debug for MultiStartSummary<T, B> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"MultiStart Summary: completed_runs={}, restarts={}, best_run={:?}",
self.completed_runs(),
self.restart_count,
self.best_run_index
)
}
}
impl<T: RealScalar, B: LinearAlgebra<T>> MultiStartSummary<T, B> {
pub fn best(&self) -> Option<&MinimizationSummary<T, B>> {
self.best_run_index.map(|index| &self.runs[index])
}
pub fn completed_runs(&self) -> usize {
self.runs.len()
}
}
pub trait RestartPolicy<T: RealScalar = f64, B: LinearAlgebra<T> = NalgebraProvider>: Send {
fn should_run(&mut self, next_run_index: usize, state: &MultiStartState<T, B>) -> bool;
}
impl<T, B, F> RestartPolicy<T, B> for F
where
T: RealScalar,
B: LinearAlgebra<T>,
F: FnMut(usize, &MultiStartState<T, B>) -> bool + Send,
{
fn should_run(&mut self, next_run_index: usize, state: &MultiStartState<T, B>) -> bool {
self(next_run_index, state)
}
}
#[derive(Debug, Clone, Copy)]
pub struct FixedRestarts {
total_runs: usize,
}
impl FixedRestarts {
pub const fn new(total_runs: usize) -> Self {
Self { total_runs }
}
}
impl<T: RealScalar, B: LinearAlgebra<T>> RestartPolicy<T, B> for FixedRestarts {
fn should_run(&mut self, next_run_index: usize, _state: &MultiStartState<T, B>) -> bool {
next_run_index < self.total_runs
}
}
pub type RestartBundle<A, P, S, U, E> = (
A,
<A as Algorithm<P, S, U, E>>::Init,
<A as Algorithm<P, S, U, E>>::Config,
Callbacks<A, P, S, U, E, <A as Algorithm<P, S, U, E>>::Config>,
);
pub trait RestartFactory<A, P, S: Status, U = (), E = Infallible, T = f64, B = NalgebraProvider>:
Send
where
T: RealScalar,
B: LinearAlgebra<T>,
A: Algorithm<P, S, U, E, Summary = MinimizationSummary<T, B>>,
{
fn create(
&mut self,
run_index: usize,
state: &MultiStartState<T, B>,
) -> RestartBundle<A, P, S, U, E>;
}
impl<A, P, S, U, E, T, B, F> RestartFactory<A, P, S, U, E, T, B> for F
where
S: Status,
T: RealScalar,
B: LinearAlgebra<T>,
A: Algorithm<P, S, U, E, Summary = MinimizationSummary<T, B>>,
F: FnMut(usize, &MultiStartState<T, B>) -> RestartBundle<A, P, S, U, E> + Send,
{
fn create(
&mut self,
run_index: usize,
state: &MultiStartState<T, B>,
) -> RestartBundle<A, P, S, U, E> {
self(run_index, state)
}
}
pub const fn restart_seed(base_seed: u64, run_index: usize) -> u64 {
base_seed.wrapping_add(run_index as u64)
}
pub fn minimize_multistart<P, U, E, A, S, F, R, T, B>(
problem: &P,
user_data: &U,
restart_factory: &mut F,
restart_policy: &mut R,
) -> Result<MultiStartSummary<T, B>, E>
where
S: Status,
T: RealScalar,
B: LinearAlgebra<T>,
A: Algorithm<P, S, U, E, Summary = MinimizationSummary<T, B>>,
F: RestartFactory<A, P, S, U, E, T, B>,
R: RestartPolicy<T, B>,
{
let mut state = MultiStartState::new();
while restart_policy.should_run(state.completed_runs(), &state) {
let run_index = state.completed_runs();
let (mut algorithm, init, config, callbacks) = restart_factory.create(run_index, &state);
let summary = algorithm.process(problem, user_data, init, config, callbacks)?;
state.push(summary);
}
let best_run_index = state.best_index();
Ok(MultiStartSummary {
restart_count: state.restart_count(),
runs: state.runs,
best_run_index,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
core::{Callbacks, Matrix, MaxSteps, Vector},
traits::{Status, StatusMessage},
};
use serde::{Deserialize, Serialize};
#[derive(Clone, Default, Serialize, Deserialize)]
struct DummyStatus {
message: StatusMessage,
}
impl Status for DummyStatus {
fn reset(&mut self) {
self.message = StatusMessage::default();
}
fn message(&self) -> &StatusMessage {
&self.message
}
fn set_message(&mut self) -> &mut StatusMessage {
&mut self.message
}
}
#[derive(Clone, Default)]
struct DummyAlgorithm;
#[derive(Clone)]
struct DummyConfig {
x: Vector<f64>,
fx: f64,
}
impl Algorithm<(), DummyStatus, (), Infallible> for DummyAlgorithm {
type Summary = MinimizationSummary<f64>;
type Config = DummyConfig;
type Init = DummyConfig;
fn initialize(
&mut self,
_problem: &(),
status: &mut DummyStatus,
_args: &(),
_init: &Self::Init,
_config: &Self::Config,
) -> Result<(), Infallible> {
status.set_message().initialize();
Ok(())
}
fn step(
&mut self,
_current_step: usize,
_problem: &(),
status: &mut DummyStatus,
_args: &(),
_config: &Self::Config,
) -> Result<(), Infallible> {
status.set_message().succeed_with_message("done");
Ok(())
}
fn summarize(
&self,
_current_step: usize,
_problem: &(),
status: &DummyStatus,
_args: &(),
_init: &Self::Init,
config: &Self::Config,
) -> Result<Self::Summary, Infallible> {
Ok(MinimizationSummary {
bounds: None,
parameter_names: None,
message: status.message.clone(),
x0: config.x.clone(),
x: config.x.clone(),
std: crate::core::summary::unknown_uncertainties(config.x.len()),
fx: config.fx,
evals: crate::core::EvalCounts::new(1, 0, 0),
covariance: Matrix::identity(config.x.len()),
})
}
fn default_callbacks() -> Callbacks<Self, (), DummyStatus, (), Infallible, Self::Config> {
Callbacks::empty().with_terminator(MaxSteps(0))
}
}
#[test]
fn fixed_restarts_runs_expected_number_of_times_and_tracks_best() {
let mut factory = |run_index: usize, _state: &MultiStartState| {
(
DummyAlgorithm,
DummyConfig {
x: Vector::from_vec(vec![run_index as f64]),
fx: (3 - run_index) as f64,
},
DummyConfig {
x: Vector::from_vec(vec![run_index as f64]),
fx: (3 - run_index) as f64,
},
DummyAlgorithm::default_callbacks(),
)
};
let mut policy = FixedRestarts::new(3);
let summary = minimize_multistart::<
(),
(),
Infallible,
DummyAlgorithm,
DummyStatus,
_,
_,
f64,
NalgebraProvider,
>(&(), &(), &mut factory, &mut policy)
.unwrap();
assert_eq!(summary.completed_runs(), 3);
assert_eq!(summary.restart_count, 2);
assert_eq!(summary.best_run_index, Some(2));
assert_eq!(summary.best().unwrap().fx, 1.0);
let display = summary.to_string();
assert!(display.contains("MULTISTART SUMMARY"));
assert!(display.contains("Best?"));
assert!(display.contains("1.00000"));
assert_eq!(
format!("{summary:?}"),
"MultiStart Summary: completed_runs=3, restarts=2, best_run=Some(2)"
);
}
#[test]
fn closure_restart_policy_can_stop_based_on_seen_runs() {
let mut factory = |run_index: usize, _state: &MultiStartState| {
(
DummyAlgorithm,
DummyConfig {
x: Vector::from_vec(vec![run_index as f64]),
fx: run_index as f64,
},
DummyConfig {
x: Vector::from_vec(vec![run_index as f64]),
fx: run_index as f64,
},
DummyAlgorithm::default_callbacks(),
)
};
let mut policy = |_: usize, state: &MultiStartState| state.completed_runs() < 2;
let summary = minimize_multistart::<
(),
(),
Infallible,
DummyAlgorithm,
DummyStatus,
_,
_,
f64,
NalgebraProvider,
>(&(), &(), &mut factory, &mut policy)
.unwrap();
assert_eq!(summary.completed_runs(), 2);
assert_eq!(summary.restart_count, 1);
}
#[test]
fn restart_seed_is_deterministic() {
assert_eq!(restart_seed(7, 0), 7);
assert_eq!(restart_seed(7, 3), 10);
}
#[test]
fn zero_run_policy_returns_empty_summary() {
let mut factory = |run_index: usize, _state: &MultiStartState| {
(
DummyAlgorithm,
DummyConfig {
x: Vector::from_vec(vec![run_index as f64]),
fx: run_index as f64,
},
DummyConfig {
x: Vector::from_vec(vec![run_index as f64]),
fx: run_index as f64,
},
DummyAlgorithm::default_callbacks(),
)
};
let mut policy = FixedRestarts::new(0);
let summary = minimize_multistart::<
(),
(),
Infallible,
DummyAlgorithm,
DummyStatus,
_,
_,
f64,
NalgebraProvider,
>(&(), &(), &mut factory, &mut policy)
.unwrap();
assert_eq!(summary.completed_runs(), 0);
assert_eq!(summary.restart_count, 0);
assert_eq!(summary.best_run_index, None);
assert!(summary.best().is_none());
assert!(summary.to_string().contains("MULTISTART SUMMARY"));
}
}