IMoTE 0.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- imote/__init__.py +28 -0
- imote/adapters/__init__.py +3 -0
- imote/adapters/adapter_utils.py +15 -0
- imote/adapters/base_adapter.py +76 -0
- imote/adapters/build_root_node_m5.py +71 -0
- imote/adapters/m5_adapter.py +18 -0
- imote/adapters/partykit_adapter.py +133 -0
- imote/adapters/pilot_adapter.py +154 -0
- imote/app.py +80 -0
- imote/benchmark_info.py +29 -0
- imote/callbacks/__init__.py +16 -0
- imote/callbacks/a_tree_info_callbacks.py +110 -0
- imote/callbacks/b_node_callbacks.py +186 -0
- imote/callbacks/c_new_tree_callbacks.py +380 -0
- imote/callbacks/d_edit_tree_callbacks.py +192 -0
- imote/callbacks/e_layout_callbacks.py +39 -0
- imote/callbacks/f_highlight_callbacks.py +79 -0
- imote/callbacks/g_elements_callbacks.py +195 -0
- imote/components/__init__.py +0 -0
- imote/components/cards/__init__.py +0 -0
- imote/components/cards/a_tree_info_card.py +14 -0
- imote/components/cards/b_node_info_card.py +79 -0
- imote/components/cards/bb_node_plot_settings_card.py +139 -0
- imote/components/cards/c_new_tree_card.py +236 -0
- imote/components/cards/d_edit_tree_card.py +159 -0
- imote/components/cards/e_layout_card.py +66 -0
- imote/components/cards/f_highlight_card.py +48 -0
- imote/components/cytoscape_graph.py +39 -0
- imote/components/stores.py +20 -0
- imote/config.py +252 -0
- imote/data/initial_tree.pkl +0 -0
- imote/dataset/__init__.py +0 -0
- imote/dataset/dataset.py +170 -0
- imote/dataset/dataset_registry.py +23 -0
- imote/ids.py +124 -0
- imote/node_metrics/__init__.py +0 -0
- imote/node_metrics/node_metric.py +167 -0
- imote/nodes/__init__.py +0 -0
- imote/nodes/base_node.py +139 -0
- imote/nodes/collapsed_node.py +46 -0
- imote/nodes/combined_lin_node.py +78 -0
- imote/nodes/internal_node.py +129 -0
- imote/nodes/leaf_node.py +74 -0
- imote/nodes/node_model.py +342 -0
- imote/nodes/none_node.py +56 -0
- imote/nodes/split_node.py +219 -0
- imote/pages/__init__.py +0 -0
- imote/pages/explain_page.py +40 -0
- imote/pages/tree_view_page.py +33 -0
- imote/plots/__init__.py +0 -0
- imote/plots/predsplot.py +485 -0
- imote/plots/predsplot2.py +540 -0
- imote/plots/regplot.py +81 -0
- imote/viz_tree/__init__.py +0 -0
- imote/viz_tree/contributions.py +113 -0
- imote/viz_tree/viz_tree.py +432 -0
- imote/viz_tree/viz_tree_cytoscape.py +293 -0
- imote-0.0.0.dist-info/METADATA +113 -0
- imote-0.0.0.dist-info/RECORD +63 -0
- imote-0.0.0.dist-info/WHEEL +5 -0
- imote-0.0.0.dist-info/entry_points.txt +2 -0
- imote-0.0.0.dist-info/licenses/LICENSE +21 -0
- imote-0.0.0.dist-info/top_level.txt +1 -0
imote/__init__.py
ADDED
|
@@ -0,0 +1,28 @@
|
|
|
1
|
+
"""IMoTE: Interactive MOdel Tree Explorer."""
|
|
2
|
+
|
|
3
|
+
from importlib.metadata import PackageNotFoundError, version
|
|
4
|
+
|
|
5
|
+
try:
|
|
6
|
+
__version__ = version("imote") # set from the git tag by setuptools-scm at install time
|
|
7
|
+
except PackageNotFoundError: # running from a source tree that was never pip installed
|
|
8
|
+
__version__ = "0+unknown"
|
|
9
|
+
|
|
10
|
+
from imote.adapters import M5Adapter, PartyKitAdapter, PilotAdapter
|
|
11
|
+
from imote.adapters.base_adapter import ADAPTERS_REGISTRY, BaseAdapter, register_adapter
|
|
12
|
+
from imote.dataset.dataset import Dataset
|
|
13
|
+
from imote.node_metrics.node_metric import NODE_METRICS_REGISTRY, BaseNodeMetric, register_metric
|
|
14
|
+
from imote.viz_tree.viz_tree import VizTree
|
|
15
|
+
|
|
16
|
+
__all__ = [
|
|
17
|
+
"VizTree",
|
|
18
|
+
"Dataset",
|
|
19
|
+
"BaseAdapter",
|
|
20
|
+
"register_adapter",
|
|
21
|
+
"ADAPTERS_REGISTRY",
|
|
22
|
+
"PilotAdapter",
|
|
23
|
+
"M5Adapter",
|
|
24
|
+
"PartyKitAdapter",
|
|
25
|
+
"BaseNodeMetric",
|
|
26
|
+
"register_metric",
|
|
27
|
+
"NODE_METRICS_REGISTRY",
|
|
28
|
+
]
|
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
def split_by_threshold(X_train, current_indices, pivot_idx, pivot_value):
|
|
4
|
+
mask = X_train[current_indices, pivot_idx] <= pivot_value
|
|
5
|
+
return _apply_mask(current_indices, mask)
|
|
6
|
+
|
|
7
|
+
def split_by_categories(X_train, current_indices, pivot_idx, left_categories):
|
|
8
|
+
mask = np.isin(X_train[current_indices, pivot_idx], left_categories)
|
|
9
|
+
return _apply_mask(current_indices, mask)
|
|
10
|
+
|
|
11
|
+
def _apply_mask(current_indices, mask):
|
|
12
|
+
left_indices, right_indices = current_indices.copy(), current_indices.copy()
|
|
13
|
+
left_indices[current_indices] = mask
|
|
14
|
+
right_indices[current_indices] = ~mask
|
|
15
|
+
return left_indices, right_indices
|
|
@@ -0,0 +1,76 @@
|
|
|
1
|
+
from abc import ABC, abstractmethod
|
|
2
|
+
import numpy as np
|
|
3
|
+
from imote.nodes.base_node import BaseNode
|
|
4
|
+
|
|
5
|
+
ADAPTERS_REGISTRY: dict[str, type["BaseAdapter"]] = {}
|
|
6
|
+
"""Maps adapter names to adapter classes, populated via register_adapter()."""
|
|
7
|
+
|
|
8
|
+
def register_adapter(name: str):
|
|
9
|
+
"""Class decorator that registers an adapter under a given name.
|
|
10
|
+
|
|
11
|
+
Args:
|
|
12
|
+
name: Name to register the adapter under.
|
|
13
|
+
|
|
14
|
+
Returns:
|
|
15
|
+
A decorator that registers the decorated class and returns it
|
|
16
|
+
unchanged.
|
|
17
|
+
"""
|
|
18
|
+
def decorator(adapter_cls):
|
|
19
|
+
ADAPTERS_REGISTRY[name] = adapter_cls
|
|
20
|
+
return adapter_cls
|
|
21
|
+
return decorator
|
|
22
|
+
|
|
23
|
+
class BaseAdapter(ABC):
|
|
24
|
+
"""Abstract base class for adapters that build a VizTree from a model.
|
|
25
|
+
|
|
26
|
+
An adapter translates a fitted model from a specific source
|
|
27
|
+
into the BaseNode linked tree structure that VizTree expects.
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
@staticmethod
|
|
31
|
+
@abstractmethod
|
|
32
|
+
def build_root_node(X_train: np.ndarray, y_train: np.ndarray, model) -> BaseNode:
|
|
33
|
+
"""Builds the BaseNode linked tree representing a fitted model.
|
|
34
|
+
|
|
35
|
+
Args:
|
|
36
|
+
X_train: Training feature matrix the model was fit on.
|
|
37
|
+
y_train: Training target values the model was fit on.
|
|
38
|
+
model: The fitted model, in whatever format this adapter
|
|
39
|
+
supports (as returned by load_model()).
|
|
40
|
+
|
|
41
|
+
Returns:
|
|
42
|
+
Root BaseNode of the reconstructed linked tree.
|
|
43
|
+
"""
|
|
44
|
+
pass
|
|
45
|
+
|
|
46
|
+
@staticmethod
|
|
47
|
+
@abstractmethod
|
|
48
|
+
def load_model(model_path: str):
|
|
49
|
+
"""Loads a fitted model from disk.
|
|
50
|
+
|
|
51
|
+
Args:
|
|
52
|
+
model_path: Path to the serialized model file.
|
|
53
|
+
|
|
54
|
+
Returns:
|
|
55
|
+
The loaded model, in the format this adapter's
|
|
56
|
+
build_root_node() expects.
|
|
57
|
+
"""
|
|
58
|
+
pass
|
|
59
|
+
|
|
60
|
+
@staticmethod
|
|
61
|
+
def predict(X: np.ndarray, model) -> np.ndarray | None:
|
|
62
|
+
"""Predicts target values directly from the underlying model, if possible.
|
|
63
|
+
|
|
64
|
+
This is optional: an adapter can override it to use the
|
|
65
|
+
original model's own prediction logic (e.g. for speed or
|
|
66
|
+
numerical parity). If not overridden, VizTree instead computes
|
|
67
|
+
predictions by traversing the built BaseNode linked tree.
|
|
68
|
+
|
|
69
|
+
Args:
|
|
70
|
+
X: Feature matrix to predict on.
|
|
71
|
+
model: The fitted model, as returned by load_model().
|
|
72
|
+
|
|
73
|
+
Returns:
|
|
74
|
+
Array of predicted values, or None if not implemented.
|
|
75
|
+
"""
|
|
76
|
+
return None
|
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
from m5py.main import ConstantLeafModel, LinRegLeafModel
|
|
3
|
+
|
|
4
|
+
from imote.adapters.adapter_utils import split_by_threshold
|
|
5
|
+
from imote.nodes.base_node import BaseNode
|
|
6
|
+
from imote.nodes.leaf_node import LeafNode
|
|
7
|
+
from imote.nodes.split_node import SplitNode
|
|
8
|
+
from imote.nodes.node_model import LinearNodeModel, NoneNodeModel, ConstantNodeModel
|
|
9
|
+
|
|
10
|
+
def build_root_node_from_m5(m5_model, X_train, y_train) -> BaseNode:
|
|
11
|
+
# builds a viz tree from a fitted m5py model (M5Base / M5Prime).
|
|
12
|
+
tree = m5_model.tree_
|
|
13
|
+
node_models = m5_model.node_models
|
|
14
|
+
root_indices = np.ones(X_train.shape[0], dtype=bool)
|
|
15
|
+
|
|
16
|
+
return _build_root_node_from_m5(tree, node_models, X_train, y_train, root_indices, node_id=0)
|
|
17
|
+
|
|
18
|
+
def _build_root_node_from_m5(tree, node_models, X_train, y_train, current_indices, node_id) -> BaseNode:
|
|
19
|
+
|
|
20
|
+
current_y_res = y_train[current_indices]
|
|
21
|
+
n_features = X_train.shape[1]
|
|
22
|
+
|
|
23
|
+
left_id = tree.children_left[node_id]
|
|
24
|
+
right_id = tree.children_right[node_id]
|
|
25
|
+
|
|
26
|
+
if left_id == -1: # sklearn's TREE_LEAF sentinel; children_left/right are both -1 at leaves
|
|
27
|
+
model = node_models[node_id]
|
|
28
|
+
|
|
29
|
+
if isinstance(model, ConstantLeafModel):
|
|
30
|
+
rss = model.error # note: this is an RMSE, not a raw RSS like PILOT's Rt - rescale/rename if needed
|
|
31
|
+
node_model = ConstantNodeModel(float(tree.value[node_id].ravel()[0]))
|
|
32
|
+
|
|
33
|
+
elif isinstance(model, LinRegLeafModel):
|
|
34
|
+
coefficients = np.zeros(n_features)
|
|
35
|
+
coefficients[model.features] = model.model.coef_
|
|
36
|
+
intercept = model.model.intercept_
|
|
37
|
+
rss = model.error # same caveat as above
|
|
38
|
+
node_model = LinearNodeModel(coefficients, intercept)
|
|
39
|
+
|
|
40
|
+
else:
|
|
41
|
+
raise ValueError(f"Unexpected leaf model type: {type(model)}")
|
|
42
|
+
|
|
43
|
+
return LeafNode(
|
|
44
|
+
indices=current_indices,
|
|
45
|
+
y_res=current_y_res,
|
|
46
|
+
rss=rss,
|
|
47
|
+
node_model=node_model,
|
|
48
|
+
)
|
|
49
|
+
|
|
50
|
+
else:
|
|
51
|
+
pivot_idx = tree.feature[node_id]
|
|
52
|
+
pivot_value = tree.threshold[node_id]
|
|
53
|
+
|
|
54
|
+
left_indices, right_indices = split_by_threshold(X_train, current_indices, pivot_idx, pivot_value)
|
|
55
|
+
|
|
56
|
+
left_child = _build_root_node_from_m5(tree, node_models, X_train, y_train, left_indices, node_id=left_id)
|
|
57
|
+
right_child = _build_root_node_from_m5(tree, node_models, X_train, y_train, right_indices, node_id=right_id)
|
|
58
|
+
|
|
59
|
+
rss = tree.impurity[node_id] * tree.n_node_samples[node_id] # sklearn stores impurity as MSE; * n -> RSS
|
|
60
|
+
|
|
61
|
+
return SplitNode(
|
|
62
|
+
indices=current_indices,
|
|
63
|
+
y_res=current_y_res,
|
|
64
|
+
rss=rss,
|
|
65
|
+
pivot_idx=pivot_idx,
|
|
66
|
+
pivot_value=pivot_value,
|
|
67
|
+
left_child=left_child,
|
|
68
|
+
right_child=right_child,
|
|
69
|
+
left_model=NoneNodeModel(),
|
|
70
|
+
right_model=NoneNodeModel(),
|
|
71
|
+
)
|
|
@@ -0,0 +1,18 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
from imote.nodes.base_node import BaseNode
|
|
3
|
+
from .build_root_node_m5 import build_root_node_from_m5
|
|
4
|
+
from .base_adapter import BaseAdapter, register_adapter
|
|
5
|
+
|
|
6
|
+
@register_adapter("M5")
|
|
7
|
+
class M5Adapter(BaseAdapter):
|
|
8
|
+
@staticmethod
|
|
9
|
+
def build_root_node(X_train, y_train, model) -> BaseNode:
|
|
10
|
+
return build_root_node_from_m5(model, X_train, y_train)
|
|
11
|
+
|
|
12
|
+
@staticmethod
|
|
13
|
+
def load_model(model_path):
|
|
14
|
+
raise NotImplementedError()
|
|
15
|
+
|
|
16
|
+
@staticmethod
|
|
17
|
+
def predict(X, m5_model) -> np.ndarray:
|
|
18
|
+
return m5_model.predict(X)
|
|
@@ -0,0 +1,133 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import json
|
|
3
|
+
import numpy as np
|
|
4
|
+
from imote.adapters.base_adapter import BaseAdapter, register_adapter
|
|
5
|
+
from imote.adapters.adapter_utils import split_by_categories, split_by_threshold
|
|
6
|
+
from imote.nodes.base_node import BaseNode
|
|
7
|
+
from imote.nodes.leaf_node import LeafNode
|
|
8
|
+
from imote.nodes.split_node import SplitNode, SplitCNode
|
|
9
|
+
from imote.nodes.node_model import LinearNodeModel, NoneNodeModel
|
|
10
|
+
|
|
11
|
+
@register_adapter('Partykit')
|
|
12
|
+
class PartyKitAdapter(BaseAdapter):
|
|
13
|
+
"""Adapter that builds a BaseNode linked tree from an R `partykit` package model export."""
|
|
14
|
+
|
|
15
|
+
@staticmethod
|
|
16
|
+
def load_model(model_path: str) -> dict:
|
|
17
|
+
"""Loads a partykit model export from a JSON file.
|
|
18
|
+
|
|
19
|
+
Args:
|
|
20
|
+
model_path: Path to the JSON file containing the exported
|
|
21
|
+
partykit model.
|
|
22
|
+
|
|
23
|
+
Returns:
|
|
24
|
+
The parsed model as a dict.
|
|
25
|
+
|
|
26
|
+
Raises:
|
|
27
|
+
FileNotFoundError: If model_path does not exist.
|
|
28
|
+
"""
|
|
29
|
+
if not os.path.exists(model_path):
|
|
30
|
+
raise FileNotFoundError(f"No such file path to load model: {model_path}") # The standard exception for a missing file
|
|
31
|
+
with open(model_path) as f:
|
|
32
|
+
model = json.load(f)
|
|
33
|
+
|
|
34
|
+
return model
|
|
35
|
+
|
|
36
|
+
@staticmethod
|
|
37
|
+
def build_root_node(X_train: np.ndarray, y_train: np.ndarray, model) -> BaseNode:
|
|
38
|
+
"""Builds the BaseNode linked tree representing a partykit model.
|
|
39
|
+
|
|
40
|
+
Args:
|
|
41
|
+
X_train: Training feature matrix the model was fit on.
|
|
42
|
+
y_train: Training target values the model was fit on.
|
|
43
|
+
model: Parsed partykit model, as returned by load_model().
|
|
44
|
+
|
|
45
|
+
Returns:
|
|
46
|
+
Root BaseNode of the reconstructed tree.
|
|
47
|
+
"""
|
|
48
|
+
root_indices = np.ones(X_train.shape[0], dtype=bool)
|
|
49
|
+
node_list = model['nodes']
|
|
50
|
+
feature_names = model['names'][1:]
|
|
51
|
+
return PartyKitAdapter._build_recursive(node_list, feature_names, node_list[0], X_train, y_train, root_indices)
|
|
52
|
+
|
|
53
|
+
@staticmethod
|
|
54
|
+
def _build_recursive(node_list, feature_names, node, X_train, y_train, current_indices) -> BaseNode:
|
|
55
|
+
"""Recursively builds a BaseNode (sub)tree from partykit node data.
|
|
56
|
+
|
|
57
|
+
Args:
|
|
58
|
+
node_list: Full flat list of partykit node dicts for the model.
|
|
59
|
+
feature_names: Feature names, in the same column order as X_train.
|
|
60
|
+
node: The partykit node dict to convert at this recursion step.
|
|
61
|
+
X_train: Training feature matrix the model was fit on.
|
|
62
|
+
y_train: Training target values the model was fit on.
|
|
63
|
+
current_indices: Boolean mask over X_train's rows selecting
|
|
64
|
+
the samples that reach this node.
|
|
65
|
+
|
|
66
|
+
Returns:
|
|
67
|
+
A BaseNode linked (sub)tree.
|
|
68
|
+
|
|
69
|
+
Raises:
|
|
70
|
+
ValueError: If node has neither a "breaks" nor an "index"
|
|
71
|
+
key, so its split type can't be determined.
|
|
72
|
+
"""
|
|
73
|
+
current_y_res = y_train[current_indices]
|
|
74
|
+
|
|
75
|
+
if node["is_terminal"]:
|
|
76
|
+
coefficients = np.array([
|
|
77
|
+
node["coefficients"].get(name, 0.0)
|
|
78
|
+
for name in feature_names
|
|
79
|
+
], dtype=float)
|
|
80
|
+
|
|
81
|
+
node_model = LinearNodeModel(coefficients, node["coefficients"]['(Intercept)'])
|
|
82
|
+
|
|
83
|
+
return LeafNode(
|
|
84
|
+
indices=current_indices,
|
|
85
|
+
y_res=current_y_res,
|
|
86
|
+
rss=-1, #TODO
|
|
87
|
+
node_model=node_model,
|
|
88
|
+
)
|
|
89
|
+
else:
|
|
90
|
+
pivot_idx = feature_names.index(node["split_var"])
|
|
91
|
+
left_node = node_list[node["kids"][0]]
|
|
92
|
+
right_node = node_list[node["kids"][1]]
|
|
93
|
+
if "breaks" in node:
|
|
94
|
+
pivot_value = node["breaks"]
|
|
95
|
+
left_indices, right_indices = split_by_threshold(X_train, current_indices, pivot_idx, pivot_value)
|
|
96
|
+
|
|
97
|
+
left_child = PartyKitAdapter._build_recursive(
|
|
98
|
+
node_list, feature_names, left_node, X_train, y_train, left_indices)
|
|
99
|
+
right_child = PartyKitAdapter._build_recursive(
|
|
100
|
+
node_list, feature_names, right_node, X_train, y_train, right_indices)
|
|
101
|
+
|
|
102
|
+
return SplitNode(
|
|
103
|
+
indices=current_indices,
|
|
104
|
+
y_res=current_y_res,
|
|
105
|
+
rss=-1, #TODO
|
|
106
|
+
pivot_idx=pivot_idx,
|
|
107
|
+
pivot_value=pivot_value,
|
|
108
|
+
left_child=left_child,
|
|
109
|
+
right_child=right_child,
|
|
110
|
+
left_model=NoneNodeModel(),
|
|
111
|
+
right_model=NoneNodeModel(),
|
|
112
|
+
)
|
|
113
|
+
elif "index" in node:
|
|
114
|
+
left_categories = [lvl for lvl, idx in zip(node["levels"], node["index"]) if idx == 1]
|
|
115
|
+
left_indices, right_indices = split_by_categories(X_train, current_indices, pivot_idx, left_categories)
|
|
116
|
+
|
|
117
|
+
left_child = PartyKitAdapter._build_recursive(
|
|
118
|
+
node_list, feature_names, left_node, X_train, y_train, left_indices)
|
|
119
|
+
right_child = PartyKitAdapter._build_recursive(
|
|
120
|
+
node_list, feature_names, right_node, X_train, y_train, right_indices)
|
|
121
|
+
return SplitCNode(
|
|
122
|
+
indices=current_indices,
|
|
123
|
+
y_res=current_y_res,
|
|
124
|
+
rss=-1, #TODO
|
|
125
|
+
pivot_idx=pivot_idx,
|
|
126
|
+
pivot_value=left_categories,
|
|
127
|
+
left_child=left_child,
|
|
128
|
+
right_child=right_child,
|
|
129
|
+
left_model=NoneNodeModel(),
|
|
130
|
+
right_model=NoneNodeModel(),
|
|
131
|
+
)
|
|
132
|
+
else:
|
|
133
|
+
raise ValueError(f"Unknown node: {node}")
|
|
@@ -0,0 +1,154 @@
|
|
|
1
|
+
from imote.nodes.base_node import BaseNode
|
|
2
|
+
from imote.nodes.leaf_node import LeafNode
|
|
3
|
+
from imote.nodes.internal_node import LinearNode
|
|
4
|
+
from imote.nodes.split_node import PconNode, BlinNode, PlinNode, PconcNode
|
|
5
|
+
from imote.nodes.node_model import LinearNodeModel, ConstantNodeModel, SimpleLinearNodeModel
|
|
6
|
+
import numpy as np
|
|
7
|
+
from .base_adapter import BaseAdapter, register_adapter
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
@register_adapter("Pilot")
|
|
11
|
+
class PilotAdapter(BaseAdapter):
|
|
12
|
+
@staticmethod
|
|
13
|
+
def build_root_node(X_train, y_train, model) -> BaseNode:
|
|
14
|
+
tree = model.tree_summary()
|
|
15
|
+
n_features = X_train.shape[1]
|
|
16
|
+
root_indices = np.ones(X_train.shape[0], dtype=bool)
|
|
17
|
+
if tree.parent_node_id[0]:
|
|
18
|
+
raise ValueError('First node is not the root node')
|
|
19
|
+
return PilotAdapter._build_recursive(
|
|
20
|
+
tree, 0, X_train, root_indices, y_train,
|
|
21
|
+
np.zeros(n_features), 0.0,
|
|
22
|
+
)
|
|
23
|
+
|
|
24
|
+
@staticmethod
|
|
25
|
+
def _build_recursive(tree, node_index, X_train, current_indices, current_y_res,
|
|
26
|
+
accumulated_coefficients, accumulated_intercept) -> BaseNode:
|
|
27
|
+
node_type = tree.node_type[node_index]
|
|
28
|
+
if node_type == 'con' or node_type == 'END':
|
|
29
|
+
leaf_intercept = accumulated_intercept + np.nan_to_num(tree.intercept_left[node_index])
|
|
30
|
+
return LeafNode(
|
|
31
|
+
indices=current_indices,
|
|
32
|
+
y_res=current_y_res,
|
|
33
|
+
rss=float(np.sum(current_y_res ** 2)),
|
|
34
|
+
node_model=LinearNodeModel(accumulated_coefficients, leaf_intercept),
|
|
35
|
+
)
|
|
36
|
+
|
|
37
|
+
children_index = np.where(tree.parent_node_id == tree.node_id[node_index])[0]
|
|
38
|
+
if node_type == 'lin':
|
|
39
|
+
pivot_idx = int(tree.feature_index[node_index])
|
|
40
|
+
coef = tree.slope_left[node_index]
|
|
41
|
+
intercept = tree.intercept_left[node_index]
|
|
42
|
+
|
|
43
|
+
new_coefficients = accumulated_coefficients.copy()
|
|
44
|
+
new_coefficients[pivot_idx] += coef
|
|
45
|
+
new_intercept = accumulated_intercept + intercept
|
|
46
|
+
|
|
47
|
+
new_y_res = current_y_res - (intercept + coef * X_train[current_indices, pivot_idx])
|
|
48
|
+
child_index = children_index[0]
|
|
49
|
+
child = PilotAdapter._build_recursive(tree, child_index, X_train, current_indices, new_y_res,
|
|
50
|
+
new_coefficients, new_intercept)
|
|
51
|
+
|
|
52
|
+
return LinearNode(
|
|
53
|
+
indices=current_indices,
|
|
54
|
+
y_res=current_y_res,
|
|
55
|
+
rss=float(np.sum(current_y_res ** 2)),
|
|
56
|
+
pivot_idx=pivot_idx,
|
|
57
|
+
linear_model=SimpleLinearNodeModel(pivot_idx, coef, intercept),
|
|
58
|
+
child=child,
|
|
59
|
+
)
|
|
60
|
+
elif node_type in ['pcon', 'pconc', 'blin', 'plin']:
|
|
61
|
+
pivot_idx = int(tree.feature_index[node_index])
|
|
62
|
+
coef_left = np.nan_to_num(tree.slope_left[node_index])
|
|
63
|
+
intercept_left = np.nan_to_num(tree.intercept_left[node_index])
|
|
64
|
+
coef_right = np.nan_to_num(tree.slope_right[node_index])
|
|
65
|
+
intercept_right = np.nan_to_num(tree.intercept_right[node_index])
|
|
66
|
+
|
|
67
|
+
if node_type == 'pconc':
|
|
68
|
+
pivot_value = tree.pivot_values[node_index]
|
|
69
|
+
left_mask = np.isin(X_train[current_indices, pivot_idx], pivot_value)
|
|
70
|
+
right_mask = ~left_mask
|
|
71
|
+
else:
|
|
72
|
+
pivot_value = tree.split_value[node_index]
|
|
73
|
+
left_mask = X_train[current_indices, pivot_idx] <= pivot_value
|
|
74
|
+
right_mask = ~left_mask
|
|
75
|
+
left_indices, right_indices = current_indices.copy(), current_indices.copy()
|
|
76
|
+
left_indices[current_indices] = left_mask
|
|
77
|
+
right_indices[current_indices] = right_mask
|
|
78
|
+
|
|
79
|
+
left_y_res = current_y_res[left_mask] - (intercept_left + coef_left * X_train[left_indices, pivot_idx])
|
|
80
|
+
left_coefficients = accumulated_coefficients.copy()
|
|
81
|
+
left_coefficients[pivot_idx] += coef_left
|
|
82
|
+
left_intercept = accumulated_intercept + intercept_left
|
|
83
|
+
|
|
84
|
+
right_y_res = current_y_res[right_mask] - (intercept_right + coef_right * X_train[right_indices, pivot_idx])
|
|
85
|
+
right_coefficients = accumulated_coefficients.copy()
|
|
86
|
+
right_coefficients[pivot_idx] += coef_right
|
|
87
|
+
right_intercept = accumulated_intercept + intercept_right
|
|
88
|
+
|
|
89
|
+
# Recurse to both children
|
|
90
|
+
left_child_index = children_index[0]
|
|
91
|
+
right_child_index = children_index[1]
|
|
92
|
+
left_child = PilotAdapter._build_recursive(tree, left_child_index, X_train, left_indices,
|
|
93
|
+
left_y_res, left_coefficients, left_intercept)
|
|
94
|
+
right_child = PilotAdapter._build_recursive(tree, right_child_index, X_train, right_indices,
|
|
95
|
+
right_y_res, right_coefficients, right_intercept)
|
|
96
|
+
|
|
97
|
+
if node_type == 'pcon':
|
|
98
|
+
return PconNode(
|
|
99
|
+
indices=current_indices,
|
|
100
|
+
y_res=current_y_res,
|
|
101
|
+
rss=float(np.sum(current_y_res ** 2)),
|
|
102
|
+
pivot_idx=pivot_idx,
|
|
103
|
+
pivot_value=pivot_value,
|
|
104
|
+
left_child=left_child,
|
|
105
|
+
right_child=right_child,
|
|
106
|
+
left_model=ConstantNodeModel(intercept_left),
|
|
107
|
+
right_model=ConstantNodeModel(intercept_right),
|
|
108
|
+
)
|
|
109
|
+
elif node_type == 'pconc':
|
|
110
|
+
return PconcNode(
|
|
111
|
+
indices=current_indices,
|
|
112
|
+
y_res=current_y_res,
|
|
113
|
+
rss=float(np.sum(current_y_res ** 2)),
|
|
114
|
+
pivot_idx=pivot_idx,
|
|
115
|
+
pivot_value=pivot_value,
|
|
116
|
+
left_child=left_child,
|
|
117
|
+
right_child=right_child,
|
|
118
|
+
left_model=ConstantNodeModel(intercept_left),
|
|
119
|
+
right_model=ConstantNodeModel(intercept_right),
|
|
120
|
+
)
|
|
121
|
+
elif node_type == 'blin':
|
|
122
|
+
return BlinNode(
|
|
123
|
+
indices=current_indices,
|
|
124
|
+
y_res=current_y_res,
|
|
125
|
+
rss=float(np.sum(current_y_res ** 2)),
|
|
126
|
+
pivot_idx=pivot_idx,
|
|
127
|
+
pivot_value=pivot_value,
|
|
128
|
+
left_child=left_child,
|
|
129
|
+
right_child=right_child,
|
|
130
|
+
left_model=SimpleLinearNodeModel(pivot_idx, coef_left, intercept_left),
|
|
131
|
+
right_model=SimpleLinearNodeModel(pivot_idx, coef_right, intercept_right),
|
|
132
|
+
)
|
|
133
|
+
else: # node_type == 'plin'
|
|
134
|
+
return PlinNode(
|
|
135
|
+
indices=current_indices,
|
|
136
|
+
y_res=current_y_res,
|
|
137
|
+
rss=float(np.sum(current_y_res ** 2)),
|
|
138
|
+
pivot_idx=pivot_idx,
|
|
139
|
+
pivot_value=pivot_value,
|
|
140
|
+
left_child=left_child,
|
|
141
|
+
right_child=right_child,
|
|
142
|
+
left_model=SimpleLinearNodeModel(pivot_idx, coef_left, intercept_left),
|
|
143
|
+
right_model=SimpleLinearNodeModel(pivot_idx, coef_right, intercept_right),
|
|
144
|
+
)
|
|
145
|
+
else:
|
|
146
|
+
raise ValueError(f"Unknown node type: {node_type}")
|
|
147
|
+
|
|
148
|
+
@staticmethod
|
|
149
|
+
def load_model(model_path):
|
|
150
|
+
raise NotImplementedError()
|
|
151
|
+
|
|
152
|
+
@staticmethod
|
|
153
|
+
def predict(X, model) -> np.ndarray:
|
|
154
|
+
return model.predict(X)
|
imote/app.py
ADDED
|
@@ -0,0 +1,80 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
|
|
3
|
+
import matplotlib
|
|
4
|
+
matplotlib.use('agg')
|
|
5
|
+
|
|
6
|
+
import dash
|
|
7
|
+
import dash_bootstrap_components as dbc
|
|
8
|
+
import dash_cytoscape as cyto
|
|
9
|
+
from dash import Dash, html
|
|
10
|
+
from flask import send_from_directory
|
|
11
|
+
|
|
12
|
+
from imote.callbacks import register_all_callbacks
|
|
13
|
+
from imote.components.stores import make_global_stores
|
|
14
|
+
from imote.config import DIR_LIVE_OUTPUT, DIR_SAVED_VIZ_TREES
|
|
15
|
+
|
|
16
|
+
cyto.load_extra_layouts()
|
|
17
|
+
|
|
18
|
+
for plot_dir in (DIR_LIVE_OUTPUT / "regplots", DIR_LIVE_OUTPUT / "predsplots"):
|
|
19
|
+
plot_dir.mkdir(parents=True, exist_ok=True)
|
|
20
|
+
for old_plot in plot_dir.glob("*.svg"):
|
|
21
|
+
old_plot.unlink(missing_ok=True)
|
|
22
|
+
DIR_SAVED_VIZ_TREES.mkdir(parents=True, exist_ok=True)
|
|
23
|
+
|
|
24
|
+
app: Dash = dash.Dash(
|
|
25
|
+
__name__,
|
|
26
|
+
use_pages=True,
|
|
27
|
+
external_stylesheets=[dbc.themes.FLATLY, dbc.icons.BOOTSTRAP],
|
|
28
|
+
suppress_callback_exceptions=True,
|
|
29
|
+
)
|
|
30
|
+
app.title = "IMoTE: Interactive MOdel Tree Explorer"
|
|
31
|
+
|
|
32
|
+
@app.server.route("/internal_regplots/<path:filename>")
|
|
33
|
+
def serve_regplots(filename):
|
|
34
|
+
return send_from_directory(str(DIR_LIVE_OUTPUT / "regplots"), filename)
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
@app.server.route("/internal_predsplots/<path:filename>")
|
|
38
|
+
def serve_predsplots(filename):
|
|
39
|
+
return send_from_directory(str(DIR_LIVE_OUTPUT / "predsplots"), filename)
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
navbar = dbc.NavbarSimple(
|
|
43
|
+
children=[
|
|
44
|
+
dbc.NavLink(page["name"], href=page["path"], active="exact")
|
|
45
|
+
for page in dash.page_registry.values()
|
|
46
|
+
],
|
|
47
|
+
brand="IMoTE: Interactive MOdel Tree Explorer",
|
|
48
|
+
color="dark",
|
|
49
|
+
dark=True,
|
|
50
|
+
)
|
|
51
|
+
|
|
52
|
+
# Stores live OUTSIDE dash.page_container so their data survives navigating
|
|
53
|
+
# between the Tree View and Explain pages.
|
|
54
|
+
app.layout = html.Div(
|
|
55
|
+
[
|
|
56
|
+
navbar,
|
|
57
|
+
make_global_stores(),
|
|
58
|
+
dbc.Container(dash.page_container, fluid=True, class_name="pt-3"),
|
|
59
|
+
]
|
|
60
|
+
)
|
|
61
|
+
|
|
62
|
+
register_all_callbacks(app)
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def _setup_logging(level: int):
|
|
66
|
+
handler = logging.StreamHandler()
|
|
67
|
+
handler.setFormatter(logging.Formatter("%(levelname)s %(name)s: %(message)s"))
|
|
68
|
+
logger = logging.getLogger("imote")
|
|
69
|
+
logger.addHandler(handler)
|
|
70
|
+
logger.setLevel(level)
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def main():
|
|
74
|
+
_setup_logging(logging.INFO)
|
|
75
|
+
app.run()
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
if __name__ == "__main__":
|
|
79
|
+
_setup_logging(logging.DEBUG)
|
|
80
|
+
app.run(debug=True)
|
imote/benchmark_info.py
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
PMLB_DATASETS_CAT_IDS = {
|
|
4
|
+
'556_analcatdata_apnea2': np.array([0,1]),
|
|
5
|
+
'557_analcatdata_apnea1': np.array([0,1]),
|
|
6
|
+
'522_pm10': np.array([-1]),
|
|
7
|
+
'1028_SWD': np.array([0,1,2,3,4,5,6,7,8,9]),
|
|
8
|
+
'485_analcatdata_vehicle': np.array([0,1,2,3]),
|
|
9
|
+
'547_no2': np.array([-1]),
|
|
10
|
+
'665_sleuth_case2002': np.array([0,1,2,3]),
|
|
11
|
+
'210_cloud': np.array([0,1]),
|
|
12
|
+
'229_pwLinear': np.array([0,1,2,3,4,5,6,7,8,9]),
|
|
13
|
+
'230_machine_cpu': np.array([-1]),
|
|
14
|
+
'656_fri_c1_100_5': np.array([-1]),
|
|
15
|
+
'192_vineyard': np.array([-1]),
|
|
16
|
+
'653_fri_c0_250_25': np.array([-1]),
|
|
17
|
+
'687_sleuth_ex1605': np.array([-1]),
|
|
18
|
+
'651_fri_c0_100_25': np.array([-1]),
|
|
19
|
+
'658_fri_c3_250_25': np.array([-1]),
|
|
20
|
+
'294_satellite_image': np.array([-1]),
|
|
21
|
+
'1199_BNG_echoMonths': np.array([0,2,8]),
|
|
22
|
+
'505_tecator': np.array([-1]),
|
|
23
|
+
'560_bodyfat': np.array([-1]),
|
|
24
|
+
'197_cpu_act': np.array([-1]),
|
|
25
|
+
'225_puma8NH': np.array([-1]),
|
|
26
|
+
'503_wind': np.array([-1]),
|
|
27
|
+
'4544_GeographicalOriginalofMusic': np.array([-1]),
|
|
28
|
+
'537_houses': np.array([-1]),
|
|
29
|
+
}
|
|
@@ -0,0 +1,16 @@
|
|
|
1
|
+
from imote.callbacks.a_tree_info_callbacks import register_callbacks as register_tree_info_callbacks
|
|
2
|
+
from imote.callbacks.b_node_callbacks import register_callbacks as register_node_callbacks
|
|
3
|
+
from imote.callbacks.c_new_tree_callbacks import register_callbacks as register_new_tree_callbacks
|
|
4
|
+
from imote.callbacks.d_edit_tree_callbacks import register_callbacks as register_edit_tree_callbacks
|
|
5
|
+
from imote.callbacks.e_layout_callbacks import register_callbacks as register_layout_callbacks
|
|
6
|
+
from imote.callbacks.f_highlight_callbacks import register_callbacks as register_highlight_callbacks
|
|
7
|
+
from imote.callbacks.g_elements_callbacks import register_callbacks as register_elements_callbacks
|
|
8
|
+
|
|
9
|
+
def register_all_callbacks(app):
|
|
10
|
+
register_tree_info_callbacks(app)
|
|
11
|
+
register_node_callbacks(app)
|
|
12
|
+
register_new_tree_callbacks(app)
|
|
13
|
+
register_edit_tree_callbacks(app)
|
|
14
|
+
register_layout_callbacks(app)
|
|
15
|
+
register_highlight_callbacks(app)
|
|
16
|
+
register_elements_callbacks(app)
|