///|
fn positive_mod(value : Int, modulus : Int) -> Int {
  let remainder = value % modulus
  if remainder < 0 {
    remainder + modulus
  } else {
    remainder
  }
}

///|
fn folds_from_tests(
  test_sets : Array[Array[Int]],
  row_count : Int,
) -> Array[Fold] {
  test_sets.map(fn(test_indices) {
    let in_test = Array::make(row_count, false)
    for index in test_indices {
      in_test[index] = true
    }
    let train_indices : Array[Int] = []
    for index = 0; index < row_count; index = index + 1 {
      if !in_test[index] {
        train_indices.push(index)
      }
    }
    { train_indices, test_indices: test_indices.copy() }
  })
}

///|
/// Creates deterministic folds from a seed-controlled cyclic order.
pub fn k_fold_indices(
  row_count : Int,
  fold_count : Int,
  seed : Int,
) -> Result[Array[Fold], TreeError] {
  if fold_count < 2 || fold_count > row_count {
    return Err(InvalidFoldCount(fold_count, row_count))
  }
  let test_sets = Array::makei(fold_count, fn(_) { [] })
  let start = positive_mod(seed, row_count)
  for position = 0; position < row_count; position = position + 1 {
    let index = (start + position) % row_count
    test_sets[position % fold_count].push(index)
  }
  Ok(folds_from_tests(test_sets, row_count))
}

///|
/// Distributes each class round-robin while retaining deterministic order.
pub fn stratified_k_fold_indices(
  labels : Array[Int],
  class_count : Int,
  fold_count : Int,
  seed : Int,
) -> Result[Array[Fold], TreeError] {
  if class_count <= 0 {
    return Err(InvalidClassCount(class_count))
  }
  if fold_count < 2 || fold_count > labels.length() {
    return Err(InvalidFoldCount(fold_count, labels.length()))
  }
  let by_class = Array::makei(class_count, fn(_) { [] })
  for index, label in labels {
    if label < 0 || label >= class_count {
      return Err(InvalidClassLabel(index, label, class_count))
    }
    by_class[label].push(index)
  }
  let test_sets = Array::makei(fold_count, fn(_) { [] })
  let mut next_fold = positive_mod(seed, fold_count)
  for class_index = 0; class_index < class_count; class_index = class_index + 1 {
    let indices = by_class[class_index]
    if !indices.is_empty() {
      let start = positive_mod(seed + class_index, indices.length())
      for position = 0; position < indices.length(); position = position + 1 {
        let source_index = indices[(start + position) % indices.length()]
        test_sets[next_fold].push(source_index)
        next_fold = (next_fold + 1) % fold_count
      }
    }
  }
  Ok(folds_from_tests(test_sets, labels.length()))
}

///|
fn classification_subset(
  dataset : ClassificationDataset,
  indices : Array[Int],
) -> ClassificationDataset {
  let rows = indices.map(fn(index) { dataset.rows[index].copy() })
  let labels = indices.map(fn(index) { dataset.labels[index] })
  classification_dataset(rows, labels, dataset.class_total).unwrap()
}

///|
fn regression_subset(
  dataset : RegressionDataset,
  indices : Array[Int],
) -> RegressionDataset {
  let rows = indices.map(fn(index) { dataset.rows[index].copy() })
  let targets = indices.map(fn(index) { dataset.targets[index] })
  regression_dataset(rows, targets).unwrap()
}

///|
pub fn cross_validate_classifier(
  dataset : ClassificationDataset,
  config : ClassificationConfig,
  fold_count : Int,
  seed : Int,
) -> Result[ClassificationValidation, TreeError] {
  let folds = match
    stratified_k_fold_indices(
      dataset.labels,
      dataset.class_total,
      fold_count,
      seed,
    ) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  let scores : Array[Double] = []
  for fold in folds {
    let model = match
      train_classifier(
        classification_subset(dataset, fold.train_indices),
        config,
      ) {
      Ok(value) => value
      Err(error) => return Err(error)
    }
    let test_rows = fold.test_indices.map(fn(index) {
      dataset.rows[index].copy()
    })
    let actual = fold.test_indices.map(fn(index) { dataset.labels[index] })
    let predicted = match model.predict_batch(test_rows) {
      Ok(value) => value
      Err(error) => return Err(error)
    }
    match classification_accuracy(actual, predicted) {
      Ok(value) => scores.push(value)
      Err(error) => return Err(error)
    }
  }
  Ok({ fold_scores: scores, mean_accuracy: mean_value(scores) })
}

///|
pub fn cross_validate_regressor(
  dataset : RegressionDataset,
  config : RegressionConfig,
  fold_count : Int,
  seed : Int,
) -> Result[RegressionValidation, TreeError] {
  let folds = match k_fold_indices(dataset.row_count(), fold_count, seed) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  let mse_scores : Array[Double] = []
  let mae_scores : Array[Double] = []
  for fold in folds {
    let model = match
      train_regressor(regression_subset(dataset, fold.train_indices), config) {
      Ok(value) => value
      Err(error) => return Err(error)
    }
    let test_rows = fold.test_indices.map(fn(index) {
      dataset.rows[index].copy()
    })
    let actual = fold.test_indices.map(fn(index) { dataset.targets[index] })
    let predicted = match model.predict_batch(test_rows) {
      Ok(value) => value
      Err(error) => return Err(error)
    }
    match mean_squared_error(actual, predicted) {
      Ok(value) => mse_scores.push(value)
      Err(error) => return Err(error)
    }
    match mean_absolute_error(actual, predicted) {
      Ok(value) => mae_scores.push(value)
      Err(error) => return Err(error)
    }
  }
  Ok({
    fold_mse: mse_scores,
    fold_mae: mae_scores,
    mean_mse: mean_value(mse_scores),
    mean_mae: mean_value(mae_scores),
  })
}