///|
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 })
}