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
|
@@ -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,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)
|