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,304 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from collections import deque
|
|
4
|
+
from typing import Any
|
|
5
|
+
|
|
6
|
+
import torch
|
|
7
|
+
import torch.nn.functional as F
|
|
8
|
+
|
|
9
|
+
from ...core.explanation import Explanation
|
|
10
|
+
from ...core.registry import register
|
|
11
|
+
from ..base import ExplanationAlgorithm
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def _softmax(logits, c: int) -> float:
|
|
15
|
+
return float(F.softmax(logits[None], dim=-1)[0, c].item())
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
@register("counterfactual", "counterfactual_explainer", "cf")
|
|
19
|
+
class Counterfactual(ExplanationAlgorithm):
|
|
20
|
+
"""Counterfactual explanation: minimal perturbation that changes the prediction.
|
|
21
|
+
|
|
22
|
+
Finds the minimal set of edges (mode='edge') or feature coordinates
|
|
23
|
+
(mode='feature') whose removal/re-scaling makes the node's prediction change
|
|
24
|
+
class (or reach `flip_to`). The search is greedy and deterministic: at each
|
|
25
|
+
step it removes the candidate element that most reduces `P(original class)`;
|
|
26
|
+
if the class does not change within `max_steps` steps it returns the current
|
|
27
|
+
state (prediction unchanged).
|
|
28
|
+
|
|
29
|
+
The returned importance marks the modified elements (edges 0/1, nodes from
|
|
30
|
+
their incidence on removed edges, changed features 0/1), with
|
|
31
|
+
`prediction_explanation` = logits after the perturbation.
|
|
32
|
+
"""
|
|
33
|
+
|
|
34
|
+
def __init__(
|
|
35
|
+
self,
|
|
36
|
+
mode: str = "edge",
|
|
37
|
+
flip_to: int | None = None,
|
|
38
|
+
max_steps: int = 10,
|
|
39
|
+
hops: int = 2,
|
|
40
|
+
eps: float = 0.0,
|
|
41
|
+
seed: int = 0,
|
|
42
|
+
):
|
|
43
|
+
if mode not in ("edge", "feature"):
|
|
44
|
+
raise ValueError("mode debe ser 'edge' o 'feature'.")
|
|
45
|
+
self.mode = mode
|
|
46
|
+
self.flip_to = flip_to
|
|
47
|
+
self.max_steps = int(max_steps)
|
|
48
|
+
self.hops = int(hops)
|
|
49
|
+
self.eps = float(eps)
|
|
50
|
+
torch.manual_seed(seed)
|
|
51
|
+
|
|
52
|
+
def explain(
|
|
53
|
+
self,
|
|
54
|
+
backend: Any,
|
|
55
|
+
model: Any,
|
|
56
|
+
data: Any,
|
|
57
|
+
index: int | list[int] | torch.Tensor | None = None,
|
|
58
|
+
target_class: int | None = None,
|
|
59
|
+
**kwargs,
|
|
60
|
+
) -> Explanation:
|
|
61
|
+
model.eval()
|
|
62
|
+
node = self._single_node(index)
|
|
63
|
+
x = backend.node_features(data).detach().clone()
|
|
64
|
+
edge_index = backend.edge_index(data)
|
|
65
|
+
edge_weight = backend.edge_weight(data)
|
|
66
|
+
num_nodes = int(x.size(0))
|
|
67
|
+
|
|
68
|
+
orig_logits = backend.forward(model, x, edge_index, edge_weight=edge_weight)
|
|
69
|
+
orig_class = int(orig_logits[node].argmax().item())
|
|
70
|
+
flip_class = target_class if self.flip_to is None else self.flip_to
|
|
71
|
+
target_new = None if flip_class is None else int(flip_class)
|
|
72
|
+
|
|
73
|
+
device = x.device
|
|
74
|
+
if self.mode == "edge":
|
|
75
|
+
edge_importance, final_logits = self._flip_edges(
|
|
76
|
+
backend,
|
|
77
|
+
model,
|
|
78
|
+
x,
|
|
79
|
+
edge_index,
|
|
80
|
+
edge_weight,
|
|
81
|
+
node,
|
|
82
|
+
orig_class,
|
|
83
|
+
target_new,
|
|
84
|
+
device,
|
|
85
|
+
)
|
|
86
|
+
return self._build(
|
|
87
|
+
backend,
|
|
88
|
+
data,
|
|
89
|
+
node,
|
|
90
|
+
orig_class,
|
|
91
|
+
orig_logits,
|
|
92
|
+
final_logits,
|
|
93
|
+
node_importance=torch.bincount(
|
|
94
|
+
torch.cat(
|
|
95
|
+
[
|
|
96
|
+
edge_index[0][edge_importance.bool()],
|
|
97
|
+
edge_index[1][edge_importance.bool()],
|
|
98
|
+
]
|
|
99
|
+
),
|
|
100
|
+
minlength=num_nodes,
|
|
101
|
+
).float(),
|
|
102
|
+
edge_importance=edge_importance,
|
|
103
|
+
feature_importance=None,
|
|
104
|
+
)
|
|
105
|
+
feature_importance, _final_x, final_logits = self._flip_features(
|
|
106
|
+
backend,
|
|
107
|
+
model,
|
|
108
|
+
x,
|
|
109
|
+
edge_index,
|
|
110
|
+
edge_weight,
|
|
111
|
+
node,
|
|
112
|
+
orig_class,
|
|
113
|
+
target_new,
|
|
114
|
+
device,
|
|
115
|
+
)
|
|
116
|
+
return self._build(
|
|
117
|
+
backend,
|
|
118
|
+
data,
|
|
119
|
+
node,
|
|
120
|
+
orig_class,
|
|
121
|
+
orig_logits,
|
|
122
|
+
final_logits,
|
|
123
|
+
node_importance=torch.zeros(num_nodes).index_fill(
|
|
124
|
+
0, torch.tensor([node], device=device), 1.0
|
|
125
|
+
),
|
|
126
|
+
edge_importance=None,
|
|
127
|
+
feature_importance=feature_importance,
|
|
128
|
+
)
|
|
129
|
+
|
|
130
|
+
# ------------------------------------------------------------------ búsquedas
|
|
131
|
+
def _flip_edges(
|
|
132
|
+
self,
|
|
133
|
+
backend,
|
|
134
|
+
model,
|
|
135
|
+
x,
|
|
136
|
+
edge_index,
|
|
137
|
+
edge_weight,
|
|
138
|
+
node,
|
|
139
|
+
orig_class,
|
|
140
|
+
target_new,
|
|
141
|
+
device,
|
|
142
|
+
):
|
|
143
|
+
num_nodes = int(x.size(0))
|
|
144
|
+
weight = (
|
|
145
|
+
edge_weight.detach().clone()
|
|
146
|
+
if edge_weight is not None
|
|
147
|
+
else torch.ones(edge_index.size(1), device=device)
|
|
148
|
+
)
|
|
149
|
+
candidates = self._candidate_edges(edge_index, num_nodes, node, self.hops)
|
|
150
|
+
removed: list[int] = []
|
|
151
|
+
pred = backend.forward(model, x, edge_index, edge_weight=weight)[node]
|
|
152
|
+
|
|
153
|
+
def flipped(logits, tc):
|
|
154
|
+
pred_cls = int(logits.argmax().item())
|
|
155
|
+
return pred_cls == tc if tc is not None else pred_cls != orig_class
|
|
156
|
+
|
|
157
|
+
for _ in range(self.max_steps):
|
|
158
|
+
if flipped(pred, target_new):
|
|
159
|
+
break
|
|
160
|
+
remaining = [e for e in candidates if e not in removed]
|
|
161
|
+
if not remaining:
|
|
162
|
+
break
|
|
163
|
+
best_e, best_logits, best_p = None, None, None
|
|
164
|
+
with torch.no_grad():
|
|
165
|
+
for e in remaining:
|
|
166
|
+
w = weight.clone()
|
|
167
|
+
w[e] = 0.0
|
|
168
|
+
lg = backend.forward(model, x, edge_index, edge_weight=w)[node]
|
|
169
|
+
p = float(_softmax(lg, orig_class))
|
|
170
|
+
if best_p is None or p < best_p:
|
|
171
|
+
best_e, best_logits, best_p = e, lg, p
|
|
172
|
+
if best_e is None or (
|
|
173
|
+
best_p is not None
|
|
174
|
+
and best_p >= float(_softmax(pred, orig_class)) - self.eps
|
|
175
|
+
):
|
|
176
|
+
break
|
|
177
|
+
removed.append(best_e)
|
|
178
|
+
weight[best_e] = 0.0
|
|
179
|
+
pred = best_logits
|
|
180
|
+
|
|
181
|
+
importance = torch.zeros(edge_index.size(1), device=device)
|
|
182
|
+
importance[removed] = 1.0
|
|
183
|
+
return importance, pred.detach()
|
|
184
|
+
|
|
185
|
+
def _flip_features(
|
|
186
|
+
self,
|
|
187
|
+
backend,
|
|
188
|
+
model,
|
|
189
|
+
x,
|
|
190
|
+
edge_index,
|
|
191
|
+
edge_weight,
|
|
192
|
+
node,
|
|
193
|
+
orig_class,
|
|
194
|
+
target_new,
|
|
195
|
+
device,
|
|
196
|
+
):
|
|
197
|
+
baseline = x.mean(dim=0, keepdim=True)
|
|
198
|
+
x_cur = x.detach().clone()
|
|
199
|
+
changed: list[int] = []
|
|
200
|
+
pred = backend.forward(model, x_cur, edge_index, edge_weight=edge_weight)[node]
|
|
201
|
+
|
|
202
|
+
def flipped(logits, tc):
|
|
203
|
+
pred_cls = int(logits.argmax().item())
|
|
204
|
+
return pred_cls == tc if tc is not None else pred_cls != orig_class
|
|
205
|
+
|
|
206
|
+
for _ in range(self.max_steps):
|
|
207
|
+
if flipped(pred, target_new):
|
|
208
|
+
break
|
|
209
|
+
best_c, best_logits, best_p = None, None, None
|
|
210
|
+
with torch.no_grad():
|
|
211
|
+
for c in range(x.size(1)):
|
|
212
|
+
if c in changed:
|
|
213
|
+
continue
|
|
214
|
+
xn = x_cur.clone()
|
|
215
|
+
xn[node, c] = baseline[0, c]
|
|
216
|
+
lg = backend.forward(
|
|
217
|
+
model, xn, edge_index, edge_weight=edge_weight
|
|
218
|
+
)[node]
|
|
219
|
+
p = float(_softmax(lg, orig_class))
|
|
220
|
+
if best_p is None or p < best_p:
|
|
221
|
+
best_c, best_logits, best_p = c, lg, p
|
|
222
|
+
if best_c is None or (
|
|
223
|
+
best_p is not None
|
|
224
|
+
and best_p >= float(_softmax(pred, orig_class)) - self.eps
|
|
225
|
+
):
|
|
226
|
+
break
|
|
227
|
+
changed.append(best_c)
|
|
228
|
+
x_cur[node, best_c] = baseline[0, best_c]
|
|
229
|
+
pred = best_logits
|
|
230
|
+
|
|
231
|
+
importance = torch.zeros(x.size(1), device=device)
|
|
232
|
+
importance[changed] = 1.0
|
|
233
|
+
return importance, x_cur, pred.detach()
|
|
234
|
+
|
|
235
|
+
# ------------------------------------------------------------------ utilidades
|
|
236
|
+
@staticmethod
|
|
237
|
+
def _single_node(index) -> int:
|
|
238
|
+
if index is None:
|
|
239
|
+
return 0
|
|
240
|
+
if isinstance(index, (list, tuple)):
|
|
241
|
+
index = index[0]
|
|
242
|
+
return int(torch.as_tensor(index).reshape(-1)[0].item())
|
|
243
|
+
|
|
244
|
+
@staticmethod
|
|
245
|
+
def _candidate_edges(edge_index, num_nodes: int, node: int, hops: int) -> list[int]:
|
|
246
|
+
adj: dict[int, list[int]] = {}
|
|
247
|
+
for u, v in zip(edge_index[0].tolist(), edge_index[1].tolist()):
|
|
248
|
+
adj.setdefault(u, []).append(v)
|
|
249
|
+
distance = {node: 0}
|
|
250
|
+
queue = deque([node])
|
|
251
|
+
while queue:
|
|
252
|
+
cur = queue.popleft()
|
|
253
|
+
if distance[cur] >= hops:
|
|
254
|
+
continue
|
|
255
|
+
for nb in adj.get(cur, []):
|
|
256
|
+
if nb not in distance:
|
|
257
|
+
distance[nb] = distance[cur] + 1
|
|
258
|
+
queue.append(nb)
|
|
259
|
+
near = set(distance)
|
|
260
|
+
edges = []
|
|
261
|
+
for e in range(edge_index.size(1)):
|
|
262
|
+
u = int(edge_index[0, e].item())
|
|
263
|
+
v = int(edge_index[1, e].item())
|
|
264
|
+
if u in near or v in near:
|
|
265
|
+
edges.append(e)
|
|
266
|
+
return edges
|
|
267
|
+
|
|
268
|
+
def _build(
|
|
269
|
+
self,
|
|
270
|
+
backend,
|
|
271
|
+
data,
|
|
272
|
+
node,
|
|
273
|
+
orig_class,
|
|
274
|
+
orig_logits,
|
|
275
|
+
final_logits,
|
|
276
|
+
node_importance,
|
|
277
|
+
edge_importance,
|
|
278
|
+
feature_importance,
|
|
279
|
+
) -> Explanation:
|
|
280
|
+
metadata = {
|
|
281
|
+
"backend": backend,
|
|
282
|
+
"backing_data": data,
|
|
283
|
+
"counterfactual": True,
|
|
284
|
+
"original_class": orig_class,
|
|
285
|
+
}
|
|
286
|
+
return Explanation(
|
|
287
|
+
node_importance=node_importance.detach().cpu(),
|
|
288
|
+
edge_importance=edge_importance.detach().cpu()
|
|
289
|
+
if edge_importance is not None
|
|
290
|
+
else None,
|
|
291
|
+
feature_importance=(
|
|
292
|
+
feature_importance.detach().cpu()
|
|
293
|
+
if feature_importance is not None
|
|
294
|
+
else None
|
|
295
|
+
),
|
|
296
|
+
prediction_original=orig_logits[node].detach().cpu(),
|
|
297
|
+
prediction_explanation=final_logits.detach().cpu(),
|
|
298
|
+
node_idx=node,
|
|
299
|
+
target_class=int(final_logits.argmax().item())
|
|
300
|
+
if final_logits.numel()
|
|
301
|
+
else None,
|
|
302
|
+
metadata=metadata,
|
|
303
|
+
mask_threshold=0.5,
|
|
304
|
+
)
|
|
@@ -0,0 +1,141 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from collections import defaultdict
|
|
4
|
+
from typing import Any
|
|
5
|
+
|
|
6
|
+
import torch
|
|
7
|
+
|
|
8
|
+
from ...core.explanation import Explanation
|
|
9
|
+
from ...core.registry import register
|
|
10
|
+
from ..base import ExplanationAlgorithm
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
@register("graph_lime", "glime", "gl")
|
|
14
|
+
class GraphLIME(ExplanationAlgorithm):
|
|
15
|
+
"""GraphLIME: feature attribution via weighted local regression.
|
|
16
|
+
|
|
17
|
+
Fits a linear regression (ridge, closed form) over the k-hop neighbors'
|
|
18
|
+
features, weighting each neighbor by its similarity to the target node's
|
|
19
|
+
feature (Gaussian kernel). The coefficients explain the probability
|
|
20
|
+
(softmax) of the target class; node importance matches the kernel
|
|
21
|
+
similarity.
|
|
22
|
+
"""
|
|
23
|
+
|
|
24
|
+
def __init__(
|
|
25
|
+
self,
|
|
26
|
+
hops: int = 2,
|
|
27
|
+
lambda_: float = 1.0,
|
|
28
|
+
sigma: float | None = None,
|
|
29
|
+
normalize: bool = True,
|
|
30
|
+
**kwargs,
|
|
31
|
+
):
|
|
32
|
+
self.hops = hops
|
|
33
|
+
self.lambda_ = lambda_
|
|
34
|
+
self.sigma = sigma
|
|
35
|
+
self.normalize = normalize
|
|
36
|
+
|
|
37
|
+
def explain(
|
|
38
|
+
self,
|
|
39
|
+
backend: Any,
|
|
40
|
+
model: Any,
|
|
41
|
+
data: Any,
|
|
42
|
+
index: int | torch.Tensor | None = None,
|
|
43
|
+
target_class: int | None = None,
|
|
44
|
+
**kwargs,
|
|
45
|
+
) -> Explanation:
|
|
46
|
+
model.eval()
|
|
47
|
+
x = backend.node_features(data)
|
|
48
|
+
edge_index = backend.edge_index(data)
|
|
49
|
+
num_nodes = int(x.size(0))
|
|
50
|
+
|
|
51
|
+
nodes = self._to_node_ids(index, num_nodes)
|
|
52
|
+
root = int(nodes[0])
|
|
53
|
+
with torch.no_grad():
|
|
54
|
+
logits = backend.forward(model, x, edge_index)
|
|
55
|
+
if logits.dim() != 2:
|
|
56
|
+
raise ValueError(
|
|
57
|
+
"GraphLIME requiere predicciones node-level (logits (N, C))."
|
|
58
|
+
)
|
|
59
|
+
target = target_class
|
|
60
|
+
if target is None:
|
|
61
|
+
target = int(logits[root].argmax().item())
|
|
62
|
+
target = max(0, min(int(target), logits.size(1) - 1))
|
|
63
|
+
|
|
64
|
+
probs = torch.softmax(logits, dim=-1)[:, target].detach().cpu()
|
|
65
|
+
|
|
66
|
+
neighbors = self._khop_neighbors(edge_index, root, self.hops)
|
|
67
|
+
if root not in neighbors:
|
|
68
|
+
neighbors.append(root)
|
|
69
|
+
nb = torch.tensor(neighbors, dtype=torch.long)
|
|
70
|
+
x_nb = x[nb].detach().cpu()
|
|
71
|
+
y = probs[nb]
|
|
72
|
+
|
|
73
|
+
dists = torch.norm(x_nb - x[root : root + 1].detach().cpu(), dim=-1)
|
|
74
|
+
sigma = self.sigma
|
|
75
|
+
if sigma is None:
|
|
76
|
+
tail = dists[1:]
|
|
77
|
+
sigma = float(tail.mean().item()) if tail.numel() else 1.0
|
|
78
|
+
sigma = max(sigma, 1e-4)
|
|
79
|
+
weights = torch.exp(-(dists**2) / (2.0 * sigma**2))
|
|
80
|
+
|
|
81
|
+
coef = self._ridge(x_nb, y, weights, self.lambda_)
|
|
82
|
+
|
|
83
|
+
node_importance = torch.zeros(num_nodes)
|
|
84
|
+
node_importance[nb] = weights
|
|
85
|
+
node_importance = node_importance.cpu()
|
|
86
|
+
|
|
87
|
+
if self.normalize and coef.numel():
|
|
88
|
+
denom = coef.abs().max()
|
|
89
|
+
if denom > 1e-12:
|
|
90
|
+
coef = coef / denom
|
|
91
|
+
|
|
92
|
+
return Explanation(
|
|
93
|
+
node_importance=node_importance,
|
|
94
|
+
edge_importance=None,
|
|
95
|
+
feature_importance=coef,
|
|
96
|
+
prediction_original=logits[root].detach().cpu(),
|
|
97
|
+
prediction_explanation=None,
|
|
98
|
+
node_idx=root,
|
|
99
|
+
target_class=target,
|
|
100
|
+
metadata={"neighborhood": neighbors},
|
|
101
|
+
)
|
|
102
|
+
|
|
103
|
+
@staticmethod
|
|
104
|
+
def _ridge(
|
|
105
|
+
x: torch.Tensor, y: torch.Tensor, weights: torch.Tensor, lambda_: float
|
|
106
|
+
) -> torch.Tensor:
|
|
107
|
+
ones = torch.ones(x.size(0), 1)
|
|
108
|
+
X = torch.cat([x, ones], dim=-1).double()
|
|
109
|
+
yw = (y * weights).double()
|
|
110
|
+
Xw = X * weights.unsqueeze(-1).double()
|
|
111
|
+
n_features = X.size(1)
|
|
112
|
+
gram = Xw.t() @ X + lambda_ * torch.eye(n_features, dtype=torch.double)
|
|
113
|
+
beta = torch.linalg.solve(gram, X.t() @ yw)
|
|
114
|
+
return beta[:-1].float()
|
|
115
|
+
|
|
116
|
+
@staticmethod
|
|
117
|
+
def _khop_neighbors(edge_index: torch.Tensor, node: int, hops: int) -> list[int]:
|
|
118
|
+
adj: dict[int, set[int]] = defaultdict(set)
|
|
119
|
+
src = edge_index[0].tolist()
|
|
120
|
+
dst = edge_index[1].tolist()
|
|
121
|
+
for s, d in zip(src, dst):
|
|
122
|
+
adj[s].add(d)
|
|
123
|
+
adj[d].add(s)
|
|
124
|
+
seen = {node}
|
|
125
|
+
frontier = {node}
|
|
126
|
+
for _ in range(hops):
|
|
127
|
+
nxt: set[int] = set()
|
|
128
|
+
for n in frontier:
|
|
129
|
+
nxt |= adj.get(n, set())
|
|
130
|
+
frontier = nxt - seen
|
|
131
|
+
seen |= frontier
|
|
132
|
+
return sorted(seen)
|
|
133
|
+
|
|
134
|
+
@staticmethod
|
|
135
|
+
def _to_node_ids(index, num_nodes: int) -> torch.Tensor:
|
|
136
|
+
if index is None:
|
|
137
|
+
return torch.zeros(1, dtype=torch.long)
|
|
138
|
+
if isinstance(index, int):
|
|
139
|
+
return torch.tensor([index], dtype=torch.long)
|
|
140
|
+
idx = torch.as_tensor(index, dtype=torch.long)
|
|
141
|
+
return idx.reshape(-1) if idx.numel() else torch.zeros(1, dtype=torch.long)
|
|
File without changes
|
|
@@ -0,0 +1,110 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
import torch
|
|
6
|
+
|
|
7
|
+
from ...core.explanation import Explanation
|
|
8
|
+
from ...core.registry import register
|
|
9
|
+
from ..base import ExplanationAlgorithm
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
@register("grad_x_input", "gradient_x_input", "gx")
|
|
13
|
+
class GradXInput(ExplanationAlgorithm):
|
|
14
|
+
"""Gradient x Input: attribution as gradient scaled by the activation.
|
|
15
|
+
|
|
16
|
+
The importance of each feature (and of each edge, if the backend supports
|
|
17
|
+
edge weights) is the gradient of the target-class logit multiplied by the
|
|
18
|
+
input-baseline difference (zero baseline by default). Node importance is the
|
|
19
|
+
sum of `abs(grad * Δx)` over features.
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
graph_level = True
|
|
23
|
+
|
|
24
|
+
def __init__(
|
|
25
|
+
self,
|
|
26
|
+
baseline: str = "zero",
|
|
27
|
+
edge_grads: bool = True,
|
|
28
|
+
node_mask_type: str | None = "attributes",
|
|
29
|
+
**kwargs,
|
|
30
|
+
):
|
|
31
|
+
self.baseline_name = baseline
|
|
32
|
+
self.edge_grads = edge_grads
|
|
33
|
+
self.node_mask_type = node_mask_type
|
|
34
|
+
|
|
35
|
+
def explain(
|
|
36
|
+
self,
|
|
37
|
+
backend: Any,
|
|
38
|
+
model: Any,
|
|
39
|
+
data: Any,
|
|
40
|
+
index: int | list[int] | torch.Tensor | None = None,
|
|
41
|
+
target_class: int | None = None,
|
|
42
|
+
**kwargs,
|
|
43
|
+
) -> Explanation:
|
|
44
|
+
model.eval()
|
|
45
|
+
x = backend.node_features(data)
|
|
46
|
+
edge_index = backend.edge_index(data)
|
|
47
|
+
num_nodes = int(x.size(0))
|
|
48
|
+
|
|
49
|
+
nodes = self._to_node_ids(index, num_nodes)
|
|
50
|
+
with torch.no_grad():
|
|
51
|
+
logits = backend.forward(model, x, edge_index)
|
|
52
|
+
if logits.dim() != 2:
|
|
53
|
+
raise ValueError(
|
|
54
|
+
"GradXInput requiere predicciones node-level (logits (N, C))."
|
|
55
|
+
)
|
|
56
|
+
target = target_class
|
|
57
|
+
if target is None:
|
|
58
|
+
target = int(logits[nodes[0]].argmax().item())
|
|
59
|
+
target_cls = max(0, min(int(target), logits.size(1) - 1))
|
|
60
|
+
|
|
61
|
+
x_in = x.detach().clone().requires_grad_(True)
|
|
62
|
+
ew = None
|
|
63
|
+
compute_edge = self.edge_grads and backend.supports_edge_weight(model)
|
|
64
|
+
if compute_edge:
|
|
65
|
+
ew = torch.ones(
|
|
66
|
+
edge_index.size(1), dtype=torch.float32, device=x.device
|
|
67
|
+
).requires_grad_(True)
|
|
68
|
+
|
|
69
|
+
out = backend.forward(model, x_in, edge_index, edge_weight=ew)
|
|
70
|
+
if out.dim() != 2:
|
|
71
|
+
raise ValueError("GradXInput requiere predicciones node-level.")
|
|
72
|
+
score = out[nodes, target_cls].sum()
|
|
73
|
+
model.zero_grad()
|
|
74
|
+
score.backward()
|
|
75
|
+
|
|
76
|
+
grad_x = x_in.grad.detach()
|
|
77
|
+
baseline = torch.zeros_like(x_in)
|
|
78
|
+
contrib = grad_x * (x_in - baseline)
|
|
79
|
+
|
|
80
|
+
node_importance = contrib.abs().sum(dim=-1)
|
|
81
|
+
feature_importance = contrib.detach()
|
|
82
|
+
|
|
83
|
+
edge_importance = None
|
|
84
|
+
if compute_edge and ew is not None and ew.grad is not None:
|
|
85
|
+
edge_importance = ew.grad.detach().abs()
|
|
86
|
+
|
|
87
|
+
return Explanation(
|
|
88
|
+
node_importance=node_importance.cpu(),
|
|
89
|
+
edge_importance=(
|
|
90
|
+
edge_importance.cpu() if edge_importance is not None else None
|
|
91
|
+
),
|
|
92
|
+
feature_importance=(
|
|
93
|
+
feature_importance.cpu()
|
|
94
|
+
if self.node_mask_type == "attributes"
|
|
95
|
+
else None
|
|
96
|
+
),
|
|
97
|
+
prediction_original=logits[nodes[0]].detach().cpu(),
|
|
98
|
+
prediction_explanation=None,
|
|
99
|
+
node_idx=int(nodes[0]) if nodes.shape[0] == 1 else index,
|
|
100
|
+
target_class=target_cls,
|
|
101
|
+
)
|
|
102
|
+
|
|
103
|
+
@staticmethod
|
|
104
|
+
def _to_node_ids(index, num_nodes: int) -> torch.Tensor:
|
|
105
|
+
if index is None:
|
|
106
|
+
return torch.zeros(1, dtype=torch.long)
|
|
107
|
+
if isinstance(index, int):
|
|
108
|
+
return torch.tensor([index], dtype=torch.long)
|
|
109
|
+
idx = torch.as_tensor(index, dtype=torch.long)
|
|
110
|
+
return idx.reshape(-1) if idx.numel() else torch.zeros(1, dtype=torch.long)
|
|
@@ -0,0 +1,117 @@
|
|
|
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
|
+
|
|
13
|
+
@register("guided_backprop", "guided-backprop", "gbp")
|
|
14
|
+
class GuidedBackprop(ExplanationAlgorithm):
|
|
15
|
+
"""Guided Backpropagation: gradients guided by the ReLU mask.
|
|
16
|
+
|
|
17
|
+
During backpropagation the gradient is filtered: it only propagates where
|
|
18
|
+
the ReLU activation was positive (negative gradients are discarded),
|
|
19
|
+
highlighting the features that positively contribute to the class. Temporary
|
|
20
|
+
hooks are registered on the `nn.ReLU` modules; if the model has none, it
|
|
21
|
+
falls back to standard gradients (metadata `guided=False`).
|
|
22
|
+
"""
|
|
23
|
+
|
|
24
|
+
graph_level = True
|
|
25
|
+
|
|
26
|
+
def __init__(self, fallback_to_gradient: bool = True, **kwargs):
|
|
27
|
+
self.fallback_to_gradient = fallback_to_gradient
|
|
28
|
+
|
|
29
|
+
def explain(
|
|
30
|
+
self,
|
|
31
|
+
backend: Any,
|
|
32
|
+
model: Any,
|
|
33
|
+
data: Any,
|
|
34
|
+
index: int | torch.Tensor | None = None,
|
|
35
|
+
target_class: int | None = None,
|
|
36
|
+
**kwargs,
|
|
37
|
+
) -> Explanation:
|
|
38
|
+
model.eval()
|
|
39
|
+
x = backend.node_features(data)
|
|
40
|
+
edge_index = backend.edge_index(data)
|
|
41
|
+
num_nodes = int(x.size(0))
|
|
42
|
+
|
|
43
|
+
nodes = self._to_node_ids(index, num_nodes)
|
|
44
|
+
root = int(nodes[0])
|
|
45
|
+
with torch.no_grad():
|
|
46
|
+
logits = backend.forward(model, x, edge_index)
|
|
47
|
+
if logits.dim() != 2:
|
|
48
|
+
raise ValueError(
|
|
49
|
+
"GuidedBackprop requiere predicciones node-level (logits (N, C))."
|
|
50
|
+
)
|
|
51
|
+
target = target_class
|
|
52
|
+
if target is None:
|
|
53
|
+
target = int(logits[root].argmax().item())
|
|
54
|
+
target = max(0, min(int(target), logits.size(1) - 1))
|
|
55
|
+
|
|
56
|
+
relus = [m for m in model.modules() if isinstance(m, nn.ReLU)]
|
|
57
|
+
hooks, guided = [], False
|
|
58
|
+
if relus and self.fallback_to_gradient:
|
|
59
|
+
guided = True
|
|
60
|
+
for module in relus:
|
|
61
|
+
fw = module.register_forward_hook(self._mask_forward)
|
|
62
|
+
bw = module.register_full_backward_hook(self._guide_backward)
|
|
63
|
+
hooks.extend((fw, bw))
|
|
64
|
+
|
|
65
|
+
try:
|
|
66
|
+
x_in = x.detach().clone().requires_grad_(True)
|
|
67
|
+
out = backend.forward(model, x_in, edge_index)
|
|
68
|
+
if out.dim() != 2:
|
|
69
|
+
raise ValueError("GuidedBackprop requiere predicciones node-level.")
|
|
70
|
+
score = out[root, target]
|
|
71
|
+
model.zero_grad()
|
|
72
|
+
score.backward()
|
|
73
|
+
grad = (
|
|
74
|
+
x_in.grad.detach() if x_in.grad is not None else torch.zeros_like(x_in)
|
|
75
|
+
)
|
|
76
|
+
finally:
|
|
77
|
+
for hook in hooks:
|
|
78
|
+
hook.remove()
|
|
79
|
+
|
|
80
|
+
node_importance = grad.abs().sum(dim=-1)
|
|
81
|
+
return Explanation(
|
|
82
|
+
node_importance=node_importance.cpu(),
|
|
83
|
+
edge_importance=None,
|
|
84
|
+
feature_importance=grad.cpu(),
|
|
85
|
+
prediction_original=logits[root].detach().cpu(),
|
|
86
|
+
prediction_explanation=None,
|
|
87
|
+
node_idx=root,
|
|
88
|
+
target_class=target,
|
|
89
|
+
metadata={"guided": guided},
|
|
90
|
+
)
|
|
91
|
+
|
|
92
|
+
@staticmethod
|
|
93
|
+
def _mask_forward(module, inp, out):
|
|
94
|
+
mask = torch.where(out > 0, torch.ones_like(out), torch.zeros_like(out))
|
|
95
|
+
module._gbp_mask = mask.detach()
|
|
96
|
+
|
|
97
|
+
@staticmethod
|
|
98
|
+
def _guide_backward(module, grad_input, grad_output):
|
|
99
|
+
mask = module._gbp_mask
|
|
100
|
+
if mask is None or not grad_output or grad_output[0] is None:
|
|
101
|
+
return grad_input
|
|
102
|
+
g = grad_output[0]
|
|
103
|
+
if g.shape != mask.shape:
|
|
104
|
+
return grad_input
|
|
105
|
+
guided = (g * mask).clamp(min=0)
|
|
106
|
+
if len(grad_input) == 1:
|
|
107
|
+
return (guided,)
|
|
108
|
+
return grad_input
|
|
109
|
+
|
|
110
|
+
@staticmethod
|
|
111
|
+
def _to_node_ids(index, num_nodes: int) -> torch.Tensor:
|
|
112
|
+
if index is None:
|
|
113
|
+
return torch.zeros(1, dtype=torch.long)
|
|
114
|
+
if isinstance(index, int):
|
|
115
|
+
return torch.tensor([index], dtype=torch.long)
|
|
116
|
+
idx = torch.as_tensor(index, dtype=torch.long)
|
|
117
|
+
return idx.reshape(-1) if idx.numel() else torch.zeros(1, dtype=torch.long)
|