///|
/// Selects a contiguous seed-rotated test segment and sorted training indices.
pub fn holdout_indices(
row_count : Int,
test_count : Int,
seed : Int,
) -> Result[Fold, TreeError] {
if test_count < 1 || test_count >= row_count {
return Err(InvalidTestCount(test_count, row_count))
}
let in_test = Array::make(row_count, false)
let test_indices : Array[Int] = []
let start = positive_mod(seed, row_count)
for offset = 0; offset < test_count; offset = offset + 1 {
let index = (start + offset) % row_count
in_test[index] = true
test_indices.push(index)
}
let train_indices : Array[Int] = []
for index = 0; index < row_count; index = index + 1 {
if !in_test[index] {
train_indices.push(index)
}
}
Ok({ train_indices, test_indices })
}
///|
/// Builds a validated Cartesian product of classification controls.
pub fn classification_config_grid(
max_depths : Array[Int],
min_samples_leaf_values : Array[Int],
criteria : Array[ClassificationCriterion],
) -> Result[Array[ClassificationConfig], TreeError] {
if max_depths.is_empty() ||
min_samples_leaf_values.is_empty() ||
criteria.is_empty() {
return Err(EmptyConfigurationGrid)
}
let configs : Array[ClassificationConfig] = []
for max_depth in max_depths {
for min_samples_leaf in min_samples_leaf_values {
for criterion in criteria {
let min_samples_split = if min_samples_leaf * 2 > 2 {
min_samples_leaf * 2
} else {
2
}
let config = {
..ClassificationConfig::default(),
max_depth,
min_samples_split,
min_samples_leaf,
criterion,
}
match validate_classification_config(config) {
Err(error) => return Err(error)
Ok(_) => configs.push(config)
}
}
}
}
Ok(configs)
}
///|
/// Builds a validated Cartesian product of regression controls.
pub fn regression_config_grid(
max_depths : Array[Int],
min_samples_leaf_values : Array[Int],
) -> Result[Array[RegressionConfig], TreeError] {
if max_depths.is_empty() || min_samples_leaf_values.is_empty() {
return Err(EmptyConfigurationGrid)
}
let configs : Array[RegressionConfig] = []
for max_depth in max_depths {
for min_samples_leaf in min_samples_leaf_values {
let min_samples_split = if min_samples_leaf * 2 > 2 {
min_samples_leaf * 2
} else {
2
}
let config = {
..RegressionConfig::default(),
max_depth,
min_samples_split,
min_samples_leaf,
}
match validate_regression_config(config) {
Err(error) => return Err(error)
Ok(_) => configs.push(config)
}
}
}
Ok(configs)
}
///|
/// Evaluates candidates in input order and keeps the first score tie.
pub fn select_classifier(
dataset : ClassificationDataset,
configs : Array[ClassificationConfig],
fold_count : Int,
seed : Int,
) -> Result[ClassificationSelection, TreeError] {
if configs.is_empty() {
return Err(EmptyConfigurationGrid)
}
let candidates : Array[ClassificationCandidate] = []
let mut best_config = configs[0]
let mut best_accuracy = -1.0
for config in configs {
let validation = match
cross_validate_classifier(dataset, config, fold_count, seed) {
Ok(value) => value
Err(error) => return Err(error)
}
candidates.push({ config, mean_accuracy: validation.mean_accuracy })
if validation.mean_accuracy > best_accuracy + 0.000000000001 {
best_config = config
best_accuracy = validation.mean_accuracy
}
}
Ok({ best_config, best_accuracy, candidates })
}
///|
/// Evaluates candidates in input order and keeps the first MSE tie.
pub fn select_regressor(
dataset : RegressionDataset,
configs : Array[RegressionConfig],
fold_count : Int,
seed : Int,
) -> Result[RegressionSelection, TreeError] {
if configs.is_empty() {
return Err(EmptyConfigurationGrid)
}
let candidates : Array[RegressionCandidate] = []
let mut best_config = configs[0]
let mut best_mse = @double.infinity
let mut best_mae = @double.infinity
for config in configs {
let validation = match
cross_validate_regressor(dataset, config, fold_count, seed) {
Ok(value) => value
Err(error) => return Err(error)
}
candidates.push({
config,
mean_mse: validation.mean_mse,
mean_mae: validation.mean_mae,
})
if validation.mean_mse < best_mse - 0.000000000001 {
best_config = config
best_mse = validation.mean_mse
best_mae = validation.mean_mae
}
}
Ok({ best_config, best_mse, best_mae, candidates })
}