///|
pub fn mean_value(values : Array[Double]) -> Double {
  if values.is_empty() {
    return 0.0
  }
  let mut sum = 0.0
  for value in values {
    sum = sum + value
  }
  sum / values.length().to_double()
}

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

///|
fn validate_regression_config(
  config : RegressionConfig,
) -> Result[Unit, TreeError] {
  if config.max_depth < 0 {
    return Err(InvalidMaxDepth(config.max_depth))
  }
  if config.min_samples_split < 2 {
    return Err(InvalidMinSamplesSplit(config.min_samples_split))
  }
  if config.min_samples_leaf < 1 {
    return Err(InvalidMinSamplesLeaf(config.min_samples_leaf))
  }
  if !finite_number(config.min_impurity_decrease) ||
    config.min_impurity_decrease < 0.0 {
    return Err(InvalidMinImpurityDecrease(config.min_impurity_decrease))
  }
  Ok(())
}

///|
fn regression_targets(
  dataset : RegressionDataset,
  indices : Array[Int],
) -> Array[Double] {
  indices.map(fn(index) { dataset.targets[index] })
}

///|
fn regression_leaf(
  dataset : RegressionDataset,
  indices : Array[Int],
) -> RegressionNode {
  let targets = regression_targets(dataset, indices)
  RegressionLeaf(mean_value(targets), indices.length(), variance_value(targets))
}

///|
fn compare_regression_indices(
  dataset : RegressionDataset,
  feature_index : Int,
  left : Int,
  right : Int,
) -> Int {
  let left_value = dataset.rows[left][feature_index]
  let right_value = dataset.rows[right][feature_index]
  if left_value < right_value {
    -1
  } else if left_value > right_value {
    1
  } else if left < right {
    -1
  } else if left > right {
    1
  } else {
    0
  }
}

///|
fn best_regression_split(
  dataset : RegressionDataset,
  indices : Array[Int],
  config : RegressionConfig,
  parent_variance : Double,
) -> RegressionSplit? {
  let total = indices.length()
  let mut best : RegressionSplit? = None
  let mut best_gain = -1.0
  for feature_index = 0
      feature_index < dataset.feature_total
      feature_index = feature_index + 1 {
    let sorted = indices.copy()
    sorted.sort_by(fn(left, right) {
      compare_regression_indices(dataset, feature_index, left, right)
    })
    let mut left_sum = 0.0
    let mut left_squared_sum = 0.0
    let mut right_sum = 0.0
    let mut right_squared_sum = 0.0
    for index in sorted {
      let target = dataset.targets[index]
      right_sum = right_sum + target
      right_squared_sum = right_squared_sum + target * target
    }
    for split_index = 1; split_index < total; split_index = split_index + 1 {
      let target = dataset.targets[sorted[split_index - 1]]
      left_sum = left_sum + target
      left_squared_sum = left_squared_sum + target * target
      right_sum = right_sum - target
      right_squared_sum = right_squared_sum - target * target
      let left_size = split_index
      let right_size = total - split_index
      if left_size < config.min_samples_leaf ||
        right_size < config.min_samples_leaf {
        continue
      }
      let left_value = dataset.rows[sorted[split_index - 1]][feature_index]
      let right_value = dataset.rows[sorted[split_index]][feature_index]
      if left_value == right_value {
        continue
      }
      let left_variance = left_squared_sum / left_size.to_double() -
        left_sum / left_size.to_double() * (left_sum / left_size.to_double())
      let right_variance = right_squared_sum / right_size.to_double() -
        right_sum /
        right_size.to_double() *
        (right_sum / right_size.to_double())
      let weighted = left_size.to_double() / total.to_double() * left_variance +
        right_size.to_double() / total.to_double() * right_variance
      let gain = parent_variance - weighted
      if gain > best_gain + 0.000000000001 {
        best_gain = gain
        best = Some({
          feature_index,
          threshold: left_value + (right_value - left_value) / 2.0,
          gain,
        })
      }
    }
  }
  best
}

///|
fn build_regression_node(
  dataset : RegressionDataset,
  indices : Array[Int],
  config : RegressionConfig,
  depth : Int,
) -> RegressionNode {
  let variance = variance_value(regression_targets(dataset, indices))
  if depth >= config.max_depth ||
    indices.length() < config.min_samples_split ||
    variance <= 0.000000000001 {
    return regression_leaf(dataset, indices)
  }
  let split = match best_regression_split(dataset, indices, config, variance) {
    Some(value) => value
    None => return regression_leaf(dataset, indices)
  }
  if split.gain + 0.000000000001 < config.min_impurity_decrease ||
    split.gain <= 0.000000000001 {
    return regression_leaf(dataset, indices)
  }
  let left_indices : Array[Int] = []
  let right_indices : Array[Int] = []
  for index in indices {
    if dataset.rows[index][split.feature_index] <= split.threshold {
      left_indices.push(index)
    } else {
      right_indices.push(index)
    }
  }
  if left_indices.is_empty() || right_indices.is_empty() {
    return regression_leaf(dataset, indices)
  }
  RegressionBranch(
    split.feature_index,
    split.threshold,
    build_regression_node(dataset, left_indices, config, depth + 1),
    build_regression_node(dataset, right_indices, config, depth + 1),
    indices.length(),
    variance,
    split.gain,
  )
}

///|
/// Trains one deterministic regression tree.
pub fn train_regressor(
  dataset : RegressionDataset,
  config : RegressionConfig,
) -> Result[RegressionTree, TreeError] {
  match validate_regression_config(config) {
    Err(error) => return Err(error)
    Ok(_) => ()
  }
  let indices = Array::makei(dataset.row_count(), fn(index) { index })
  Ok({
    root: build_regression_node(dataset, indices, config, 0),
    feature_total: dataset.feature_total,
    config,
  })
}

///|
fn regression_prediction(
  node : RegressionNode,
  features : Array[Double],
) -> Double {
  match node {
    RegressionLeaf(prediction, _, _) => prediction
    RegressionBranch(feature, threshold, left, right, _, _, _) =>
      if features[feature] <= threshold {
        regression_prediction(left, features)
      } else {
        regression_prediction(right, features)
      }
  }
}

///|
pub fn RegressionTree::predict(
  self : RegressionTree,
  features : Array[Double],
) -> Result[Double, TreeError] {
  if features.length() != self.feature_total {
    return Err(
      PredictionFeatureCountMismatch(features.length(), self.feature_total),
    )
  }
  Ok(regression_prediction(self.root, features))
}

///|
pub fn RegressionTree::predict_batch(
  self : RegressionTree,
  rows : Array[Array[Double]],
) -> Result[Array[Double], TreeError] {
  let predictions : Array[Double] = []
  for row in rows {
    match self.predict(row) {
      Ok(value) => predictions.push(value)
      Err(error) => return Err(error)
    }
  }
  Ok(predictions)
}