///|
fn score_standard_deviation(values : Array[Double], mean : Double) -> Double {
  if values.is_empty() {
    return 0.0
  }
  let mut squared_sum = 0.0
  for value in values {
    let difference = value - mean
    squared_sum = squared_sum + difference * difference
  }
  square_root(squared_sum / values.length().to_double())
}

///|
fn score_minimum(values : Array[Double]) -> Double {
  let mut minimum = values[0]
  for index = 1; index < values.length(); index = index + 1 {
    if values[index] < minimum {
      minimum = values[index]
    }
  }
  minimum
}

///|
fn score_maximum(values : Array[Double]) -> Double {
  let mut maximum = values[0]
  for index = 1; index < values.length(); index = index + 1 {
    if values[index] > maximum {
      maximum = values[index]
    }
  }
  maximum
}

///|
/// Repeats stratified cross-validation with deterministic derived seeds.
pub fn repeated_cross_validate_classifier(
  dataset : ClassificationDataset,
  config : ClassificationConfig,
  fold_count : Int,
  repeat_count : Int,
  seed : Int,
) -> Result[RepeatedClassificationValidation, TreeError] {
  if repeat_count < 1 {
    return Err(InvalidRepeatCount(repeat_count))
  }
  let repeat_means : Array[Double] = []
  let fold_scores : Array[Double] = []
  for repeat_index = 0
      repeat_index < repeat_count
      repeat_index = repeat_index + 1 {
    let repeat_seed = seed + repeat_index * 104729
    let result = match
      cross_validate_classifier(dataset, config, fold_count, repeat_seed) {
      Ok(value) => value
      Err(error) => return Err(error)
    }
    repeat_means.push(result.mean_accuracy)
    for score in result.fold_scores {
      fold_scores.push(score)
    }
  }
  let mean_accuracy = mean_value(fold_scores)
  Ok({
    repeat_means,
    fold_scores,
    mean_accuracy,
    standard_deviation: score_standard_deviation(fold_scores, mean_accuracy),
    minimum_accuracy: score_minimum(fold_scores),
    maximum_accuracy: score_maximum(fold_scores),
  })
}

///|
/// Repeats regression cross-validation with deterministic derived seeds.
pub fn repeated_cross_validate_regressor(
  dataset : RegressionDataset,
  config : RegressionConfig,
  fold_count : Int,
  repeat_count : Int,
  seed : Int,
) -> Result[RepeatedRegressionValidation, TreeError] {
  if repeat_count < 1 {
    return Err(InvalidRepeatCount(repeat_count))
  }
  let repeat_mse : Array[Double] = []
  let repeat_mae : Array[Double] = []
  let fold_mse : Array[Double] = []
  let fold_mae : Array[Double] = []
  for repeat_index = 0
      repeat_index < repeat_count
      repeat_index = repeat_index + 1 {
    let repeat_seed = seed + repeat_index * 104729
    let result = match
      cross_validate_regressor(dataset, config, fold_count, repeat_seed) {
      Ok(value) => value
      Err(error) => return Err(error)
    }
    repeat_mse.push(result.mean_mse)
    repeat_mae.push(result.mean_mae)
    for score in result.fold_mse {
      fold_mse.push(score)
    }
    for score in result.fold_mae {
      fold_mae.push(score)
    }
  }
  let mean_mse = mean_value(fold_mse)
  let mean_mae = mean_value(fold_mae)
  Ok({
    repeat_mse,
    repeat_mae,
    fold_mse,
    fold_mae,
    mean_mse,
    mean_mae,
    mse_standard_deviation: score_standard_deviation(fold_mse, mean_mse),
    mae_standard_deviation: score_standard_deviation(fold_mae, mean_mae),
  })
}

///|
fn validate_pruning_alphas(alphas : Array[Double]) -> Result[Unit, TreeError] {
  if alphas.is_empty() {
    return Err(EmptyPruningGrid)
  }
  for alpha in alphas {
    if !finite_number(alpha) || alpha < 0.0 {
      return Err(InvalidPruningAlpha(alpha))
    }
  }
  Ok(())
}

///|
fn classification_pruning_score(
  dataset : ClassificationDataset,
  config : ClassificationConfig,
  folds : Array[Fold],
  alpha : Double,
) -> Result[Double, TreeError] {
  let scores : Array[Double] = []
  for fold in folds {
    let trained = match
      train_classifier(
        classification_subset(dataset, fold.train_indices),
        config,
      ) {
      Ok(value) => value
      Err(error) => return Err(error)
    }
    let model = trained.prune(alpha).unwrap()
    let rows = fold.test_indices.map(fn(index) { dataset.rows[index].copy() })
    let labels = fold.test_indices.map(fn(index) { dataset.labels[index] })
    let predicted = model.predict_batch(rows).unwrap()
    scores.push(classification_accuracy(labels, predicted).unwrap())
  }
  Ok(mean_value(scores))
}

