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.
Files changed (42) hide show
  1. graph_explain/__init__.py +79 -0
  2. graph_explain/backends/__init__.py +4 -0
  3. graph_explain/backends/base.py +103 -0
  4. graph_explain/backends/dgl.py +121 -0
  5. graph_explain/benchmarks/__init__.py +3 -0
  6. graph_explain/benchmarks/synthetic.py +246 -0
  7. graph_explain/cli.py +459 -0
  8. graph_explain/core/__init__.py +14 -0
  9. graph_explain/core/benchmark.py +284 -0
  10. graph_explain/core/evaluation.py +391 -0
  11. graph_explain/core/explainer.py +83 -0
  12. graph_explain/core/explanation.py +72 -0
  13. graph_explain/core/model_utils.py +44 -0
  14. graph_explain/core/registry.py +55 -0
  15. graph_explain/methods/__init__.py +39 -0
  16. graph_explain/methods/attention/attention.py +147 -0
  17. graph_explain/methods/base.py +25 -0
  18. graph_explain/methods/baseline/random_baseline.py +78 -0
  19. graph_explain/methods/counterfactual/counterfactual.py +304 -0
  20. graph_explain/methods/feature/graph_lime.py +141 -0
  21. graph_explain/methods/gradient/__init__.py +0 -0
  22. graph_explain/methods/gradient/grad_x_input.py +110 -0
  23. graph_explain/methods/gradient/guided_backprop.py +117 -0
  24. graph_explain/methods/gradient/integrated_gradients.py +115 -0
  25. graph_explain/methods/gradient/saliency.py +93 -0
  26. graph_explain/methods/perturbation/__init__.py +0 -0
  27. graph_explain/methods/perturbation/gnn_explainer.py +265 -0
  28. graph_explain/methods/perturbation/node_mask.py +136 -0
  29. graph_explain/methods/perturbation/pg_explainer.py +162 -0
  30. graph_explain/methods/perturbation/subgraphx.py +393 -0
  31. graph_explain/methods/relevance/deeplift.py +262 -0
  32. graph_explain/methods/relevance/gnn_lrp.py +219 -0
  33. graph_explain/narration/__init__.py +3 -0
  34. graph_explain/narration/narrator.py +185 -0
  35. graph_explain/visualization/__init__.py +4 -0
  36. graph_explain/visualization/interactive.py +73 -0
  37. graph_explain/visualization/static.py +90 -0
  38. graph_explain-0.7.0.dist-info/METADATA +332 -0
  39. graph_explain-0.7.0.dist-info/RECORD +42 -0
  40. graph_explain-0.7.0.dist-info/WHEEL +5 -0
  41. graph_explain-0.7.0.dist-info/entry_points.txt +2 -0
  42. graph_explain-0.7.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,262 @@
