///|
/// Selects the impurity measure used for classification splits.
pub(all) enum ClassificationCriterion {
Gini
Entropy
} derive(Eq, Debug)
///|
/// Configures deterministic classification-tree training.
pub(all) struct ClassificationConfig {
max_depth : Int
min_samples_split : Int
min_samples_leaf : Int
min_impurity_decrease : Double
criterion : ClassificationCriterion
} derive(Eq, Debug)
///|
pub fn ClassificationConfig::default() -> ClassificationConfig {
{
max_depth: 10,
min_samples_split: 2,
min_samples_leaf: 1,
min_impurity_decrease: 0.0,
criterion: Gini,
}
}
///|
/// Retains all statistics needed for prediction, export, and pruning.
pub(all) enum ClassificationNode {
ClassificationLeaf(Int, Array[Double], Int, Double)
ClassificationBranch(
Int,
Double,
ClassificationNode,
ClassificationNode,
Int,
Double,
Double
)
} derive(Eq, Debug)
///|
/// Stores one immutable trained classification tree.
pub(all) struct ClassificationTree {
root : ClassificationNode
feature_total : Int
class_total : Int
config : ClassificationConfig
} derive(Eq, Debug)
///|
priv struct ClassificationSplit {
feature_index : Int
threshold : Double
gain : Double
}