graph-explain 0.7.0__tar.gz → 0.7.2__tar.gz

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 (61) hide show
  1. graph_explain-0.7.2/LICENSE +21 -0
  2. {graph_explain-0.7.0 → graph_explain-0.7.2}/PKG-INFO +8 -1
  3. {graph_explain-0.7.0 → graph_explain-0.7.2}/pyproject.toml +12 -2
  4. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/__init__.py +1 -1
  5. graph_explain-0.7.2/src/graph_explain/narration/narrator.py +246 -0
  6. graph_explain-0.7.2/src/graph_explain/py.typed +0 -0
  7. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain.egg-info/PKG-INFO +8 -1
  8. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain.egg-info/SOURCES.txt +2 -0
  9. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain.egg-info/requires.txt +2 -0
  10. {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_narration.py +46 -0
  11. graph_explain-0.7.0/src/graph_explain/narration/narrator.py +0 -185
  12. {graph_explain-0.7.0 → graph_explain-0.7.2}/README.md +0 -0
  13. {graph_explain-0.7.0 → graph_explain-0.7.2}/setup.cfg +0 -0
  14. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/backends/__init__.py +0 -0
  15. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/backends/base.py +0 -0
  16. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/backends/dgl.py +0 -0
  17. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/benchmarks/__init__.py +0 -0
  18. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/benchmarks/synthetic.py +0 -0
  19. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/cli.py +0 -0
  20. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/core/__init__.py +0 -0
  21. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/core/benchmark.py +0 -0
  22. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/core/evaluation.py +0 -0
  23. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/core/explainer.py +0 -0
  24. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/core/explanation.py +0 -0
  25. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/core/model_utils.py +0 -0
  26. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/core/registry.py +0 -0
  27. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/__init__.py +0 -0
  28. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/attention/attention.py +0 -0
  29. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/base.py +0 -0
  30. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/baseline/random_baseline.py +0 -0
  31. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/counterfactual/counterfactual.py +0 -0
  32. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/feature/graph_lime.py +0 -0
  33. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/gradient/__init__.py +0 -0
  34. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/gradient/grad_x_input.py +0 -0
  35. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/gradient/guided_backprop.py +0 -0
  36. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/gradient/integrated_gradients.py +0 -0
  37. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/gradient/saliency.py +0 -0
  38. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/perturbation/__init__.py +0 -0
  39. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/perturbation/gnn_explainer.py +0 -0
  40. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/perturbation/node_mask.py +0 -0
  41. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/perturbation/pg_explainer.py +0 -0
  42. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/perturbation/subgraphx.py +0 -0
  43. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/relevance/deeplift.py +0 -0
  44. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/methods/relevance/gnn_lrp.py +0 -0
  45. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/narration/__init__.py +0 -0
  46. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/visualization/__init__.py +0 -0
  47. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/visualization/interactive.py +0 -0
  48. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain/visualization/static.py +0 -0
  49. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain.egg-info/dependency_links.txt +0 -0
  50. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain.egg-info/entry_points.txt +0 -0
  51. {graph_explain-0.7.0 → graph_explain-0.7.2}/src/graph_explain.egg-info/top_level.txt +0 -0
  52. {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_cli.py +0 -0
  53. {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_core.py +0 -0
  54. {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_counterfactual.py +0 -0
  55. {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_dgl_integration.py +0 -0
  56. {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_phase10.py +0 -0
  57. {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_phase2.py +0 -0
  58. {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_phase3.py +0 -0
  59. {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_phase4.py +0 -0
  60. {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_phase6.py +0 -0
  61. {graph_explain-0.7.0 → graph_explain-0.7.2}/tests/test_phase8.py +0 -0
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2026 graph-explain contributors
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
@@ -1,9 +1,12 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: graph-explain
3
- Version: 0.7.0
3
+ Version: 0.7.2
4
4
  Summary: Explainability library for graph-based models (GNN)
5
5
  Author: graph-explain contributors
6
6
  License-Expression: MIT
7
+ Project-URL: Homepage, https://github.com/Tzinny-dev/graph-explain
8
+ Project-URL: Repository, https://github.com/Tzinny-dev/graph-explain
9
+ Project-URL: Documentation, https://github.com/Tzinny-dev/graph-explain#readme
7
10
  Keywords: gnn,explainability,xai,graph,neural-networks,interpretability
8
11
  Classifier: Development Status :: 4 - Beta
9
12
  Classifier: Intended Audience :: Science/Research
@@ -16,9 +19,11 @@ Classifier: Programming Language :: Python :: 3.11
16
19
  Classifier: Programming Language :: Python :: 3.12
17
20
  Requires-Python: >=3.10
18
21
  Description-Content-Type: text/markdown
22
+ License-File: LICENSE
19
23
  Requires-Dist: numpy>=1.24
20
24
  Requires-Dist: networkx>=3.0
21
25
  Requires-Dist: matplotlib>=3.6
26
+ Requires-Dist: torch>=2.0
22
27
  Provides-Extra: pyg
23
28
  Requires-Dist: torch>=2.0; extra == "pyg"
24
29
  Requires-Dist: torch-geometric>=2.5; extra == "pyg"
@@ -31,6 +36,7 @@ Requires-Dist: pyvis>=0.3; extra == "interactive"
31
36
  Provides-Extra: all
32
37
  Requires-Dist: torch>=2.0; extra == "all"
33
38
  Requires-Dist: torch-geometric>=2.5; extra == "all"
39
+ Requires-Dist: dgl>=2.0; extra == "all"
34
40
  Requires-Dist: plotly>=5.15; extra == "all"
35
41
  Requires-Dist: pyvis>=0.3; extra == "all"
36
42
  Provides-Extra: dev
@@ -42,6 +48,7 @@ Requires-Dist: sphinx>=7.2; extra == "docs"
42
48
  Requires-Dist: sphinx-rtd-theme>=2.0; extra == "docs"
43
49
  Provides-Extra: publish
44
50
  Requires-Dist: twine>=5.0; extra == "publish"
51
+ Dynamic: license-file
45
52
 
46
53
  # graph-explain
47
54
 
@@ -4,11 +4,12 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "graph-explain"
7
- version = "0.7.0"
7
+ version = "0.7.2"
8
8
  description = "Explainability library for graph-based models (GNN)"
9
9
  readme = "README.md"
10
10
  requires-python = ">=3.10"
11
11
  license = "MIT"
12
+ license-files = ["LICENSE"]
12
13
  authors = [{ name = "graph-explain contributors" }]
13
14
  keywords = [
14
15
  "gnn",
@@ -34,13 +35,19 @@ dependencies = [
34
35
  "numpy>=1.24",
35
36
  "networkx>=3.0",
36
37
  "matplotlib>=3.6",
38
+ "torch>=2.0",
37
39
  ]
38
40
 
41
+ [project.urls]
42
+ Homepage = "https://github.com/Tzinny-dev/graph-explain"
43
+ Repository = "https://github.com/Tzinny-dev/graph-explain"
44
+ Documentation = "https://github.com/Tzinny-dev/graph-explain#readme"
45
+
39
46
  [project.optional-dependencies]
40
47
  pyg = ["torch>=2.0", "torch-geometric>=2.5"]
41
48
  dgl = ["torch>=2.0", "dgl>=2.0"]
42
49
  interactive = ["plotly>=5.15", "pyvis>=0.3"]
43
- all = ["torch>=2.0", "torch-geometric>=2.5", "plotly>=5.15", "pyvis>=0.3"]
50
+ all = ["torch>=2.0", "torch-geometric>=2.5", "dgl>=2.0", "plotly>=5.15", "pyvis>=0.3"]
44
51
  dev = ["pytest>=7.0", "ruff>=0.5", "build>=1.0"]
45
52
  docs = ["sphinx>=7.2", "sphinx-rtd-theme>=2.0"]
46
53
  publish = ["twine>=5.0"]
@@ -51,6 +58,9 @@ graph-explain = "graph_explain.cli:main"
51
58
  [tool.setuptools.packages.find]
52
59
  where = ["src"]
53
60
 
61
+ [tool.setuptools.package-data]
62
+ graph_explain = ["py.typed"]
63
+
54
64
  [tool.pytest.ini_options]
55
65
  testpaths = ["tests"]
56
66
 
@@ -35,7 +35,7 @@ from .methods import (
35
35
  from .narration import Narrator, describe, narrate, summarize
36
36
  from .visualization import show, visualize_interactive, visualize_static
37
37
 
38
- __version__ = "0.7.0"
38
+ __version__ = "0.7.2"
39
39
 
40
40
  __all__ = [
41
41
  "AttentionExplainer",
@@ -0,0 +1,246 @@
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ from collections.abc import Callable
5
+ from typing import Any
6
+
7
+ import torch
8
+
9
+
10
+ def _top_values(importance, k: int) -> list[tuple[int, float]]:
11
+ imp = importance.detach().reshape(-1)
12
+ n = int(imp.numel())
13
+ k = max(1, min(int(k), n))
14
+ idx = imp.argsort(descending=True)[:k]
15
+ return [(int(i), float(imp[i])) for i in idx.tolist()]
16
+
17
+
18
+ def _data_context(explanation, data: Any | None):
19
+ if data is None:
20
+ data = explanation.metadata.get("backing_data")
21
+ backend = explanation.metadata.get("backend")
22
+ return backend, data
23
+
24
+
25
+ def _labels(explanation, data: Any | None) -> Any | None:
26
+ backend, data = _data_context(explanation, data)
27
+ if backend is None or data is None:
28
+ return None
29
+ try:
30
+ return backend.node_labels(data)
31
+ except Exception: # noqa: BLE001
32
+ return None
33
+
34
+
35
+ def _edge_index(data, backend) -> Any | None:
36
+ if backend is not None and data is not None:
37
+ try:
38
+ return backend.edge_index(data)
39
+ except Exception: # noqa: BLE001
40
+ return None
41
+ return None
42
+
43
+
44
+ _SUPPORTED_LANGS = ("es", "en")
45
+
46
+ _TEMPLATES: dict[str, dict[str, str]] = {
47
+ "es": {
48
+ "graph_head": "Explicación a nivel de grafo.",
49
+ "node_head": "Explicación del nodo {node}.",
50
+ "target": " La clase objetivo es {target}.",
51
+ "correct": " La predicción del modelo es correcta.",
52
+ "incorrect": " La predicción del modelo difiere de la etiqueta real.",
53
+ "nodes": "Los nodos más relevantes son {nodes}.",
54
+ "edges": " Las aristas más relevantes son {edges}.",
55
+ "node_item": "nodo {i} (importancia {v:.3f})",
56
+ "edge_item": "arista {u}-{v} ({w:.3f})",
57
+ "edge_item_float": "arista con relevancia {v:.3f}",
58
+ "cf_intro": " Se necesitaron {n} cambios (aristas/features eliminadas) "
59
+ "para cambiar la predicción de {change}.",
60
+ "cf_change": "la clase {orig} a la clase {new}",
61
+ "cf_change_solo": "la clase {orig}",
62
+ "nodata": "Sin datos suficientes para describir la explicación.",
63
+ "prompt": "Eres un asistente que explica predicciones de GNNs en lenguaje "
64
+ "natural. Dado este resumen de una explicación (JSON), escribe un párrafo "
65
+ "breve en español (2-4 oraciones) describiendo qué hace el modelo y qué "
66
+ "evidencia respalda su predicción. Resumen:\n",
67
+ "llm_fallback": "[LLM no disponible: {exc}]",
68
+ },
69
+ "en": {
70
+ "graph_head": "Graph-level explanation.",
71
+ "node_head": "Explanation of node {node}.",
72
+ "target": " The target class is {target}.",
73
+ "correct": " The model prediction is correct.",
74
+ "incorrect": " The model prediction differs from the true label.",
75
+ "nodes": "The most relevant nodes are {nodes}.",
76
+ "edges": " The most relevant edges are {edges}.",
77
+ "node_item": "node {i} (importance {v:.3f})",
78
+ "edge_item": "edge {u}-{v} ({w:.3f})",
79
+ "edge_item_float": "edge with relevance {v:.3f}",
80
+ "cf_intro": " {n} changes (removed edges/features) were needed to change the "
81
+ "prediction from {change}.",
82
+ "cf_change": "class {orig} to class {new}",
83
+ "cf_change_solo": "class {orig}",
84
+ "nodata": "Not enough data to describe the explanation.",
85
+ "prompt": "You are an assistant that explains GNN predictions in natural "
86
+ "language. Given this JSON summary of an explanation, write a brief "
87
+ "paragraph (2-4 sentences) in English describing what the model does and "
88
+ "what evidence supports its prediction. Summary:\n",
89
+ "llm_fallback": "[LLM unavailable: {exc}]",
90
+ },
91
+ }
92
+
93
+
94
+ def _templates(lang: str) -> dict[str, str]:
95
+ if lang not in _SUPPORTED_LANGS:
96
+ raise ValueError(f"lang must be one of {_SUPPORTED_LANGS}, got {lang!r}")
97
+ return _TEMPLATES[lang]
98
+
99
+
100
+ def summarize(explanation, data: Any | None = None, top_k: int = 5) -> dict[str, Any]:
101
+ """Structured summary of an explanation (for narration or JSON)."""
102
+ backend, data = _data_context(explanation, data)
103
+ node = explanation.node_idx
104
+ target = explanation.target_class
105
+ pred = None
106
+ if explanation.prediction_original is not None:
107
+ p = explanation.prediction_original
108
+ if torch.is_tensor(p):
109
+ pred = int(p.reshape(-1).argmax().item())
110
+ labels = _labels(explanation, data)
111
+
112
+ true_label = None
113
+ if labels is not None and node is not None:
114
+ try:
115
+ true_label = int(labels[node].item())
116
+ except Exception: # noqa: BLE001
117
+ true_label = None
118
+
119
+ summary: dict[str, Any] = {
120
+ "node": None if node is None else int(node),
121
+ "target_class": target,
122
+ "predicted_class": pred,
123
+ "true_class": true_label,
124
+ "correct": (
125
+ None if pred is None or true_label is None else bool(pred == true_label)
126
+ ),
127
+ "important_nodes": (
128
+ _top_values(explanation.node_importance, top_k)
129
+ if explanation.node_importance is not None
130
+ else []
131
+ ),
132
+ "important_edges": [],
133
+ "counterfactual": bool(explanation.metadata.get("counterfactual", False)),
134
+ }
135
+ if explanation.edge_importance is not None:
136
+ ei = _edge_index(data, backend)
137
+ top = _top_values(explanation.edge_importance, top_k)
138
+ if ei is None:
139
+ summary["important_edges"] = [v for _, v in top]
140
+ else:
141
+ summary["important_edges"] = [
142
+ (int(ei[0, i]), int(ei[1, i]), v) for i, v in top
143
+ ]
144
+ if summary["counterfactual"]:
145
+ summary["original_class"] = explanation.metadata.get("original_class")
146
+ return summary
147
+
148
+
149
+ def describe(
150
+ explanation, data: Any | None = None, top_k: int = 5, lang: str = "es"
151
+ ) -> str:
152
+ """Deterministic template-based narration of an explanation.
153
+
154
+ Args:
155
+ lang: Template language, ``"es"`` (default) or ``"en"``.
156
+ """
157
+ _T = _templates(lang)
158
+ s = summarize(explanation, data, top_k)
159
+ node = s["node"]
160
+ target = (
161
+ s["target_class"] if s["target_class"] is not None else s["predicted_class"]
162
+ )
163
+
164
+ if node is None:
165
+ head = _T["graph_head"]
166
+ else:
167
+ head = _T["node_head"].format(node=node)
168
+ if target is not None:
169
+ head += _T["target"].format(target=target)
170
+ if s["correct"] is True:
171
+ head += _T["correct"]
172
+ elif s["correct"] is False:
173
+ head += _T["incorrect"]
174
+
175
+ nodes_txt = ", ".join(
176
+ _T["node_item"].format(i=i, v=v) for i, v in s["important_nodes"]
177
+ )
178
+ tail = _T["nodes"].format(nodes=nodes_txt) if nodes_txt else ""
179
+
180
+ if s["important_edges"]:
181
+ pieces = []
182
+ for e in s["important_edges"]:
183
+ if len(e) == 3:
184
+ u, v, w = e
185
+ pieces.append(_T["edge_item"].format(u=u, v=v, w=w))
186
+ else:
187
+ pieces.append(_T["edge_item_float"].format(v=float(e)))
188
+ tail += _T["edges"].format(edges=", ".join(pieces))
189
+
190
+ if s["counterfactual"]:
191
+ n = len(s["important_edges"])
192
+ new_class = s["predicted_class"]
193
+ if new_class is not None:
194
+ change = _T["cf_change"].format(orig=s["original_class"], new=new_class)
195
+ else:
196
+ change = _T["cf_change_solo"].format(orig=s["original_class"])
197
+ tail += _T["cf_intro"].format(n=n, change=change)
198
+ if tail:
199
+ head += " " + tail.strip()
200
+ return head.strip() or _T["nodata"]
201
+
202
+
203
+ def _prompt(summary: dict[str, Any], lang: str = "es") -> str:
204
+ _T = _templates(lang)
205
+ return _T["prompt"] + json.dumps(summary, ensure_ascii=False, indent=2)
206
+
207
+
208
+ def narrate(
209
+ explanation,
210
+ llm: Callable[[str], str] | None = None,
211
+ data: Any | None = None,
212
+ top_k: int = 5,
213
+ lang: str = "es",
214
+ ) -> str:
215
+ """Narrates an explanation. With `llm` (a `prompt -> text` callable) it uses the
216
+ generative model's output; otherwise it falls back to deterministic
217
+ template-based narration."""
218
+ _T = _templates(lang)
219
+ summary = summarize(explanation, data, top_k)
220
+ deterministic = describe(explanation, data, top_k, lang=lang)
221
+ if llm is None:
222
+ return deterministic
223
+ try:
224
+ return llm(_prompt(summary, lang)).strip()
225
+ except Exception as exc: # noqa: BLE001
226
+ return f"{deterministic}\n\n{_T['llm_fallback'].format(exc=exc)}"
227
+
228
+
229
+ class Narrator:
230
+ """Reusable narrator; lets you inject the LLM just once."""
231
+
232
+ def __init__(
233
+ self,
234
+ llm: Callable[[str], str] | None = None,
235
+ top_k: int = 5,
236
+ lang: str = "es",
237
+ ):
238
+ self.llm = llm
239
+ self.top_k = top_k
240
+ self.lang = lang
241
+
242
+ def describe(self, explanation, data: Any | None = None) -> str:
243
+ return describe(explanation, data, self.top_k, lang=self.lang)
244
+
245
+ def narrate(self, explanation, data: Any | None = None) -> str:
246
+ return narrate(explanation, self.llm, data, self.top_k, lang=self.lang)
File without changes
@@ -1,9 +1,12 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: graph-explain
3
- Version: 0.7.0
3
+ Version: 0.7.2
4
4
  Summary: Explainability library for graph-based models (GNN)
5
5
  Author: graph-explain contributors
6
6
  License-Expression: MIT
7
+ Project-URL: Homepage, https://github.com/Tzinny-dev/graph-explain
8
+ Project-URL: Repository, https://github.com/Tzinny-dev/graph-explain
9
+ Project-URL: Documentation, https://github.com/Tzinny-dev/graph-explain#readme
7
10
  Keywords: gnn,explainability,xai,graph,neural-networks,interpretability
8
11
  Classifier: Development Status :: 4 - Beta
9
12
  Classifier: Intended Audience :: Science/Research
@@ -16,9 +19,11 @@ Classifier: Programming Language :: Python :: 3.11
16
19
  Classifier: Programming Language :: Python :: 3.12
17
20
  Requires-Python: >=3.10
18
21
  Description-Content-Type: text/markdown
22
+ License-File: LICENSE
19
23
  Requires-Dist: numpy>=1.24
20
24
  Requires-Dist: networkx>=3.0
21
25
  Requires-Dist: matplotlib>=3.6
26
+ Requires-Dist: torch>=2.0
22
27
  Provides-Extra: pyg
23
28
  Requires-Dist: torch>=2.0; extra == "pyg"
24
29
  Requires-Dist: torch-geometric>=2.5; extra == "pyg"
@@ -31,6 +36,7 @@ Requires-Dist: pyvis>=0.3; extra == "interactive"
31
36
  Provides-Extra: all
32
37
  Requires-Dist: torch>=2.0; extra == "all"
33
38
  Requires-Dist: torch-geometric>=2.5; extra == "all"
39
+ Requires-Dist: dgl>=2.0; extra == "all"
34
40
  Requires-Dist: plotly>=5.15; extra == "all"
35
41
  Requires-Dist: pyvis>=0.3; extra == "all"
36
42
  Provides-Extra: dev
@@ -42,6 +48,7 @@ Requires-Dist: sphinx>=7.2; extra == "docs"
42
48
  Requires-Dist: sphinx-rtd-theme>=2.0; extra == "docs"
43
49
  Provides-Extra: publish
44
50
  Requires-Dist: twine>=5.0; extra == "publish"
51
+ Dynamic: license-file
45
52
 
46
53
  # graph-explain
47
54
 
@@ -1,7 +1,9 @@
1
+ LICENSE
1
2
  README.md
2
3
  pyproject.toml
3
4
  src/graph_explain/__init__.py
4
5
  src/graph_explain/cli.py
6
+ src/graph_explain/py.typed
5
7
  src/graph_explain.egg-info/PKG-INFO
6
8
  src/graph_explain.egg-info/SOURCES.txt
7
9
  src/graph_explain.egg-info/dependency_links.txt
@@ -1,10 +1,12 @@
1
1
  numpy>=1.24
2
2
  networkx>=3.0
3
3
  matplotlib>=3.6
4
+ torch>=2.0
4
5
 
5
6
  [all]
6
7
  torch>=2.0
7
8
  torch-geometric>=2.5
9
+ dgl>=2.0
8
10
  plotly>=5.15
9
11
  pyvis>=0.3
10
12
 
@@ -74,3 +74,49 @@ class TestNarration:
74
74
  _, expl = self._setup()
75
75
  narrator = Narrator(llm=None)
76
76
  assert narrator.describe(expl) == describe(expl)
77
+
78
+ def test_describe_english(self):
79
+ data, expl = self._setup()
80
+ text = describe(expl, data=data, lang="en")
81
+ assert "node" in text
82
+ assert "class" in text
83
+ assert "most relevant" in text
84
+ assert text == describe(expl, data=data, lang="en") # determinista
85
+
86
+ def test_english_vs_spanish(self):
87
+ data, expl = self._setup()
88
+ es = describe(expl, data=data, lang="es")
89
+ en = describe(expl, data=data, lang="en")
90
+ assert es != en
91
+ assert "nodo" in es and "node" in en
92
+
93
+ def test_narrate_with_llm_english(self):
94
+ _, expl = self._setup()
95
+
96
+ def fake_llm(prompt):
97
+ return "CUSTOM NARRATION"
98
+
99
+ text = narrate(expl, llm=fake_llm, lang="en")
100
+ assert text == "CUSTOM NARRATION"
101
+
102
+ def test_narrate_fallback_english(self):
103
+ _, expl = self._setup()
104
+
105
+ def broken_llm(_prompt):
106
+ raise RuntimeError("no api")
107
+
108
+ text = narrate(expl, llm=broken_llm, lang="en")
109
+ assert "LLM unavailable" in text
110
+ assert "Explanation of node" in text
111
+
112
+ def test_invalid_lang(self):
113
+ import pytest
114
+
115
+ _, expl = self._setup()
116
+ with pytest.raises(ValueError):
117
+ describe(expl, lang="fr")
118
+
119
+ def test_narrator_class_english(self):
120
+ _, expl = self._setup()
121
+ narrator = Narrator(llm=None, lang="en")
122
+ assert narrator.describe(expl) == describe(expl, lang="en")
@@ -1,185 +0,0 @@
1
- from __future__ import annotations
2
-
3
- import json
4
- from collections.abc import Callable
5
- from typing import Any
6
-
7
- import torch
8
-
9
-
10
- def _top_values(importance, k: int) -> list[tuple[int, float]]:
11
- imp = importance.detach().reshape(-1)
12
- n = int(imp.numel())
13
- k = max(1, min(int(k), n))
14
- idx = imp.argsort(descending=True)[:k]
15
- return [(int(i), float(imp[i])) for i in idx.tolist()]
16
-
17
-
18
- def _data_context(explanation, data: Any | None):
19
- if data is None:
20
- data = explanation.metadata.get("backing_data")
21
- backend = explanation.metadata.get("backend")
22
- return backend, data
23
-
24
-
25
- def _labels(explanation, data: Any | None) -> Any | None:
26
- backend, data = _data_context(explanation, data)
27
- if backend is None or data is None:
28
- return None
29
- try:
30
- return backend.node_labels(data)
31
- except Exception: # noqa: BLE001
32
- return None
33
-
34
-
35
- def _edge_index(data, backend) -> Any | None:
36
- if backend is not None and data is not None:
37
- try:
38
- return backend.edge_index(data)
39
- except Exception: # noqa: BLE001
40
- return None
41
- return None
42
-
43
-
44
- def summarize(explanation, data: Any | None = None, top_k: int = 5) -> dict[str, Any]:
45
- """Structured summary of an explanation (for narration or JSON)."""
46
- backend, data = _data_context(explanation, data)
47
- node = explanation.node_idx
48
- target = explanation.target_class
49
- pred = None
50
- if explanation.prediction_original is not None:
51
- p = explanation.prediction_original
52
- if torch.is_tensor(p):
53
- pred = int(p.reshape(-1).argmax().item())
54
- labels = _labels(explanation, data)
55
-
56
- true_label = None
57
- if labels is not None and node is not None:
58
- try:
59
- true_label = int(labels[node].item())
60
- except Exception: # noqa: BLE001
61
- true_label = None
62
-
63
- summary: dict[str, Any] = {
64
- "node": None if node is None else int(node),
65
- "target_class": target,
66
- "predicted_class": pred,
67
- "true_class": true_label,
68
- "correct": (
69
- None if pred is None or true_label is None else bool(pred == true_label)
70
- ),
71
- "important_nodes": (
72
- _top_values(explanation.node_importance, top_k)
73
- if explanation.node_importance is not None
74
- else []
75
- ),
76
- "important_edges": [],
77
- "counterfactual": bool(explanation.metadata.get("counterfactual", False)),
78
- }
79
- if explanation.edge_importance is not None:
80
- ei = _edge_index(data, backend)
81
- top = _top_values(explanation.edge_importance, top_k)
82
- if ei is None:
83
- summary["important_edges"] = [v for _, v in top]
84
- else:
85
- summary["important_edges"] = [
86
- (int(ei[0, i]), int(ei[1, i]), v) for i, v in top
87
- ]
88
- if summary["counterfactual"]:
89
- summary["original_class"] = explanation.metadata.get("original_class")
90
- return summary
91
-
92
-
93
- def describe(explanation, data: Any | None = None, top_k: int = 5) -> str:
94
- """Deterministic template-based narration of an explanation (Spanish by default)."""
95
- s = summarize(explanation, data, top_k)
96
- node = s["node"]
97
- target = (
98
- s["target_class"] if s["target_class"] is not None else s["predicted_class"]
99
- )
100
-
101
- if node is None:
102
- head = "Explicación a nivel de grafo."
103
- else:
104
- head = f"Explicación del nodo {node}."
105
- if target is not None:
106
- head += f" La clase objetivo es {target}."
107
- if s["correct"] is True:
108
- head += " La predicción del modelo es correcta."
109
- elif s["correct"] is False:
110
- head += " La predicción del modelo difiere de la etiqueta real."
111
-
112
- nodes_txt = ", ".join(
113
- f"nodo {i} (importancia {v:.3f})" for i, v in s["important_nodes"]
114
- )
115
- tail = f"Los nodos más relevantes son {nodes_txt}." if nodes_txt else ""
116
-
117
- if s["important_edges"]:
118
- pieces = []
119
- for e in s["important_edges"]:
120
- if len(e) == 3:
121
- u, v, w = e
122
- pieces.append(f"arista {u}-{v} ({w:.3f})")
123
- else:
124
- pieces.append(f"arista con relevancia {float(e):.3f}")
125
- tail += " Las aristas más relevantes son " + ", ".join(pieces) + "."
126
-
127
- if s["counterfactual"]:
128
- n = len(s["important_edges"])
129
- new_class = s["predicted_class"]
130
- change = (
131
- f"la clase {s['original_class']} a la clase {new_class}"
132
- if new_class is not None
133
- else f"la clase {s['original_class']}"
134
- )
135
- tail += (
136
- f" Se necesitaron {n} cambios"
137
- f" (aristas/features eliminadas) para cambiar la predicción de"
138
- f" {change}."
139
- )
140
- if tail:
141
- head += " " + tail.strip()
142
- return head.strip() or "Sin datos suficientes para describir la explicación."
143
-
144
-
145
- def _prompt(summary: dict[str, Any]) -> str:
146
- return (
147
- "Eres un asistente que explica predicciones de GNNs en lenguaje natural. "
148
- "Dado este resumen de una explicación (JSON), escribe un párrafo breve "
149
- "en español (2-4 oraciones) describiendo qué hace el modelo y qué "
150
- "evidencia respalda su predicción. Resumen:\n"
151
- + json.dumps(summary, ensure_ascii=False, indent=2)
152
- )
153
-
154
-
155
- def narrate(
156
- explanation,
157
- llm: Callable[[str], str] | None = None,
158
- data: Any | None = None,
159
- top_k: int = 5,
160
- ) -> str:
161
- """Narrates an explanation. With `llm` (a `prompt -> text` callable) it uses the
162
- generative model's output; otherwise it falls back to deterministic
163
- template-based narration."""
164
- summary = summarize(explanation, data, top_k)
165
- deterministic = describe(explanation, data, top_k)
166
- if llm is None:
167
- return deterministic
168
- try:
169
- return llm(_prompt(summary)).strip()
170
- except Exception as exc: # noqa: BLE001
171
- return f"{deterministic}\n\n[LLM no disponible: {exc}]"
172
-
173
-
174
- class Narrator:
175
- """Reusable narrator; lets you inject the LLM just once."""
176
-
177
- def __init__(self, llm: Callable[[str], str] | None = None, top_k: int = 5):
178
- self.llm = llm
179
- self.top_k = top_k
180
-
181
- def describe(self, explanation, data: Any | None = None) -> str:
182
- return describe(explanation, data, self.top_k)
183
-
184
- def narrate(self, explanation, data: Any | None = None) -> str:
185
- return narrate(explanation, self.llm, data, self.top_k)
File without changes
File without changes