///|
fn validate_feature_names(
actual : Int,
expected : Int,
) -> Result[Unit, TreeError] {
if actual != expected {
Err(InvalidFeatureNameCount(actual, expected))
} else {
Ok(())
}
}
///|
fn validate_class_names(
actual : Int,
expected : Int,
) -> Result[Unit, TreeError] {
if actual != expected {
Err(InvalidClassNameCount(actual, expected))
} else {
Ok(())
}
}
///|
fn escaped_string(value : String) -> String {
let mut result = ""
for character in value {
let escaped = match character {
'\\' => "\\\\"
'\"' => "\\\""
'\b' => "\\b"
'\n' => "\\n"
'\r' => "\\r"
'\t' => "\\t"
_ =>
if character.to_int() < 0x20 {
let digits = [
"0", "1", "2", "3", "4", "5", "6", "7", "8", "9", "a", "b", "c", "d",
"e", "f",
]
let code = character.to_int()
"\\u00" + digits[code / 16] + digits[code % 16]
} else {
character.to_string()
}
}
result = result + escaped
}
result
}
///|
fn indentation(depth : Int) -> String {
let mut result = ""
for _ in 0.. String {
let prefix = indentation(depth)
match node {
ClassificationLeaf(prediction, _, samples, impurity) =>
"\{prefix}predict=\{class_names[prediction]}, samples=\{samples}, impurity=\{impurity}\n"
ClassificationBranch(
feature,
threshold,
left,
right,
samples,
impurity,
gain
) =>
"\{prefix}\{feature_names[feature]} <= \{threshold}, samples=\{samples}, impurity=\{impurity}, gain=\{gain}\n" +
"\{prefix}left:\n" +
classification_text_node(left, feature_names, class_names, depth + 1) +
"\{prefix}right:\n" +
classification_text_node(right, feature_names, class_names, depth + 1)
}
}
///|
fn regression_text_node(
node : RegressionNode,
feature_names : Array[String],
depth : Int,
) -> String {
let prefix = indentation(depth)
match node {
RegressionLeaf(prediction, samples, variance) =>
"\{prefix}predict=\{prediction}, samples=\{samples}, variance=\{variance}\n"
RegressionBranch(feature, threshold, left, right, samples, variance, gain) =>
"\{prefix}\{feature_names[feature]} <= \{threshold}, samples=\{samples}, variance=\{variance}, gain=\{gain}\n" +
"\{prefix}left:\n" +
regression_text_node(left, feature_names, depth + 1) +
"\{prefix}right:\n" +
regression_text_node(right, feature_names, depth + 1)
}
}
///|
fn classification_json_node(
node : ClassificationNode,
feature_names : Array[String],
class_names : Array[String],
) -> String {
match node {
ClassificationLeaf(prediction, probabilities, samples, impurity) => {
let mut probability_text = "["
for index = 0; index < probabilities.length(); index = index + 1 {
if index > 0 {
probability_text = probability_text + ","
}
probability_text = probability_text + "\{probabilities[index]}"
}
probability_text = probability_text + "]"
"{\"type\":\"leaf\",\"prediction\":\{prediction},\"class_name\":\"\{escaped_string(class_names[prediction])}\",\"probabilities\":\{probability_text},\"samples\":\{samples},\"impurity\":\{impurity}}"
}
ClassificationBranch(
feature,
threshold,
left,
right,
samples,
impurity,
gain
) =>
"{\"type\":\"branch\",\"feature\":\{feature},\"feature_name\":\"\{escaped_string(feature_names[feature])}\",\"threshold\":\{threshold},\"samples\":\{samples},\"impurity\":\{impurity},\"gain\":\{gain},\"left\":\{classification_json_node(left, feature_names, class_names)},\"right\":\{classification_json_node(right, feature_names, class_names)}}"
}
}
///|
fn regression_json_node(
node : RegressionNode,
feature_names : Array[String],
) -> String {
match node {
RegressionLeaf(prediction, samples, variance) =>
"{\"type\":\"leaf\",\"prediction\":\{prediction},\"samples\":\{samples},\"variance\":\{variance}}"
RegressionBranch(feature, threshold, left, right, samples, variance, gain) =>
"{\"type\":\"branch\",\"feature\":\{feature},\"feature_name\":\"\{escaped_string(feature_names[feature])}\",\"threshold\":\{threshold},\"samples\":\{samples},\"variance\":\{variance},\"gain\":\{gain},\"left\":\{regression_json_node(left, feature_names)},\"right\":\{regression_json_node(right, feature_names)}}"
}
}
///|
fn classification_dot_node(
node : ClassificationNode,
feature_names : Array[String],
class_names : Array[String],
node_id : Int,
) -> (String, Int) {
match node {
ClassificationLeaf(prediction, _, samples, impurity) =>
(
" n\{node_id} [shape=box,label=\"predict=\{escaped_string(class_names[prediction])}\\nsamples=\{samples}\\nimpurity=\{impurity}\"];\n",
node_id + 1,
)
ClassificationBranch(feature, threshold, left, right, samples, impurity, _) => {
let left_id = node_id + 1
let (left_text, next_id) = classification_dot_node(
left, feature_names, class_names, left_id,
)
let right_id = next_id
let (right_text, final_id) = classification_dot_node(
right, feature_names, class_names, right_id,
)
let current = " n\{node_id} [label=\"\{escaped_string(feature_names[feature])} <= \{threshold}\\nsamples=\{samples}\\nimpurity=\{impurity}\"];\n"
let edges = " n\{node_id} -> n\{left_id} [label=\"true\"];\n n\{node_id} -> n\{right_id} [label=\"false\"];\n"
(current + left_text + right_text + edges, final_id)
}
}
}
///|
fn regression_dot_node(
node : RegressionNode,
feature_names : Array[String],
node_id : Int,
) -> (String, Int) {
match node {
RegressionLeaf(prediction, samples, variance) =>
(
" n\{node_id} [shape=box,label=\"predict=\{prediction}\\nsamples=\{samples}\\nvariance=\{variance}\"];\n",
node_id + 1,
)
RegressionBranch(feature, threshold, left, right, samples, variance, _) => {
let left_id = node_id + 1
let (left_text, next_id) = regression_dot_node(
left, feature_names, left_id,
)
let right_id = next_id
let (right_text, final_id) = regression_dot_node(
right, feature_names, right_id,
)
let current = " n\{node_id} [label=\"\{escaped_string(feature_names[feature])} <= \{threshold}\\nsamples=\{samples}\\nvariance=\{variance}\"];\n"
let edges = " n\{node_id} -> n\{left_id} [label=\"true\"];\n n\{node_id} -> n\{right_id} [label=\"false\"];\n"
(current + left_text + right_text + edges, final_id)
}
}
}
///|
pub fn ClassificationTree::to_text(
self : ClassificationTree,
feature_names : Array[String],
class_names : Array[String],
) -> Result[String, TreeError] {
match validate_feature_names(feature_names.length(), self.feature_total) {
Err(error) => return Err(error)
Ok(_) => ()
}
match validate_class_names(class_names.length(), self.class_total) {
Err(error) => return Err(error)
Ok(_) => ()
}
Ok(classification_text_node(self.root, feature_names, class_names, 0))
}
///|
pub fn ClassificationTree::to_json(
self : ClassificationTree,
feature_names : Array[String],
class_names : Array[String],
) -> Result[String, TreeError] {
match validate_feature_names(feature_names.length(), self.feature_total) {
Err(error) => return Err(error)
Ok(_) => ()
}
match validate_class_names(class_names.length(), self.class_total) {
Err(error) => return Err(error)
Ok(_) => ()
}
Ok(classification_json_node(self.root, feature_names, class_names))
}
///|
pub fn ClassificationTree::to_dot(
self : ClassificationTree,
feature_names : Array[String],
class_names : Array[String],
) -> Result[String, TreeError] {
match validate_feature_names(feature_names.length(), self.feature_total) {
Err(error) => return Err(error)
Ok(_) => ()
}
match validate_class_names(class_names.length(), self.class_total) {
Err(error) => return Err(error)
Ok(_) => ()
}
let (body, _) = classification_dot_node(
self.root,
feature_names,
class_names,
0,
)
Ok("digraph DecisionTree {\n" + body + "}\n")
}
///|
pub fn RegressionTree::to_text(
self : RegressionTree,
feature_names : Array[String],
) -> Result[String, TreeError] {
match validate_feature_names(feature_names.length(), self.feature_total) {
Err(error) => return Err(error)
Ok(_) => ()
}
Ok(regression_text_node(self.root, feature_names, 0))
}
///|
pub fn RegressionTree::to_json(
self : RegressionTree,
feature_names : Array[String],
) -> Result[String, TreeError] {
match validate_feature_names(feature_names.length(), self.feature_total) {
Err(error) => return Err(error)
Ok(_) => ()
}
Ok(regression_json_node(self.root, feature_names))
}
///|
pub fn RegressionTree::to_dot(
self : RegressionTree,
feature_names : Array[String],
) -> Result[String, TreeError] {
match validate_feature_names(feature_names.length(), self.feature_total) {
Err(error) => return Err(error)
Ok(_) => ()
}
let (body, _) = regression_dot_node(self.root, feature_names, 0)
Ok("digraph DecisionTree {\n" + body + "}\n")
}