1
+ from __future__ import annotations
2
+
3
+ from typing import Any
4
+
5
+ import torch
6
+ from torch import nn
7
+
8
+ from ...core.explanation import Explanation
9
+ from ...core.registry import register
10
+ from ..base import ExplanationAlgorithm
11
+
12
+ _ACTIVATION_GATES = (
13
+ nn.ReLU,
14
+ nn.ReLU6,
15
+ nn.LeakyReLU,
16
+ )
17
+
18
+
19
+ @register("deep_lift", "deeplift", "dl")
20
+ class DeepLift(ExplanationAlgorithm):
21
+ """DeepLIFT (rescale rule) for GCNs + ReLU + Linear.
22
+
23
+ It is an additive rule: each input feature receives a contribution (delta)
24
+ proportional to how much the target-class output changes when moving from a
25
+ baseline (zero, by default) to the actual instance. The multiplier is
26
+ propagated backwards layer by layer: exact for linear layers and GCN
27
+ messages, and with the rescale rule (delta_out / delta_in) for elementwise
28
+ nonlinearities.
29
+
30
+ Returns `node_importance` (absolute contribution per node), `edge_importance`
31
+ (contributions through the message passing of each GCNConv, per directed
32
+ edge) and `feature_importance` (contribution per feature).
33
+ """
34
+
35
+ def __init__(
36
+ self,
37
+ eps: float = 1e-7,
38
+ normalize: bool = False,
39
+ node_mask_type: str | None = "attributes",
40
+ ):
41
+ self.eps = float(eps)
42
+ self.normalize = bool(normalize)
43
+ self.node_mask_type = node_mask_type
44
+
45
+ def explain(
46
+ self,
47
+ backend: Any,
48
+ model: Any,
49
+ data: Any,
50
+ index: int | list[int] | torch.Tensor | None = None,
51
+ target_class: int | None = None,
52
+ **kwargs,
53
+ ) -> Explanation:
54
+ model.eval()
55
+ x = backend.node_features(data).detach()
56
+ edge_index = backend.edge_index(data)
57
+ edge_weight = backend.edge_weight(data)
58
+ num_nodes = int(x.size(0))
59
+
60
+ nodes = self._to_node_ids(index, num_nodes)
61
+
62
+ order, order0, logits, logits0 = self._capture(
63
+ model, backend, x, edge_index, edge_weight
64
+ )
65
+ if logits.dim() != 2:
66
+ raise ValueError(
67
+ "DeepLift requiere predicciones node-level (logits (N, C))."
68
+ )
69
+ target = target_class
70
+ if target is None:
71
+ target = int(logits[nodes[0]].argmax().item())
72
+ target_cls = max(0, min(int(target), logits.size(1) - 1))
73
+
74
+ mul = torch.zeros_like(logits)
75
+ mul[nodes, target_cls] = 1.0
76
+ edge_rel_full = torch.zeros(edge_index.size(1), device=x.device)
77
+ eps = self.eps
78
+
79
+ for m, m0 in reversed(list(zip(order, order0))):
80
+ in_x = m["args"][0]
81
+ in_x0 = m0["args"][0]
82
+ delta_in = in_x - in_x0
83
+ name = m["name"]
84
+ mod = m["module"]
85
+ if name == "linear":
86
+ mul = mul @ mod.weight
87
+ elif name == "gcn":
88
+ mul, edge_rel = self._conv_back(
89
+ mod, in_x, in_x0, edge_index, edge_weight, mul
90
+ )
91
+ edge_rel_full = edge_rel_full + edge_rel[: edge_index.size(1)]
92
+ elif name == "activation":
93
+ delta_out = mod(in_x) - mod(in_x0)
94
+ mul = mul * self._rescale_ratio(delta_in, delta_out, eps)
95
+ else:
96
+ continue
97
+
98
+ delta_total = logits[nodes, target_cls].sum() - logits0[nodes, target_cls].sum()
99
+ contrib = mul * x
100
+ node_importance = contrib.abs().sum(dim=-1)
101
+ edge_importance = edge_rel_full.abs()
102
+ self._last_delta_total = float(delta_total.item())
103
+
104
+ if self.normalize:
105
+ total = float(node_importance.sum().item())
106
+ if total > 0:
107
+ node_importance = node_importance / total
108
+ total_e = float(edge_importance.sum().item())
109
+ if total_e > 0:
110
+ edge_importance = edge_importance / total_e
111
+
112
+ return Explanation(
113
+ node_importance=node_importance.detach().cpu(),
114
+ edge_importance=edge_importance.detach().cpu(),
115
+ feature_importance=(
116
+ contrib.detach().cpu() if self.node_mask_type == "attributes" else None
117
+ ),
118
+ prediction_original=logits[nodes[0]].detach().cpu(),
119
+ prediction_explanation=None,
120
+ node_idx=int(nodes[0]) if nodes.shape[0] == 1 else index,
121
+ target_class=target_cls,
122
+ )
123
+
124
+ # ------------------------------------------------------------------ utils
125
+ @staticmethod
126
+ def _to_node_ids(index, num_nodes: int) -> torch.Tensor:
127
+ if index is None:
128
+ return torch.zeros(1, dtype=torch.long)
129
+ if isinstance(index, int):
130
+ return torch.tensor([index], dtype=torch.long)
131
+ idx = torch.as_tensor(index, dtype=torch.long)
132
+ return idx.reshape(-1) if idx.numel() else torch.zeros(1, dtype=torch.long)
133
+
134
+ @staticmethod
135
+ def _is_linear(module: nn.Module) -> bool:
136
+ return isinstance(module, nn.Linear)
137
+
138
+ @staticmethod
139
+ def _is_activation(module: nn.Module) -> bool:
140
+ return isinstance(module, _ACTIVATION_GATES)
141
+
142
+ @staticmethod
143
+ def _is_gcn(module: nn.Module) -> bool:
144
+ try:
145
+ from torch_geometric.nn import GCNConv
146
+
147
+ return isinstance(module, GCNConv)
148
+ except ImportError:
149
+ return False
150
+
151
+ def _capture(self, model, backend, x, edge_index, edge_weight):
152
+ order: list[dict[str, Any]] = []
153
+ order0: list[dict[str, Any]] = []
154
+
155
+ def _pre(module, args):
156
+ order.append(
157
+ {
158
+ "module": module,
159
+ "name": self._kind(module),
160
+ "args": tuple(
161
+ a.detach() if torch.is_tensor(a) else a for a in args
162
+ ),
163
+ }
164
+ )
165
+
166
+ def _pre0(module, args):
167
+ order0.append(
168
+ {
169
+ "module": module,
170
+ "name": self._kind(module),
171
+ "args": tuple(
172
+ a.detach() if torch.is_tensor(a) else a for a in args
173
+ ),
174
+ }
175
+ )
176
+
177
+ handles = [
178
+ module.register_forward_pre_hook(_pre)
179
+ for module in model.modules()
180
+ if module is not model
181
+ ]
182
+ with torch.no_grad():
183
+ logits = backend.forward(model, x, edge_index, edge_weight=edge_weight)
184
+ for handle in handles:
185
+ handle.remove()
186
+
187
+ handles0 = [
188
+ module.register_forward_pre_hook(_pre0)
189
+ for module in model.modules()
190
+ if module is not model
191
+ ]
192
+ baseline = torch.zeros_like(x)
193
+ with torch.no_grad():
194
+ logits0 = backend.forward(
195
+ model, baseline, edge_index, edge_weight=edge_weight
196
+ )
197
+ for handle in handles0:
198
+ handle.remove()
199
+ return order, order0, logits, logits0
200
+
201
+ @staticmethod
202
+ def _kind(module: nn.Module) -> str:
203
+ if DeepLift._is_linear(module):
204
+ return "linear"
205
+ if DeepLift._is_gcn(module):
206
+ return "gcn"
207
+ if DeepLift._is_activation(module):
208
+ return "activation"
209
+ return "other"
210
+
211
+ def _conv_back(
212
+ self,
213
+ conv: nn.Module,
214
+ x: torch.Tensor,
215
+ x0: torch.Tensor,
216
+ edge_index: torch.Tensor,
217
+ edge_weight: torch.Tensor | None,
218
+ mul: torch.Tensor,
219
+ ):
220
+ from torch_geometric.nn.conv.gcn_conv import gcn_norm
221
+ from torch_geometric.utils import add_self_loops
222
+
223
+ num_nodes = int(x.size(0))
224
+ if getattr(conv, "normalize", True):
225
+ ei, norm = gcn_norm(
226
+ edge_index,
227
+ edge_weight=edge_weight,
228
+ num_nodes=num_nodes,
229
+ improved=getattr(conv, "improved", False),
230
+ add_self_loops=getattr(conv, "add_self_loops", True),
231
+ flow=getattr(conv, "flow", "source_to_target"),
232
+ )
233
+ else:
234
+ ei = (
235
+ add_self_loops(edge_index, num_nodes=num_nodes)[0]
236
+ if getattr(conv, "add_self_loops", True)
237
+ else edge_index
238
+ )
239
+ norm = (
240
+ edge_weight
241
+ if edge_weight is not None
242
+ else torch.ones(ei.size(1), device=x.device)
243
+ )
244
+ src, dst = ei[0], ei[1]
245
+ W = conv.lin.weight # (out, in)
246
+
247
+ mul_agg = mul @ W # (N, F_in)
248
+
249
+ delta_in = x - x0
250
+ edge_contrib = (mul_agg[dst] * (norm[:, None] * delta_in[src])).sum(dim=-1)
251
+
252
+ mul_src = torch.zeros(num_nodes, delta_in.size(1), device=x.device)
253
+ mul_src.index_add_(0, src, mul_agg[dst] * norm[:, None])
254
+ return mul_src, edge_contrib
255
+
256
+ @staticmethod
257
+ def _rescale_ratio(delta_in: torch.Tensor, delta_out: torch.Tensor, eps: float):
258
+ denom = delta_in.abs()
259
+ safe = denom > eps
260
+ ratio = torch.ones_like(delta_in)
261
+ ratio[safe] = delta_out[safe] / delta_in[safe]
262
+ return ratio
@@ -0,0 +1,219 @@
1
+ from __future__ import annotations
2
+
3
+ from typing import Any
4
+
5
+ import torch
6
+ from torch import nn
7
+
8
+ from ...core.explanation import Explanation
9
+ from ...core.registry import register
10
+ from ..base import ExplanationAlgorithm
11
+
12
+ _ACTIVATION_GATES = (
13
+ nn.ReLU,
14
+ nn.ReLU6,
15
+ nn.LeakyReLU,
16
+ )
17
+
18
+
19
+ @register("gnn_lrp", "gnn-lrp", "lrp")
20
+ class GNNGatedLRP(ExplanationAlgorithm):
21
+ """GNN-LRP (Layer-wise Relevance Propagation for GNNs).
22
+
23
+ Propagates the relevance from the target-class logit backwards, layer by
24
+ layer, redistributing it according to the positive contributions of each
25
+ neuron (LRP-0 / z+ rules). For each `GCNConv` the relevance is split into
26
+ two steps: (a) the linear transform `W` over the aggregated features and (b)
27
+ the convolution, attributing relevance to the neighboring nodes/edges in
28
+ proportion to their contribution to the message-passing step (GCN norm
29
+ included). Supports GCN architectures (`GCNConv` + `ReLU` + `Linear`).
30
+
31
+ The resulting relevance is non-negative (positive rules) and is returned as
32
+ `node_importance` (sum per node) and `edge_importance` (per directed edge,
33
+ aligned with the `edge_index` indices).
34
+ """
35
+
36
+ def __init__(
37
+ self,
38
+ eps: float = 1e-6,
39
+ normalize: bool = False,
40
+ node_mask_type: str | None = None,
41
+ ):
42
+ self.eps = float(eps)
43
+ self.normalize = bool(normalize)
44
+ self.node_mask_type = node_mask_type
45
+
46
+ def explain(
47
+ self,
48
+ backend: Any,
49
+ model: Any,
50
+ data: Any,
51
+ index: int | list[int] | torch.Tensor | None = None,
52
+ target_class: int | None = None,
53
+ **kwargs,
54
+ ) -> Explanation:
55
+ model.eval()
56
+ x = backend.node_features(data).detach().requires_grad_(False)
57
+ edge_index = backend.edge_index(data)
58
+ edge_weight = backend.edge_weight(data)
59
+ num_nodes = int(x.size(0))
60
+
61
+ nodes = self._to_node_ids(index, num_nodes)
62
+ order: list[tuple[nn.Module, tuple[Any, ...]]] = []
63
+
64
+ def _pre(module: nn.Module, args: tuple[Any, ...]):
65
+ order.append((module, args))
66
+
67
+ handles = [
68
+ module.register_forward_pre_hook(_pre)
69
+ for module in model.modules()
70
+ if module is not model
71
+ ]
72
+
73
+ out = backend.forward(model, x, edge_index, edge_weight=edge_weight)
74
+ for handle in handles:
75
+ handle.remove()
76
+
77
+ logits = out
78
+ if logits.dim() != 2:
79
+ raise ValueError(
80
+ "GNN-LRP requiere predicciones node-level (logits (N, C))."
81
+ )
82
+ target = target_class
83
+ if target is None:
84
+ target = int(logits[nodes[0]].argmax().item())
85
+ target_cls = max(0, min(int(target), logits.size(1) - 1))
86
+
87
+ seed = torch.zeros_like(logits)
88
+ seed[nodes, target_cls] = 1.0
89
+ relevance = seed # (N, C)
90
+
91
+ edge_rel = torch.zeros(edge_index.size(1), device=x.device)
92
+ eps = self.eps
93
+
94
+ for module, args in reversed(order):
95
+ if isinstance(module, nn.Linear):
96
+ relevance = self._linear_lrp(module.weight, args[0], relevance, eps)
97
+ elif self._is_gcn(module):
98
+ rel_out, rel_edge = self._conv_lrp(
99
+ module, args[0], edge_index, edge_weight, relevance, eps
100
+ )
101
+ relevance = rel_out
102
+ num_expanded = int(rel_edge.numel())
103
+ if num_expanded >= edge_index.size(1):
104
+ edge_rel = edge_rel + rel_edge[: edge_index.size(1)]
105
+ elif self._is_activation(module):
106
+ gate = (args[0] > 0).to(relevance.dtype)
107
+ relevance = relevance * gate
108
+ else:
109
+ continue
110
+
111
+ node_importance = relevance.sum(dim=-1)
112
+ if self.normalize:
113
+ total = float(node_importance.sum().item())
114
+ if total > 0:
115
+ node_importance = node_importance / total
116
+ total_e = float(edge_rel.sum().item())
117
+ if total_e > 0:
118
+ edge_rel = edge_rel / total_e
119
+
120
+ return Explanation(
121
+ node_importance=node_importance.detach().cpu(),
122
+ edge_importance=edge_rel.detach().cpu(),
123
+ feature_importance=(
124
+ relevance.detach().cpu()
125
+ if self.node_mask_type == "attributes"
126
+ else None
127
+ ),
128
+ prediction_original=logits[nodes[0]].detach().cpu(),
129
+ prediction_explanation=None,
130
+ node_idx=int(nodes[0]) if nodes.shape[0] == 1 else index,
131
+ target_class=target_cls,
132
+ )
133
+
134
+ # ------------------------------------------------------------------ utils
135
+ @staticmethod
136
+ def _to_node_ids(index, num_nodes: int) -> torch.Tensor:
137
+ if index is None:
138
+ return torch.zeros(1, dtype=torch.long)
139
+ if isinstance(index, int):
140
+ return torch.tensor([index], dtype=torch.long)
141
+ idx = torch.as_tensor(index, dtype=torch.long)
142
+ return idx.reshape(-1) if idx.numel() else torch.zeros(1, dtype=torch.long)
143
+
144
+ @staticmethod
145
+ def _is_gcn(module: nn.Module) -> bool:
146
+ try:
147
+ from torch_geometric.nn import GCNConv
148
+
149
+ return isinstance(module, GCNConv)
150
+ except ImportError:
151
+ return False
152
+
153
+ @staticmethod
154
+ def _is_activation(module: nn.Module) -> bool:
155
+ return isinstance(module, _ACTIVATION_GATES)
156
+
157
+ @staticmethod
158
+ def _linear_lrp(
159
+ weight: torch.Tensor,
160
+ x: torch.Tensor,
161
+ r: torch.Tensor,
162
+ eps: float,
163
+ ) -> torch.Tensor:
164
+ """z+ rule (positive LRP-0) for a linear transform y = Wx."""
165
+ wp = weight.clamp(min=0) # (out, in)
166
+ xp = x.clamp(min=0) # (N, in)
167
+ contrib = xp[:, None, :] * wp[None, :, :] # (N, out, in)
168
+ denom = contrib.sum(dim=-1).clamp(min=eps) # (N, out)
169
+ return (contrib / denom[:, :, None] * r[:, :, None]).sum(dim=1) # (N, in)
170
+
171
+ def _conv_lrp(
172
+ self,
173
+ conv: nn.Module,
174
+ x: torch.Tensor,
175
+ edge_index: torch.Tensor,
176
+ edge_weight: torch.Tensor | None,
177
+ r: torch.Tensor,
178
+ eps: float,
179
+ ):
180
+ """Relevance through a GCNConv: linear `W` + messages (GCN norm)."""
181
+ from torch_geometric.nn.conv.gcn_conv import gcn_norm
182
+
183
+ num_nodes = int(x.size(0))
184
+ if getattr(conv, "normalize", True):
185
+ ei, norm = gcn_norm(
186
+ edge_index,
187
+ edge_weight=edge_weight,
188
+ num_nodes=num_nodes,
189
+ improved=getattr(conv, "improved", False),
190
+ add_self_loops=getattr(conv, "add_self_loops", True),
191
+ flow=getattr(conv, "flow", "source_to_target"),
192
+ )
193
+ else:
194
+ if getattr(conv, "add_self_loops", True):
195
+ from torch_geometric.utils import add_self_loops
196
+
197
+ ei = add_self_loops(edge_index, num_nodes=num_nodes)[0]
198
+ else:
199
+ ei = edge_index
200
+ norm = (
201
+ edge_weight
202
+ if edge_weight is not None
203
+ else torch.ones(ei.size(1), device=x.device)
204
+ )
205
+ src, dst = ei[0], ei[1]
206
+
207
+ xp = x.clamp(min=0)
208
+ msg = norm[:, None] * xp[src] # (E, F)
209
+ agg = torch.zeros(num_nodes, x.size(1), device=x.device)
210
+ agg.index_add_(0, dst, norm[:, None] * x[src])
211
+ agg_pos = agg.clamp(min=0)
212
+
213
+ r_agg = self._linear_lrp(conv.lin.weight, agg_pos, r, eps) # (N, F)
214
+
215
+ frac = msg / agg_pos[dst].clamp(min=eps) # (E, F)
216
+ r_msg = (r_agg[dst] * frac).sum(dim=-1) # (E,)
217
+ r_to_src = torch.zeros(num_nodes, x.size(1), device=x.device)
218
+ r_to_src.index_add_(0, src, r_agg[dst] * frac)
219
+ return r_to_src, r_msg
@@ -0,0 +1,3 @@
1
+ from .narrator import Narrator, describe, narrate, summarize
2
+
3
+ __all__ = ["Narrator", "describe", "narrate", "summarize"]
@@ -0,0 +1,185 @@
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)
@@ -0,0 +1,4 @@
1
+ from .interactive import visualize_interactive
2
+ from .static import show, visualize_static
3
+
4
+ __all__ = ["show", "visualize_interactive", "visualize_static"]