///|
/// Selects alpha by mean validation accuracy, then by fewer full-data nodes.
pub fn select_classifier_pruning_alpha(
  dataset : ClassificationDataset,
  config : ClassificationConfig,
  alphas : Array[Double],
  fold_count : Int,
  seed : Int,
) -> Result[ClassificationPruningSelection, TreeError] {
  match validate_pruning_alphas(alphas) {
    Err(error) => return Err(error)
    Ok(_) => ()
  }
  let folds = match
    stratified_k_fold_indices(
      dataset.labels,
      dataset.class_total,
      fold_count,
      seed,
    ) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  let full = match train_classifier(dataset, config) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  let candidates : Array[ClassificationPruningCandidate] = []
  let mut best_alpha = alphas[0]
  let mut best_accuracy = -1.0
  let mut best_nodes = 2147483647
  let mut best_model = full
  for alpha in alphas {
    let mean_accuracy = match
      classification_pruning_score(dataset, config, folds, alpha) {
      Ok(value) => value
      Err(error) => return Err(error)
    }
    let model = full.prune(alpha).unwrap()
    let node_count = model.node_count()
    let leaf_count = model.leaf_count()
    candidates.push({ alpha, mean_accuracy, node_count, leaf_count })
    if mean_accuracy > best_accuracy + 0.000000000001 ||
      ((mean_accuracy - best_accuracy).is_close(0.0) && node_count < best_nodes) {
      best_alpha = alpha
      best_accuracy = mean_accuracy
      best_nodes = node_count
      best_model = model
    }
  }
  Ok({ best_alpha, best_accuracy, model: best_model, candidates })
}

///|
fn regression_pruning_score(
  dataset : RegressionDataset,
  config : RegressionConfig,
  folds : Array[Fold],
  alpha : Double,
) -> Result[(Double, Double), TreeError] {
  let mse_scores : Array[Double] = []
  let mae_scores : Array[Double] = []
  for fold in folds {
    let trained = match
      train_regressor(regression_subset(dataset, fold.train_indices), config) {
      Ok(value) => value
      Err(error) => return Err(error)
    }
    let model = trained.prune(alpha).unwrap()
    let rows = fold.test_indices.map(fn(index) { dataset.rows[index].copy() })
    let targets = fold.test_indices.map(fn(index) { dataset.targets[index] })
    let predicted = model.predict_batch(rows).unwrap()
    mse_scores.push(mean_squared_error(targets, predicted).unwrap())
    mae_scores.push(mean_absolute_error(targets, predicted).unwrap())
  }
  Ok((mean_value(mse_scores), mean_value(mae_scores)))
}

///|
/// Selects alpha by mean validation MSE, then by fewer full-data nodes.
pub fn select_regressor_pruning_alpha(
  dataset : RegressionDataset,
  config : RegressionConfig,
  alphas : Array[Double],
  fold_count : Int,
  seed : Int,
) -> Result[RegressionPruningSelection, TreeError] {
  match validate_pruning_alphas(alphas) {
    Err(error) => return Err(error)
    Ok(_) => ()
  }
  let folds = match k_fold_indices(dataset.row_count(), fold_count, seed) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  let full = match train_regressor(dataset, config) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  let candidates : Array[RegressionPruningCandidate] = []
  let mut best_alpha = alphas[0]
  let mut best_mse = @double.infinity
  let mut best_mae = @double.infinity
  let mut best_nodes = 2147483647
  let mut best_model = full
  for alpha in alphas {
    let (mean_mse, mean_mae) = match
      regression_pruning_score(dataset, config, folds, alpha) {
      Ok(value) => value
      Err(error) => return Err(error)
    }
    let model = full.prune(alpha).unwrap()
    let node_count = model.node_count()
    let leaf_count = model.leaf_count()
    candidates.push({ alpha, mean_mse, mean_mae, node_count, leaf_count })
    if mean_mse < best_mse - 0.000000000001 ||
      ((mean_mse - best_mse).is_close(0.0) && node_count < best_nodes) {
      best_alpha = alpha
      best_mse = mean_mse
      best_mae = mean_mae
      best_nodes = node_count
      best_model = model
    }
  }
  Ok({ best_alpha, best_mse, best_mae, model: best_model, candidates })
}