///|
fn log2_positive(value : Double) -> Double {
if value <= 0.0 {
return 0.0
}
let mut scaled = value
let mut exponent = 0
while scaled < 0.5 {
scaled = scaled * 2.0
exponent = exponent - 1
}
while scaled >= 1.0 {
scaled = scaled / 2.0
exponent = exponent + 1
}
let y = (scaled - 1.0) / (scaled + 1.0)
let y_squared = y * y
let mut power = y
let mut denominator = 1
let mut series = 0.0
let mut term_index = 0
while term_index < 24 {
series = series + power / denominator.to_double()
power = power * y_squared
denominator = denominator + 2
term_index = term_index + 1
}
let natural_log = 2.0 * series
exponent.to_double() + natural_log / 0.6931471805599453
}
///|
fn impurity_from_counts(
counts : Array[Int],
sample_count : Int,
criterion : ClassificationCriterion,
) -> Double {
if sample_count == 0 {
return 0.0
}
match criterion {
Gini => {
let mut result = 1.0
for count in counts {
if count > 0 {
let probability = count.to_double() / sample_count.to_double()
result = result - probability * probability
}
}
result
}
Entropy => {
let mut result = 0.0
for count in counts {
if count > 0 {
let probability = count.to_double() / sample_count.to_double()
result = result - probability * log2_positive(probability)
}
}
result
}
}
}
///|
pub fn gini_impurity(labels : Array[Int], class_count : Int) -> Double {
if labels.is_empty() || class_count <= 0 {
return 0.0
}
let counts = Array::make(class_count, 0)
for label in labels {
if label >= 0 && label < class_count {
counts[label] = counts[label] + 1
}
}
impurity_from_counts(counts, labels.length(), Gini)
}
///|
pub fn entropy_impurity(labels : Array[Int], class_count : Int) -> Double {
if labels.is_empty() || class_count <= 0 {
return 0.0
}
let counts = Array::make(class_count, 0)
for label in labels {
if label >= 0 && label < class_count {
counts[label] = counts[label] + 1
}
}
impurity_from_counts(counts, labels.length(), Entropy)
}
///|
fn validate_impurity_labels(
labels : Array[Int],
class_count : Int,
) -> Result[Unit, TreeError] {
if class_count <= 0 {
return Err(InvalidClassCount(class_count))
}
for index, label in labels {
if label < 0 || label >= class_count {
return Err(InvalidClassLabel(index, label, class_count))
}
}
Ok(())
}
///|
/// Validates labels before calculating Gini impurity.
pub fn checked_gini_impurity(
labels : Array[Int],
class_count : Int,
) -> Result[Double, TreeError] {
match validate_impurity_labels(labels, class_count) {
Err(error) => Err(error)
Ok(_) => Ok(gini_impurity(labels, class_count))
}
}
///|
/// Validates labels before calculating entropy impurity.
pub fn checked_entropy_impurity(
labels : Array[Int],
class_count : Int,
) -> Result[Double, TreeError] {
match validate_impurity_labels(labels, class_count) {
Err(error) => Err(error)
Ok(_) => Ok(entropy_impurity(labels, class_count))
}
}
///|
fn compare_feature_indices(
dataset : ClassificationDataset,
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
}
}