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.
Files changed (63) hide show
  1. imote/__init__.py +28 -0
  2. imote/adapters/__init__.py +3 -0
  3. imote/adapters/adapter_utils.py +15 -0
  4. imote/adapters/base_adapter.py +76 -0
  5. imote/adapters/build_root_node_m5.py +71 -0
  6. imote/adapters/m5_adapter.py +18 -0
  7. imote/adapters/partykit_adapter.py +133 -0
  8. imote/adapters/pilot_adapter.py +154 -0
  9. imote/app.py +80 -0
  10. imote/benchmark_info.py +29 -0
  11. imote/callbacks/__init__.py +16 -0
  12. imote/callbacks/a_tree_info_callbacks.py +110 -0
  13. imote/callbacks/b_node_callbacks.py +186 -0
  14. imote/callbacks/c_new_tree_callbacks.py +380 -0
  15. imote/callbacks/d_edit_tree_callbacks.py +192 -0
  16. imote/callbacks/e_layout_callbacks.py +39 -0
  17. imote/callbacks/f_highlight_callbacks.py +79 -0
  18. imote/callbacks/g_elements_callbacks.py +195 -0
  19. imote/components/__init__.py +0 -0
  20. imote/components/cards/__init__.py +0 -0
  21. imote/components/cards/a_tree_info_card.py +14 -0
  22. imote/components/cards/b_node_info_card.py +79 -0
  23. imote/components/cards/bb_node_plot_settings_card.py +139 -0
  24. imote/components/cards/c_new_tree_card.py +236 -0
  25. imote/components/cards/d_edit_tree_card.py +159 -0
  26. imote/components/cards/e_layout_card.py +66 -0
  27. imote/components/cards/f_highlight_card.py +48 -0
  28. imote/components/cytoscape_graph.py +39 -0
  29. imote/components/stores.py +20 -0
  30. imote/config.py +252 -0
  31. imote/data/initial_tree.pkl +0 -0
  32. imote/dataset/__init__.py +0 -0
  33. imote/dataset/dataset.py +170 -0
  34. imote/dataset/dataset_registry.py +23 -0
  35. imote/ids.py +124 -0
  36. imote/node_metrics/__init__.py +0 -0
  37. imote/node_metrics/node_metric.py +167 -0
  38. imote/nodes/__init__.py +0 -0
  39. imote/nodes/base_node.py +139 -0
  40. imote/nodes/collapsed_node.py +46 -0
  41. imote/nodes/combined_lin_node.py +78 -0
  42. imote/nodes/internal_node.py +129 -0
  43. imote/nodes/leaf_node.py +74 -0
  44. imote/nodes/node_model.py +342 -0
  45. imote/nodes/none_node.py +56 -0
  46. imote/nodes/split_node.py +219 -0
  47. imote/pages/__init__.py +0 -0
  48. imote/pages/explain_page.py +40 -0
  49. imote/pages/tree_view_page.py +33 -0
  50. imote/plots/__init__.py +0 -0
  51. imote/plots/predsplot.py +485 -0
  52. imote/plots/predsplot2.py +540 -0
  53. imote/plots/regplot.py +81 -0
  54. imote/viz_tree/__init__.py +0 -0
  55. imote/viz_tree/contributions.py +113 -0
  56. imote/viz_tree/viz_tree.py +432 -0
  57. imote/viz_tree/viz_tree_cytoscape.py +293 -0
  58. imote-0.0.0.dist-info/METADATA +113 -0
  59. imote-0.0.0.dist-info/RECORD +63 -0
  60. imote-0.0.0.dist-info/WHEEL +5 -0
  61. imote-0.0.0.dist-info/entry_points.txt +2 -0
  62. imote-0.0.0.dist-info/licenses/LICENSE +21 -0
  63. 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,3 @@
1
+ from .pilot_adapter import PilotAdapter
2
+ from .m5_adapter import M5Adapter
3
+ from .partykit_adapter import PartyKitAdapter
@@ -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)
@@ -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)