///|
fn validate_feature_index(
  feature_index : Int,
  feature_count : Int,
) -> Result[Unit, TreeError] {
  if feature_index < 0 || feature_index >= feature_count {
    Err(InvalidFeatureIndex(feature_index, feature_count))
  } else {
    Ok(())
  }
}

///|
fn rotated_classification_rows(
  dataset : ClassificationDataset,
  feature_index : Int,
  shift : Int,
) -> Array[Array[Double]] {
  Array::makei(dataset.row_count(), fn(row_index) {
    let row = dataset.rows[row_index].copy()
    let source_index = (row_index + shift) % dataset.row_count()
    row[feature_index] = dataset.rows[source_index][feature_index]
    row
  })
}

///|
fn rotated_regression_rows(
  dataset : RegressionDataset,
  feature_index : Int,
  shift : Int,
) -> Array[Array[Double]] {
  Array::makei(dataset.row_count(), fn(row_index) {
    let row = dataset.rows[row_index].copy()
    let source_index = (row_index + shift) % dataset.row_count()
    row[feature_index] = dataset.rows[source_index][feature_index]
    row
  })
}

///|
fn diagnostic_shift(seed : Int, feature_index : Int, row_count : Int) -> Int {
  if row_count <= 1 {
    0
  } else {
    1 + positive_mod(seed + feature_index, row_count - 1)
  }
}

///|
/// Returns baseline accuracy minus accuracy after deterministic column rotation.
pub fn classifier_permutation_importance(
  model : ClassificationTree,
  dataset : ClassificationDataset,
  seed : Int,
) -> Result[Array[Double], TreeError] {
  if model.feature_total != dataset.feature_total {
    return Err(
      PredictionFeatureCountMismatch(dataset.feature_total, model.feature_total),
    )
  }
  let baseline_predictions = match model.predict_batch(dataset.rows) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  let baseline = classification_accuracy(dataset.labels, baseline_predictions).unwrap()
  let importance = Array::make(dataset.feature_total, 0.0)
  for feature_index = 0
      feature_index < dataset.feature_total
      feature_index = feature_index + 1 {
    let shift = diagnostic_shift(seed, feature_index, dataset.row_count())
    let rows = rotated_classification_rows(dataset, feature_index, shift)
    let predicted = model.predict_batch(rows).unwrap()
    let shuffled = classification_accuracy(dataset.labels, predicted).unwrap()
    let decrease = baseline - shuffled
    importance[feature_index] = if decrease > 0.0 { decrease } else { 0.0 }
  }
  Ok(importance)
}

///|
/// Returns shuffled-feature MSE minus baseline MSE.
pub fn regressor_permutation_importance(
  model : RegressionTree,
  dataset : RegressionDataset,
  seed : Int,
) -> Result[Array[Double], TreeError] {
  if model.feature_total != dataset.feature_total {
    return Err(
      PredictionFeatureCountMismatch(dataset.feature_total, model.feature_total),
    )
  }
  let baseline_predictions = match model.predict_batch(dataset.rows) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  let baseline = mean_squared_error(dataset.targets, baseline_predictions).unwrap()
  let importance = Array::make(dataset.feature_total, 0.0)
  for feature_index = 0
      feature_index < dataset.feature_total
      feature_index = feature_index + 1 {
    let shift = diagnostic_shift(seed, feature_index, dataset.row_count())
    let rows = rotated_regression_rows(dataset, feature_index, shift)
    let predicted = model.predict_batch(rows).unwrap()
    let shuffled = mean_squared_error(dataset.targets, predicted).unwrap()
    let increase = shuffled - baseline
    importance[feature_index] = if increase > 0.0 { increase } else { 0.0 }
  }
  Ok(importance)
}

///|
fn classification_dependence_at(
  model : ClassificationTree,
  dataset : ClassificationDataset,
  feature_index : Int,
  feature_value : Double,
) -> Result[Array[Double], TreeError] {
  if !finite_number(feature_value) {
    return Err(NonFiniteFeature(0, feature_index))
  }
  let totals = Array::make(model.class_total, 0.0)
  for source_row in dataset.rows {
    let row = source_row.copy()
    row[feature_index] = feature_value
    let probabilities = model.predict_proba(row).unwrap()
    for class_index = 0
        class_index < model.class_total
        class_index = class_index + 1 {
      totals[class_index] = totals[class_index] + probabilities[class_index]
    }
  }
  Ok(totals.map(fn(total) { total / dataset.row_count().to_double() }))
}

///|
pub fn classifier_partial_dependence(
  model : ClassificationTree,
  dataset : ClassificationDataset,
  feature_index : Int,
  values : Array[Double],
) -> Result[ClassificationPartialDependence, TreeError] {
  match validate_feature_index(feature_index, dataset.feature_total) {
    Err(error) => return Err(error)
    Ok(_) => ()
  }
  if model.feature_total != dataset.feature_total {
    return Err(
      PredictionFeatureCountMismatch(dataset.feature_total, model.feature_total),
    )
  }
  let averages : Array[Array[Double]] = []
  for value in values {
    match classification_dependence_at(model, dataset, feature_index, value) {
      Ok(result) => averages.push(result)
      Err(error) => return Err(error)
    }
  }
  Ok({ feature_index, values: values.copy(), average_probabilities: averages })
}

///|
fn regression_dependence_at(
  model : RegressionTree,
  dataset : RegressionDataset,
  feature_index : Int,
  feature_value : Double,
) -> Result[Double, TreeError] {
  if !finite_number(feature_value) {
    return Err(NonFiniteFeature(0, feature_index))
  }
  let mut total = 0.0
  for source_row in dataset.rows {
    let row = source_row.copy()
    row[feature_index] = feature_value
    total = total + model.predict(row).unwrap()
  }
  Ok(total / dataset.row_count().to_double())
}

///|
pub fn regressor_partial_dependence(
  model : RegressionTree,
  dataset : RegressionDataset,
  feature_index : Int,
  values : Array[Double],
) -> Result[RegressionPartialDependence, TreeError] {
  match validate_feature_index(feature_index, dataset.feature_total) {
    Err(error) => return Err(error)
    Ok(_) => ()
  }
  if model.feature_total != dataset.feature_total {
    return Err(
      PredictionFeatureCountMismatch(dataset.feature_total, model.feature_total),
    )
  }
  let averages : Array[Double] = []
  for value in values {
    match regression_dependence_at(model, dataset, feature_index, value) {
      Ok(result) => averages.push(result)
      Err(error) => return Err(error)
    }
  }
  Ok({ feature_index, values: values.copy(), average_predictions: averages })
}