use std::num::NonZeroUsize;
use serde::{Deserialize, Serialize};
use super::scenario_source::RawScenarioSourceConfig;
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct TrainingConfig {
#[serde(default = "TrainingConfig::default_enabled")]
pub enabled: bool,
#[serde(default)]
pub tree_seed: Option<i64>,
pub forward_passes: Option<u32>,
pub stopping_rules: Option<Vec<StoppingRuleConfig>>,
#[serde(default = "TrainingConfig::default_stopping_mode")]
pub stopping_mode: String,
#[serde(default)]
pub cut_selection: RowSelectionConfig,
#[serde(default)]
pub solver: TrainingSolverConfig,
#[serde(default)]
pub parallelism: ParallelismConfig,
#[serde(default)]
pub scenario_source: Option<RawScenarioSourceConfig>,
}
impl TrainingConfig {
pub(super) fn default_enabled() -> bool {
true
}
pub(super) fn default_stopping_mode() -> String {
"any".to_string()
}
}
#[derive(Debug, Clone, Deserialize, Serialize, Default)]
#[serde(default, deny_unknown_fields)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct RowSelectionConfig {
#[serde(default)]
pub row_activity_tolerance: Option<f64>,
#[serde(default)]
pub max_active_per_stage: Option<u32>,
#[serde(default)]
pub selection: Option<SelectionMethod>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(tag = "method", rename_all = "snake_case", deny_unknown_fields)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub enum SelectionMethod {
Level1 {
#[serde(default = "default_tie_tolerance")]
tie_tolerance: f64,
#[serde(default = "default_check_frequency")]
check_frequency: u32,
},
Lml1 {
#[serde(default = "default_tie_tolerance")]
tie_tolerance: f64,
#[serde(default = "default_check_frequency")]
check_frequency: u32,
},
Domination {
domination_tolerance: f64,
#[serde(default = "default_check_frequency")]
check_frequency: u32,
},
Dynamic {
#[serde(default = "default_start_iteration")]
start_iteration: u32,
#[serde(default = "default_seed_window")]
seed_window: u32,
#[serde(default)]
candidate_recency: Option<u32>,
#[serde(default = "default_max_added_per_round")]
max_added_per_round: u32,
#[serde(default = "default_violation_tolerance")]
violation_tolerance: f64,
},
}
fn default_tie_tolerance() -> f64 {
1e-10
}
fn default_check_frequency() -> u32 {
5
}
fn default_start_iteration() -> u32 {
2
}
fn default_seed_window() -> u32 {
5
}
fn default_max_added_per_round() -> u32 {
10
}
fn default_violation_tolerance() -> f64 {
1e-10
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(default, deny_unknown_fields)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct TrainingSolverConfig {
pub retry_max_attempts: u32,
pub retry_time_budget_seconds: f64,
#[serde(default)]
pub backward: Option<PhaseSolverProfileConfig>,
#[serde(default)]
pub forward: Option<PhaseSolverProfileConfig>,
}
impl Default for TrainingSolverConfig {
fn default() -> Self {
Self {
retry_max_attempts: 5,
retry_time_budget_seconds: 30.0,
backward: None,
forward: None,
}
}
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct PhaseSolverProfileConfig {
#[serde(default)]
pub dual_edge_weight: Option<DualEdgeWeight>,
#[serde(default)]
pub scale: Option<ScaleStrategy>,
#[serde(default)]
pub price: Option<PriceStrategy>,
#[serde(default)]
pub primal_feasibility_tolerance: Option<f64>,
#[serde(default)]
pub dual_feasibility_tolerance: Option<f64>,
#[serde(default)]
pub presolve: Option<PresolveMode>,
#[serde(default)]
pub simplex_update_limit: Option<u32>,
#[serde(default)]
pub cost_perturbation: Option<f64>,
#[serde(default)]
pub refactor_error_tolerance: Option<f64>,
#[serde(default)]
pub factor_pivot_threshold: Option<f64>,
#[serde(default)]
pub use_warm_start: Option<bool>,
#[serde(default)]
pub steepest_edge_devex_fallback_threshold: Option<f64>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize)]
#[serde(rename_all = "snake_case")]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub enum PresolveMode {
On,
Off,
Choose,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Deserialize, Serialize)]
#[serde(default, deny_unknown_fields)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct ParallelismConfig {
pub backward_scheduler: BackwardScheduler,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize)]
#[serde(tag = "method", rename_all = "snake_case", deny_unknown_fields)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub enum BackwardScheduler {
TrialPoint {},
OpeningBlock {
#[serde(default)]
block_size: Option<NonZeroUsize>,
},
}
impl Default for BackwardScheduler {
fn default() -> Self {
Self::TrialPoint {}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize)]
#[serde(rename_all = "snake_case")]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub enum DualEdgeWeight {
Devex,
SteepestEdge,
Dantzig,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize)]
#[serde(rename_all = "snake_case")]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub enum ScaleStrategy {
Off,
SolverScaling,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize)]
#[serde(rename_all = "snake_case")]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub enum PriceStrategy {
Row,
RowHyperSparse,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(tag = "type", rename_all = "snake_case", deny_unknown_fields)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub enum StoppingRuleConfig {
IterationLimit {
limit: u32,
},
TimeLimit {
seconds: f64,
},
BoundStalling {
iterations: u32,
tolerance: f64,
},
Simulation {
replications: u32,
period: u32,
bound_window: u32,
distance_tol: f64,
bound_tol: f64,
},
}
#[derive(Debug, Clone, Deserialize, Serialize, Default)]
#[serde(default, deny_unknown_fields)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct UpperBoundEvaluationConfig {
#[serde(default)]
pub enabled: Option<bool>,
#[serde(default)]
pub initial_iteration: Option<u32>,
#[serde(default)]
pub interval_iterations: Option<u32>,
#[serde(default)]
pub lipschitz: LipschitzConfig,
}
#[derive(Debug, Clone, Deserialize, Serialize, Default)]
#[serde(default, deny_unknown_fields)]
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
pub struct LipschitzConfig {
#[serde(default)]
pub mode: Option<String>,
#[serde(default)]
pub fallback_value: Option<f64>,
#[serde(default)]
pub scale_factor: Option<f64>,
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::{
BackwardScheduler, DualEdgeWeight, NonZeroUsize, PresolveMode, PriceStrategy,
ScaleStrategy, SelectionMethod, TrainingConfig,
};
#[test]
fn dynamic_selection_block_round_trips() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"cut_selection": {
"row_activity_tolerance": 1e-6,
"max_active_per_stage": 4000,
"selection": {
"method": "dynamic",
"start_iteration": 5,
"seed_window": 0,
"candidate_recency": 20,
"max_added_per_round": 3,
"violation_tolerance": 1e-9
}
}
}"#;
let cfg: TrainingConfig = serde_json::from_str(json).unwrap();
let cs = &cfg.cut_selection;
assert_eq!(cs.row_activity_tolerance, Some(1e-6));
assert_eq!(cs.max_active_per_stage, Some(4000));
match cs.selection.as_ref().expect("selection present") {
SelectionMethod::Dynamic {
start_iteration,
seed_window,
candidate_recency,
max_added_per_round,
violation_tolerance,
} => {
assert_eq!(*start_iteration, 5);
assert_eq!(*seed_window, 0);
assert_eq!(*candidate_recency, Some(20));
assert_eq!(*max_added_per_round, 3);
assert!((*violation_tolerance - 1e-9).abs() < f64::EPSILON);
}
other => panic!("expected Dynamic, got {other:?}"),
}
}
#[test]
fn level1_selection_block_round_trips_with_defaults() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"cut_selection": { "selection": { "method": "level1" } }
}"#;
let cfg: TrainingConfig = serde_json::from_str(json).unwrap();
match cfg
.cut_selection
.selection
.as_ref()
.expect("selection present")
{
SelectionMethod::Level1 {
tie_tolerance,
check_frequency,
} => {
assert!((*tie_tolerance - 1e-10).abs() < 1e-20);
assert_eq!(*check_frequency, 5);
}
other => panic!("expected Level1, got {other:?}"),
}
}
#[test]
fn omitting_selection_disables_row_selection() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"cut_selection": {}
}"#;
let cfg: TrainingConfig = serde_json::from_str(json).unwrap();
assert!(cfg.cut_selection.selection.is_none());
}
#[test]
fn wrong_method_field_is_deserialize_error() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"cut_selection": {
"selection": { "method": "level1", "max_added_per_round": 3 }
}
}"#;
let result = serde_json::from_str::<TrainingConfig>(json);
assert!(
result.is_err(),
"a Dynamic-only field under level1 must be rejected"
);
}
#[test]
fn bad_method_string_is_deserialize_error() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"cut_selection": { "selection": { "method": "dynmic" } }
}"#;
let result = serde_json::from_str::<TrainingConfig>(json);
assert!(result.is_err(), "an unknown method tag must be rejected");
}
#[test]
fn domination_without_tolerance_is_missing_field_error() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"cut_selection": { "selection": { "method": "domination" } }
}"#;
let result = serde_json::from_str::<TrainingConfig>(json);
assert!(
result.is_err(),
"domination requires domination_tolerance; absence must be rejected"
);
}
#[test]
fn backward_solver_profile_block_round_trips() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"solver": {
"backward": {
"dual_edge_weight": "steepest_edge",
"scale": "solver_scaling",
"price": "row",
"primal_feasibility_tolerance": 1e-7
}
}
}"#;
let cfg: TrainingConfig = serde_json::from_str(json).unwrap();
let backward = cfg.solver.backward.as_ref().expect("backward present");
assert_eq!(
backward.dual_edge_weight,
Some(DualEdgeWeight::SteepestEdge)
);
assert_eq!(backward.scale, Some(ScaleStrategy::SolverScaling));
assert_eq!(backward.price, Some(PriceStrategy::Row));
assert_eq!(backward.primal_feasibility_tolerance, Some(1e-7));
assert!(cfg.solver.forward.is_none());
}
#[test]
fn backward_solver_profile_new_fields_round_trip() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"solver": {
"backward": {
"presolve": "off",
"use_warm_start": false,
"factor_pivot_threshold": 0.2
}
}
}"#;
let cfg: TrainingConfig = serde_json::from_str(json).unwrap();
let backward = cfg.solver.backward.as_ref().expect("backward present");
assert_eq!(backward.presolve, Some(PresolveMode::Off));
assert_eq!(backward.use_warm_start, Some(false));
assert_eq!(backward.factor_pivot_threshold, Some(0.2));
assert!(backward.dual_feasibility_tolerance.is_none());
assert!(backward.simplex_update_limit.is_none());
assert!(backward.cost_perturbation.is_none());
assert!(backward.refactor_error_tolerance.is_none());
assert!(backward.steepest_edge_devex_fallback_threshold.is_none());
}
#[test]
fn forward_solver_profile_block_round_trips() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"solver": {
"forward": {
"price": "row_hyper_sparse",
"dual_edge_weight": "dantzig"
}
}
}"#;
let cfg: TrainingConfig = serde_json::from_str(json).unwrap();
let forward = cfg.solver.forward.as_ref().expect("forward present");
assert_eq!(forward.price, Some(PriceStrategy::RowHyperSparse));
assert_eq!(forward.dual_edge_weight, Some(DualEdgeWeight::Dantzig));
assert!(cfg.solver.backward.is_none());
}
#[test]
fn backward_solver_profile_unknown_field_is_deserialize_error() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"solver": { "backward": { "dual_edge_weght": "devex" } }
}"#;
let result = serde_json::from_str::<TrainingConfig>(json);
assert!(
result.is_err(),
"an unknown field under backward must be rejected"
);
}
#[test]
fn backward_solver_profile_presolv_typo_is_deserialize_error() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"solver": { "backward": { "presolv": "off" } }
}"#;
let result = serde_json::from_str::<TrainingConfig>(json);
assert!(
result.is_err(),
"the presolv typo under backward must be rejected"
);
}
#[test]
fn backward_solver_profile_bad_enum_value_is_deserialize_error() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"solver": { "backward": { "scale": "curtis_reid" } }
}"#;
let result = serde_json::from_str::<TrainingConfig>(json);
assert!(result.is_err(), "an unknown scale value must be rejected");
}
#[test]
fn backward_scheduler_defaults_to_trial_point_when_absent() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }]
}"#;
let cfg: TrainingConfig = serde_json::from_str(json).unwrap();
assert_eq!(
cfg.parallelism.backward_scheduler,
BackwardScheduler::TrialPoint {}
);
}
#[test]
fn opening_block_scheduler_and_block_size_round_trip() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"parallelism": {
"backward_scheduler": { "method": "opening_block", "block_size": 4 }
}
}"#;
let cfg: TrainingConfig = serde_json::from_str(json).unwrap();
assert_eq!(
cfg.parallelism.backward_scheduler,
BackwardScheduler::OpeningBlock {
block_size: NonZeroUsize::new(4)
}
);
}
#[test]
fn opening_block_scheduler_without_block_size_round_trips() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"parallelism": {
"backward_scheduler": { "method": "opening_block" }
}
}"#;
let cfg: TrainingConfig = serde_json::from_str(json).unwrap();
assert_eq!(
cfg.parallelism.backward_scheduler,
BackwardScheduler::OpeningBlock { block_size: None }
);
}
#[test]
fn block_size_under_trial_point_is_deserialize_error() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"parallelism": {
"backward_scheduler": { "method": "trial_point", "block_size": 4 }
}
}"#;
let result = serde_json::from_str::<TrainingConfig>(json);
assert!(
result.is_err(),
"block_size under trial_point must be rejected"
);
}
#[test]
fn unknown_scheduler_method_is_deserialize_error() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"parallelism": {
"backward_scheduler": { "method": "openin_block" }
}
}"#;
let result = serde_json::from_str::<TrainingConfig>(json);
assert!(result.is_err(), "an unknown method tag must be rejected");
}
#[test]
fn block_size_zero_is_deserialize_error() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [{ "type": "iteration_limit", "limit": 100 }],
"parallelism": {
"backward_scheduler": { "method": "opening_block", "block_size": 0 }
}
}"#;
let result = serde_json::from_str::<TrainingConfig>(json);
assert!(
result.is_err(),
"block_size = 0 must be rejected by NonZeroUsize"
);
}
#[test]
fn removed_root_scheduler_keys_are_rejected() {
for stale in [
r#""backward_scheduler": "opening_block""#,
r#""opening_block_size": 4"#,
r#""backward_opening_order": "sigma_key""#,
] {
let json = format!(
r#"{{
"forward_passes": 4,
"stopping_rules": [{{ "type": "iteration_limit", "limit": 100 }}],
{stale}
}}"#
);
let result = serde_json::from_str::<TrainingConfig>(&json);
assert!(
result.is_err(),
"removed root key must be rejected, got Ok for: {stale}"
);
}
}
#[test]
fn wrong_stopping_rule_field_is_deserialize_error() {
let json = r#"{
"forward_passes": 4,
"stopping_rules": [
{ "type": "iteration_limit", "limit": 100, "seconds": 60.0 }
]
}"#;
let result = serde_json::from_str::<TrainingConfig>(json);
assert!(
result.is_err(),
"a time_limit-only field under iteration_limit must be rejected"
);
}
}