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 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 enabled: Option<bool>,
#[serde(default)]
pub method: Option<String>,
#[serde(default)]
pub threshold: Option<u32>,
#[serde(default)]
pub memory_window: Option<u32>,
#[serde(default)]
pub domination_epsilon: Option<f64>,
#[serde(default)]
pub check_frequency: Option<u32>,
#[serde(default)]
pub cut_activity_tolerance: Option<f64>,
#[serde(default)]
pub basis_activity_window: Option<u32>,
#[serde(default)]
pub max_active_per_stage: Option<u32>,
}
#[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,
}
impl Default for TrainingSolverConfig {
fn default() -> Self {
Self {
retry_max_attempts: 5,
retry_time_budget_seconds: 30.0,
}
}
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(tag = "type", rename_all = "snake_case")]
#[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>,
}