///|
fn classification_leaf_from_subtree(
node : ClassificationNode,
class_count : Int,
impurity : Double,
) -> ClassificationNode {
let weighted = Array::make(class_count, 0.0)
let sample_count = collect_classification_probabilities(node, weighted)
let probabilities = weighted.map(fn(value) {
value / sample_count.to_double()
})
let mut prediction = 0
for class_index = 1; class_index < class_count; class_index = class_index + 1 {
if probabilities[class_index] > probabilities[prediction] {
prediction = class_index
}
}
ClassificationLeaf(prediction, probabilities, sample_count, impurity)
}
///|
fn collect_classification_probabilities(
node : ClassificationNode,
weighted : Array[Double],
) -> Int {
match node {
ClassificationLeaf(_, probabilities, samples, _) => {
for class_index = 0
class_index < probabilities.length()
class_index = class_index + 1 {
weighted[class_index] = weighted[class_index] +
probabilities[class_index] * samples.to_double()
}
samples
}
ClassificationBranch(_, _, left, right, _, _, _) =>
collect_classification_probabilities(left, weighted) +
collect_classification_probabilities(right, weighted)
}
}
///|
fn prune_classification_node(
node : ClassificationNode,
class_count : Int,
alpha : Double,
) -> (ClassificationNode, Double, Int) {
match node {
ClassificationLeaf(_, _, samples, impurity) =>
(node, samples.to_double() * impurity, 1)
ClassificationBranch(
feature,
threshold,
left,
right,
samples,
impurity,
gain
) => {
let (pruned_left, left_risk, left_leaves) = prune_classification_node(
left, class_count, alpha,
)
let (pruned_right, right_risk, right_leaves) = prune_classification_node(
right, class_count, alpha,
)
let subtree_risk = left_risk + right_risk
let leaf_total = left_leaves + right_leaves
let collapsed_risk = samples.to_double() * impurity
let branch = ClassificationBranch(
feature, threshold, pruned_left, pruned_right, samples, impurity, gain,
)
if collapsed_risk + alpha <= subtree_risk + alpha * leaf_total.to_double() {
(
classification_leaf_from_subtree(branch, class_count, impurity),
collapsed_risk,
1,
)
} else {
(branch, subtree_risk, leaf_total)
}
}
}
}
///|
fn regression_leaf_from_subtree(
node : RegressionNode,
variance : Double,
) -> RegressionNode {
let (weighted_sum, samples) = collect_regression_predictions(node)
RegressionLeaf(weighted_sum / samples.to_double(), samples, variance)
}
///|
fn collect_regression_predictions(node : RegressionNode) -> (Double, Int) {
match node {
RegressionLeaf(prediction, samples, _) =>
(prediction * samples.to_double(), samples)
RegressionBranch(_, _, left, right, _, _, _) => {
let (left_sum, left_samples) = collect_regression_predictions(left)
let (right_sum, right_samples) = collect_regression_predictions(right)
(left_sum + right_sum, left_samples + right_samples)
}
}
}
///|
fn prune_regression_node(
node : RegressionNode,
alpha : Double,
) -> (RegressionNode, Double, Int) {
match node {
RegressionLeaf(_, samples, variance) =>
(node, samples.to_double() * variance, 1)
RegressionBranch(feature, threshold, left, right, samples, variance, gain) => {
let (pruned_left, left_risk, left_leaves) = prune_regression_node(
left, alpha,
)
let (pruned_right, right_risk, right_leaves) = prune_regression_node(
right, alpha,
)
let subtree_risk = left_risk + right_risk
let leaf_total = left_leaves + right_leaves
let collapsed_risk = samples.to_double() * variance
let branch = RegressionBranch(
feature, threshold, pruned_left, pruned_right, samples, variance, gain,
)
if collapsed_risk + alpha <= subtree_risk + alpha * leaf_total.to_double() {
(regression_leaf_from_subtree(branch, variance), collapsed_risk, 1)
} else {
(branch, subtree_risk, leaf_total)
}
}
}
}
///|
/// Returns a pruned copy using leaf-risk plus alpha times leaf count.
pub fn ClassificationTree::prune(
self : ClassificationTree,
alpha : Double,
) -> Result[ClassificationTree, TreeError] {
if !finite_number(alpha) || alpha < 0.0 {
return Err(InvalidPruningAlpha(alpha))
}
if alpha == 0.0 {
return Ok(self)
}
let (root, _, _) = prune_classification_node(
self.root,
self.class_total,
alpha,
)
Ok({ ..self, root, })
}
///|
/// Returns a pruned copy using leaf-risk plus alpha times leaf count.
pub fn RegressionTree::prune(
self : RegressionTree,
alpha : Double,
) -> Result[RegressionTree, TreeError] {
if !finite_number(alpha) || alpha < 0.0 {
return Err(InvalidPruningAlpha(alpha))
}
if alpha == 0.0 {
return Ok(self)
}
let (root, _, _) = prune_regression_node(self.root, alpha)
Ok({ ..self, root, })
}