graph-explain 0.7.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.
- graph_explain/__init__.py +79 -0
- graph_explain/backends/__init__.py +4 -0
- graph_explain/backends/base.py +103 -0
- graph_explain/backends/dgl.py +121 -0
- graph_explain/benchmarks/__init__.py +3 -0
- graph_explain/benchmarks/synthetic.py +246 -0
- graph_explain/cli.py +459 -0
- graph_explain/core/__init__.py +14 -0
- graph_explain/core/benchmark.py +284 -0
- graph_explain/core/evaluation.py +391 -0
- graph_explain/core/explainer.py +83 -0
- graph_explain/core/explanation.py +72 -0
- graph_explain/core/model_utils.py +44 -0
- graph_explain/core/registry.py +55 -0
- graph_explain/methods/__init__.py +39 -0
- graph_explain/methods/attention/attention.py +147 -0
- graph_explain/methods/base.py +25 -0
- graph_explain/methods/baseline/random_baseline.py +78 -0
- graph_explain/methods/counterfactual/counterfactual.py +304 -0
- graph_explain/methods/feature/graph_lime.py +141 -0
- graph_explain/methods/gradient/__init__.py +0 -0
- graph_explain/methods/gradient/grad_x_input.py +110 -0
- graph_explain/methods/gradient/guided_backprop.py +117 -0
- graph_explain/methods/gradient/integrated_gradients.py +115 -0
- graph_explain/methods/gradient/saliency.py +93 -0
- graph_explain/methods/perturbation/__init__.py +0 -0
- graph_explain/methods/perturbation/gnn_explainer.py +265 -0
- graph_explain/methods/perturbation/node_mask.py +136 -0
- graph_explain/methods/perturbation/pg_explainer.py +162 -0
- graph_explain/methods/perturbation/subgraphx.py +393 -0
- graph_explain/methods/relevance/deeplift.py +262 -0
- graph_explain/methods/relevance/gnn_lrp.py +219 -0
- graph_explain/narration/__init__.py +3 -0
- graph_explain/narration/narrator.py +185 -0
- graph_explain/visualization/__init__.py +4 -0
- graph_explain/visualization/interactive.py +73 -0
- graph_explain/visualization/static.py +90 -0
- graph_explain-0.7.0.dist-info/METADATA +332 -0
- graph_explain-0.7.0.dist-info/RECORD +42 -0
- graph_explain-0.7.0.dist-info/WHEEL +5 -0
- graph_explain-0.7.0.dist-info/entry_points.txt +2 -0
- graph_explain-0.7.0.dist-info/top_level.txt +1 -0
graph_explain/cli.py
ADDED
|
@@ -0,0 +1,459 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import argparse
|
|
4
|
+
import json
|
|
5
|
+
import sys
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
from graph_explain import __version__
|
|
9
|
+
from graph_explain.core.registry import get_algorithm, instantiate
|
|
10
|
+
|
|
11
|
+
_METHODS = [
|
|
12
|
+
"gnn_explainer",
|
|
13
|
+
"gnnexplainer",
|
|
14
|
+
"saliency",
|
|
15
|
+
"gradient",
|
|
16
|
+
"grad",
|
|
17
|
+
"pg_explainer",
|
|
18
|
+
"pgexplainer",
|
|
19
|
+
"subgraphx",
|
|
20
|
+
"subgraph_x",
|
|
21
|
+
"integrated_gradients",
|
|
22
|
+
"ig",
|
|
23
|
+
"gnn_lrp",
|
|
24
|
+
"gnn-lrp",
|
|
25
|
+
"lrp",
|
|
26
|
+
"deep_lift",
|
|
27
|
+
"deeplift",
|
|
28
|
+
"dl",
|
|
29
|
+
"attention",
|
|
30
|
+
"gat",
|
|
31
|
+
"attention_explainer",
|
|
32
|
+
"grad_x_input",
|
|
33
|
+
"gradient_x_input",
|
|
34
|
+
"gx",
|
|
35
|
+
"graph_lime",
|
|
36
|
+
"glime",
|
|
37
|
+
"gl",
|
|
38
|
+
"node_mask",
|
|
39
|
+
"nodemask",
|
|
40
|
+
"nm",
|
|
41
|
+
"guided_backprop",
|
|
42
|
+
"guided-backprop",
|
|
43
|
+
"gbp",
|
|
44
|
+
"random",
|
|
45
|
+
"random_baseline",
|
|
46
|
+
"rand",
|
|
47
|
+
"counterfactual",
|
|
48
|
+
"cf",
|
|
49
|
+
]
|
|
50
|
+
|
|
51
|
+
_METRICS = ["fidelity", "fidelity_plus", "fidelity_minus", "gea", "stability"]
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def build_parser() -> argparse.ArgumentParser:
|
|
55
|
+
parser = argparse.ArgumentParser(
|
|
56
|
+
prog="graph-explain",
|
|
57
|
+
description="Explainability for graph-based models (GNNs).",
|
|
58
|
+
)
|
|
59
|
+
parser.add_argument(
|
|
60
|
+
"--version", action="version", version=f"graph-explain {__version__}"
|
|
61
|
+
)
|
|
62
|
+
sub = parser.add_subparsers(dest="command")
|
|
63
|
+
|
|
64
|
+
explain = sub.add_parser("explain", help="Explain the prediction of a node/graph")
|
|
65
|
+
explain.add_argument("--model", required=True, help="Path to the saved model (.pt)")
|
|
66
|
+
explain.add_argument("--data", required=True, help="Path to the saved Data (.pt)")
|
|
67
|
+
explain.add_argument("--method", default="gnn_explainer", choices=_METHODS)
|
|
68
|
+
explain.add_argument(
|
|
69
|
+
"--node",
|
|
70
|
+
type=int,
|
|
71
|
+
default=None,
|
|
72
|
+
help="Node index to explain (node-level)",
|
|
73
|
+
)
|
|
74
|
+
explain.add_argument("--target-class", type=int, default=None)
|
|
75
|
+
explain.add_argument("--epochs", type=int, default=200)
|
|
76
|
+
explain.add_argument("--lr", type=float, default=None)
|
|
77
|
+
explain.add_argument(
|
|
78
|
+
"--mode",
|
|
79
|
+
default="edge",
|
|
80
|
+
choices=["edge", "feature"],
|
|
81
|
+
help="Counterfactual mode",
|
|
82
|
+
)
|
|
83
|
+
explain.add_argument("--hops", type=int, default=2)
|
|
84
|
+
explain.add_argument("--max-steps", type=int, default=10)
|
|
85
|
+
explain.add_argument("--eps", type=float, default=None)
|
|
86
|
+
explain.add_argument("--steps", type=int, default=50)
|
|
87
|
+
explain.add_argument("--normalize", action="store_true", help="Normalized GNN-LRP")
|
|
88
|
+
explain.add_argument("--backend", default="pyg", choices=["pyg", "dgl"])
|
|
89
|
+
explain.add_argument("--output", default=None, help="Save the explanation to .pt")
|
|
90
|
+
explain.add_argument("--plot", default=None, help="Save visualization to .png/.pdf")
|
|
91
|
+
explain.add_argument(
|
|
92
|
+
"--html",
|
|
93
|
+
default=None,
|
|
94
|
+
help="Save interactive visualization to .html",
|
|
95
|
+
)
|
|
96
|
+
explain.add_argument("--threshold", type=float, default=0.5)
|
|
97
|
+
explain.add_argument(
|
|
98
|
+
"--top-k", type=int, default=5, help="Top-k for GEA/stability/JSON"
|
|
99
|
+
)
|
|
100
|
+
explain.add_argument(
|
|
101
|
+
"--metrics",
|
|
102
|
+
default="",
|
|
103
|
+
help=f"Comma-separated list: {', '.join(_METRICS)}",
|
|
104
|
+
)
|
|
105
|
+
explain.add_argument("--num-perturbations", type=int, default=10)
|
|
106
|
+
explain.add_argument("--noise-std", type=float, default=0.05)
|
|
107
|
+
explain.add_argument("--describe", action="store_true", help="Print the narration")
|
|
108
|
+
explain.add_argument(
|
|
109
|
+
"--json", default=None, help="Export summary + metrics to .json"
|
|
110
|
+
)
|
|
111
|
+
|
|
112
|
+
bench = sub.add_parser("bench", help="Comparative benchmark of methods over a node")
|
|
113
|
+
bench.add_argument("--model", required=True, help="Path to the saved model (.pt)")
|
|
114
|
+
bench.add_argument("--data", required=True, help="Path to the saved Data (.pt)")
|
|
115
|
+
bench.add_argument(
|
|
116
|
+
"--node",
|
|
117
|
+
type=int,
|
|
118
|
+
default=None,
|
|
119
|
+
help="Node index (node-level); omit for graph-level",
|
|
120
|
+
)
|
|
121
|
+
bench.add_argument("--target-class", type=int, default=None)
|
|
122
|
+
bench.add_argument(
|
|
123
|
+
"--methods",
|
|
124
|
+
default="all",
|
|
125
|
+
help=f"Comma-separated methods (or 'all'). Aliases: {', '.join(_METHODS)}",
|
|
126
|
+
)
|
|
127
|
+
bench.add_argument("--backend", default="pyg", choices=["pyg", "dgl"])
|
|
128
|
+
bench.add_argument("--epochs", type=int, default=200)
|
|
129
|
+
bench.add_argument("--lr", type=float, default=None)
|
|
130
|
+
bench.add_argument("--top-k", type=int, default=5)
|
|
131
|
+
bench.add_argument("--num-perturbations", type=int, default=5)
|
|
132
|
+
bench.add_argument("--noise-std", type=float, default=0.05)
|
|
133
|
+
bench.add_argument("--threshold", type=float, default=0.5)
|
|
134
|
+
bench.add_argument("--seed", type=int, default=0)
|
|
135
|
+
bench.add_argument("--no-stability", action="store_true")
|
|
136
|
+
bench.add_argument("--json", default=None, help="Export results to .json")
|
|
137
|
+
bench.add_argument(
|
|
138
|
+
"--html", default=None, help="Export comparative report to .html"
|
|
139
|
+
)
|
|
140
|
+
return parser
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
def _instantiate(name: str, args: argparse.Namespace):
|
|
144
|
+
kw: dict[str, Any] = {}
|
|
145
|
+
for attr, param in (
|
|
146
|
+
("epochs", "epochs"),
|
|
147
|
+
("lr", "lr"),
|
|
148
|
+
("mode", "mode"),
|
|
149
|
+
("hops", "hops"),
|
|
150
|
+
("max_steps", "max_steps"),
|
|
151
|
+
("eps", "eps"),
|
|
152
|
+
("steps", "steps"),
|
|
153
|
+
):
|
|
154
|
+
val = getattr(args, attr)
|
|
155
|
+
if val is not None:
|
|
156
|
+
kw[param] = val
|
|
157
|
+
if "normalize" in dir(args) and args.normalize:
|
|
158
|
+
kw["normalize"] = True
|
|
159
|
+
return instantiate(name, **kw)
|
|
160
|
+
|
|
161
|
+
|
|
162
|
+
def _make_explainer(args: argparse.Namespace):
|
|
163
|
+
from graph_explain import Explainer
|
|
164
|
+
|
|
165
|
+
algorithm = _instantiate(args.method, args)
|
|
166
|
+
return Explainer(
|
|
167
|
+
algorithm=algorithm,
|
|
168
|
+
backend=args.backend,
|
|
169
|
+
mask_threshold=args.threshold,
|
|
170
|
+
)
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
def _fmt(value: Any) -> Any:
|
|
174
|
+
if hasattr(value, "tolist"):
|
|
175
|
+
return [round(float(v), 4) for v in value.reshape(-1).tolist()]
|
|
176
|
+
return value
|
|
177
|
+
|
|
178
|
+
|
|
179
|
+
def _json_safe(value: Any) -> Any:
|
|
180
|
+
if isinstance(value, dict):
|
|
181
|
+
return {str(k): _json_safe(v) for k, v in value.items()}
|
|
182
|
+
if isinstance(value, (list, tuple)):
|
|
183
|
+
return [_json_safe(v) for v in value]
|
|
184
|
+
import torch
|
|
185
|
+
|
|
186
|
+
if torch.is_tensor(value):
|
|
187
|
+
return _json_safe(value.tolist())
|
|
188
|
+
if isinstance(value, (int, float, str, bool)) or value is None:
|
|
189
|
+
return value
|
|
190
|
+
return str(value)
|
|
191
|
+
|
|
192
|
+
|
|
193
|
+
def _eval_metric(name: str, args, model, data, explanation) -> float | None:
|
|
194
|
+
from graph_explain.core.evaluation import (
|
|
195
|
+
evaluate_fidelity_minus,
|
|
196
|
+
evaluate_fidelity_plus,
|
|
197
|
+
evaluate_gea,
|
|
198
|
+
evaluate_stability,
|
|
199
|
+
)
|
|
200
|
+
|
|
201
|
+
try:
|
|
202
|
+
if name in ("fidelity", "fidelity_plus"):
|
|
203
|
+
return float(evaluate_fidelity_plus(model, explanation))
|
|
204
|
+
if name == "fidelity_minus":
|
|
205
|
+
return float(evaluate_fidelity_minus(model, explanation))
|
|
206
|
+
if name == "gea":
|
|
207
|
+
if getattr(model, "task_level", "node") == "graph":
|
|
208
|
+
from graph_explain.core.evaluation import evaluate_gea_graph
|
|
209
|
+
|
|
210
|
+
return float(
|
|
211
|
+
evaluate_gea_graph(explanation, data=data, top_k=args.top_k)
|
|
212
|
+
)
|
|
213
|
+
return float(evaluate_gea(explanation, data=data, top_k=args.top_k))
|
|
214
|
+
if name == "stability":
|
|
215
|
+
if args.node is None:
|
|
216
|
+
raise ValueError("stability requires --node")
|
|
217
|
+
|
|
218
|
+
def _again(d):
|
|
219
|
+
return _make_explainer(args).explain_node(d, model, args.node)
|
|
220
|
+
|
|
221
|
+
return float(
|
|
222
|
+
evaluate_stability(
|
|
223
|
+
_again,
|
|
224
|
+
data,
|
|
225
|
+
num_perturbations=args.num_perturbations,
|
|
226
|
+
noise_std=args.noise_std,
|
|
227
|
+
top_k=args.top_k,
|
|
228
|
+
)
|
|
229
|
+
)
|
|
230
|
+
except Exception as exc: # noqa: BLE001
|
|
231
|
+
print(f" * metric {name} unavailable: {exc}", file=sys.stderr)
|
|
232
|
+
return None
|
|
233
|
+
raise ValueError(f"unknown metric: {name}")
|
|
234
|
+
|
|
235
|
+
|
|
236
|
+
def _cmd_explain(args: argparse.Namespace) -> int:
|
|
237
|
+
import torch
|
|
238
|
+
|
|
239
|
+
from graph_explain.core.registry import get_algorithm
|
|
240
|
+
|
|
241
|
+
model = torch.load(args.model, map_location="cpu", weights_only=False)
|
|
242
|
+
data = torch.load(args.data, map_location="cpu", weights_only=False)
|
|
243
|
+
model.eval()
|
|
244
|
+
|
|
245
|
+
task = getattr(model, "task_level", "node")
|
|
246
|
+
if args.node is None:
|
|
247
|
+
if task == "graph":
|
|
248
|
+
if not get_algorithm(args.method).graph_level:
|
|
249
|
+
print(
|
|
250
|
+
f"Error: {args.method} does not support graph-level explanations "
|
|
251
|
+
"(node-level only).",
|
|
252
|
+
file=sys.stderr,
|
|
253
|
+
)
|
|
254
|
+
return 2
|
|
255
|
+
index = None
|
|
256
|
+
else:
|
|
257
|
+
print(
|
|
258
|
+
"To explain a node you must pass --node (per-node explanation).",
|
|
259
|
+
file=sys.stderr,
|
|
260
|
+
)
|
|
261
|
+
return 2
|
|
262
|
+
else:
|
|
263
|
+
index = args.node
|
|
264
|
+
|
|
265
|
+
algorithm = _instantiate(args.method, args)
|
|
266
|
+
explainer = _make_explainer(args)
|
|
267
|
+
try:
|
|
268
|
+
explanation = explainer.explain(
|
|
269
|
+
data,
|
|
270
|
+
model,
|
|
271
|
+
index=index,
|
|
272
|
+
target_class=args.target_class,
|
|
273
|
+
)
|
|
274
|
+
except ValueError as exc:
|
|
275
|
+
print(f"Error with {args.method}: {exc}", file=sys.stderr)
|
|
276
|
+
return 2
|
|
277
|
+
algorithm_class = get_algorithm(args.method)
|
|
278
|
+
|
|
279
|
+
print(f"Method: {args.method} ({algorithm_class.__name__})")
|
|
280
|
+
print(f"Original prediction: {_fmt(explanation.prediction_original)}")
|
|
281
|
+
if explanation.prediction_explanation is not None:
|
|
282
|
+
print(
|
|
283
|
+
f"Prediction after explanation: {_fmt(explanation.prediction_explanation)}"
|
|
284
|
+
)
|
|
285
|
+
|
|
286
|
+
metrics: dict[str, float | None] = {}
|
|
287
|
+
if args.metrics:
|
|
288
|
+
for name in args.metrics.split(","):
|
|
289
|
+
name = name.strip()
|
|
290
|
+
if not name:
|
|
291
|
+
continue
|
|
292
|
+
metrics[name] = _eval_metric(name, args, model, data, explanation)
|
|
293
|
+
print(f"Metrics: {metrics}")
|
|
294
|
+
|
|
295
|
+
if args.describe:
|
|
296
|
+
from graph_explain import describe
|
|
297
|
+
|
|
298
|
+
print(f"Narration: {describe(explanation, data=data, top_k=args.top_k)}")
|
|
299
|
+
|
|
300
|
+
if args.json:
|
|
301
|
+
from graph_explain import summarize
|
|
302
|
+
|
|
303
|
+
report = {
|
|
304
|
+
"version": __version__,
|
|
305
|
+
"method": algorithm.name,
|
|
306
|
+
"backend": args.backend,
|
|
307
|
+
"node": args.node,
|
|
308
|
+
"target_class": args.target_class,
|
|
309
|
+
"threshold": args.threshold,
|
|
310
|
+
"prediction_original": _fmt(explanation.prediction_original),
|
|
311
|
+
"prediction_explanation": _fmt(explanation.prediction_explanation),
|
|
312
|
+
"metrics": metrics,
|
|
313
|
+
"summary": summarize(explanation, data=data, top_k=args.top_k),
|
|
314
|
+
}
|
|
315
|
+
with open(args.json, "w", encoding="utf-8") as fh:
|
|
316
|
+
json.dump(_json_safe(report), fh, indent=2, ensure_ascii=False)
|
|
317
|
+
print(f"JSON report saved to {args.json}")
|
|
318
|
+
|
|
319
|
+
if args.output:
|
|
320
|
+
torch.save(explanation, args.output)
|
|
321
|
+
print(f"Explanation saved to {args.output}")
|
|
322
|
+
if args.plot:
|
|
323
|
+
from graph_explain.visualization import visualize_static
|
|
324
|
+
|
|
325
|
+
visualize_static(explanation, threshold=args.threshold)
|
|
326
|
+
import matplotlib.pyplot as plt
|
|
327
|
+
|
|
328
|
+
plt.savefig(args.plot, bbox_inches="tight")
|
|
329
|
+
print(f"Visualization saved to {args.plot}")
|
|
330
|
+
if args.html:
|
|
331
|
+
from graph_explain.visualization import visualize_interactive
|
|
332
|
+
|
|
333
|
+
visualize_interactive(
|
|
334
|
+
explanation, output_path=args.html, threshold=args.threshold
|
|
335
|
+
)
|
|
336
|
+
print(f"Interactive visualization saved to {args.html}")
|
|
337
|
+
return 0
|
|
338
|
+
|
|
339
|
+
|
|
340
|
+
def _cmd_bench(args: argparse.Namespace) -> int:
|
|
341
|
+
import torch
|
|
342
|
+
|
|
343
|
+
from graph_explain.core.benchmark import DEFAULT_METHODS, compare, report_html
|
|
344
|
+
|
|
345
|
+
model = torch.load(args.model, map_location="cpu", weights_only=False)
|
|
346
|
+
data = torch.load(args.data, map_location="cpu", weights_only=False)
|
|
347
|
+
model.eval()
|
|
348
|
+
|
|
349
|
+
task = getattr(model, "task_level", "node")
|
|
350
|
+
if args.node is None and task != "graph":
|
|
351
|
+
print("For node-level you must pass --node.", file=sys.stderr)
|
|
352
|
+
return 2
|
|
353
|
+
|
|
354
|
+
if args.methods.strip().lower() == "all":
|
|
355
|
+
methods = list(DEFAULT_METHODS)
|
|
356
|
+
else:
|
|
357
|
+
methods = []
|
|
358
|
+
for item in args.methods.split(","):
|
|
359
|
+
item = item.strip()
|
|
360
|
+
if not item:
|
|
361
|
+
continue
|
|
362
|
+
try:
|
|
363
|
+
methods.append(get_algorithm(item).name)
|
|
364
|
+
except ValueError as exc:
|
|
365
|
+
print(f"Error: {exc}", file=sys.stderr)
|
|
366
|
+
return 2
|
|
367
|
+
|
|
368
|
+
methods = list(dict.fromkeys(methods))
|
|
369
|
+
|
|
370
|
+
print(
|
|
371
|
+
f"Benchmark {'over node ' + str(args.node) if args.node is not None else 'graph-level'}"
|
|
372
|
+
f" ({len(methods)} methods) - backend {args.backend}\n"
|
|
373
|
+
)
|
|
374
|
+
results = compare(
|
|
375
|
+
data,
|
|
376
|
+
model,
|
|
377
|
+
node=args.node,
|
|
378
|
+
target_class=args.target_class,
|
|
379
|
+
backend=args.backend,
|
|
380
|
+
methods=methods,
|
|
381
|
+
top_k=args.top_k,
|
|
382
|
+
num_perturbations=args.num_perturbations,
|
|
383
|
+
noise_std=args.noise_std,
|
|
384
|
+
epochs=args.epochs,
|
|
385
|
+
lr=args.lr,
|
|
386
|
+
seed=args.seed,
|
|
387
|
+
mask_threshold=args.threshold,
|
|
388
|
+
stability=not args.no_stability,
|
|
389
|
+
)
|
|
390
|
+
|
|
391
|
+
headers = ("Method", "fid+", "fid-", "GEA", "sparsity", "stab")
|
|
392
|
+
widths = [len(h) for h in headers]
|
|
393
|
+
rows: list[tuple[Any, ...]] = []
|
|
394
|
+
for name, entry in results.items():
|
|
395
|
+
if name.startswith("_"):
|
|
396
|
+
continue
|
|
397
|
+
if entry["skipped"]:
|
|
398
|
+
rows.append((name, "not applicable", "", "", "", ""))
|
|
399
|
+
continue
|
|
400
|
+
m = entry["metrics"]
|
|
401
|
+
rows.append(
|
|
402
|
+
(
|
|
403
|
+
name,
|
|
404
|
+
"" if m["fidelity_plus"] is None else f"{m['fidelity_plus']:.3f}",
|
|
405
|
+
"" if m["fidelity_minus"] is None else f"{m['fidelity_minus']:.3f}",
|
|
406
|
+
"" if m["gea"] is None else f"{m['gea']:.3f}",
|
|
407
|
+
"" if m["sparsity"] is None else f"{m['sparsity']:.3f}",
|
|
408
|
+
"" if m["stability"] is None else f"{m['stability']:.3f}",
|
|
409
|
+
)
|
|
410
|
+
)
|
|
411
|
+
|
|
412
|
+
widths[0] = max(widths[0], max((len(r[0]) for r in rows), default=0))
|
|
413
|
+
for i in range(1, len(headers)):
|
|
414
|
+
widths[i] = max(widths[i], max((len(r[i]) for r in rows if r[i]), default=0))
|
|
415
|
+
|
|
416
|
+
for i, h in enumerate(headers):
|
|
417
|
+
widths[i] = max(widths[i], len(h))
|
|
418
|
+
line = " ".join(h.ljust(widths[i]) for i, h in enumerate(headers))
|
|
419
|
+
print(line)
|
|
420
|
+
print(" ".join("-" * w for w in widths))
|
|
421
|
+
for row in rows:
|
|
422
|
+
print(" ".join(str(c).ljust(widths[i]) for i, c in enumerate(row)))
|
|
423
|
+
|
|
424
|
+
skipped = results["_meta"]["skipped"]
|
|
425
|
+
if skipped:
|
|
426
|
+
print("\nSkipped:")
|
|
427
|
+
for name, reason in skipped.items():
|
|
428
|
+
print(f" {name}: {reason}")
|
|
429
|
+
|
|
430
|
+
if args.json:
|
|
431
|
+
out = {
|
|
432
|
+
name: entry for name, entry in results.items() if not name.startswith("_")
|
|
433
|
+
}
|
|
434
|
+
out["_meta"] = results["_meta"]
|
|
435
|
+
out["_meta"]["version"] = __version__
|
|
436
|
+
with open(args.json, "w", encoding="utf-8") as fh:
|
|
437
|
+
json.dump(_json_safe(out), fh, indent=2, ensure_ascii=False)
|
|
438
|
+
print(f"\nResults saved to JSON at {args.json}")
|
|
439
|
+
if args.html:
|
|
440
|
+
report_html(results, args.html)
|
|
441
|
+
print(f"HTML report saved to {args.html}")
|
|
442
|
+
return 0
|
|
443
|
+
|
|
444
|
+
|
|
445
|
+
def main(argv: list[str] | None = None) -> int:
|
|
446
|
+
parser = build_parser()
|
|
447
|
+
args = parser.parse_args(argv)
|
|
448
|
+
if args.command is None:
|
|
449
|
+
parser.print_help()
|
|
450
|
+
return 1
|
|
451
|
+
if args.command == "explain":
|
|
452
|
+
return _cmd_explain(args)
|
|
453
|
+
if args.command == "bench":
|
|
454
|
+
return _cmd_bench(args)
|
|
455
|
+
return 1
|
|
456
|
+
|
|
457
|
+
|
|
458
|
+
if __name__ == "__main__":
|
|
459
|
+
sys.exit(main())
|
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
from .benchmark import compare, report_html
|
|
2
|
+
from .explainer import Explainer
|
|
3
|
+
from .explanation import Explanation
|
|
4
|
+
from .registry import get_algorithm, instantiate, register
|
|
5
|
+
|
|
6
|
+
__all__ = [
|
|
7
|
+
"Explainer",
|
|
8
|
+
"Explanation",
|
|
9
|
+
"compare",
|
|
10
|
+
"get_algorithm",
|
|
11
|
+
"instantiate",
|
|
12
|
+
"register",
|
|
13
|
+
"report_html",
|
|
14
|
+
]
|