autoLRP 0.1.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.
autoLRP/__init__.py ADDED
@@ -0,0 +1,55 @@
1
+ r"""autoLRP: layer-wise relevance propagation on the autograd graph.
2
+
3
+ Wrap the input, run the model, pick the output scalar, call ``.lrp()``::
4
+
5
+ import autoLRP
6
+ from autoLRP import LRPConfig, BASE
7
+
8
+ x = autoLRP.tensor(image)
9
+ out = model(x)
10
+ out[0, pred].lrp() # BASE: epsilon on linear
11
+ heatmap = x.relevance # families, proportional elsewhere
12
+
13
+ Every rule-bearing node in the graph is addressed by its name without
14
+ the version digit, or by a fact an analyzer attached to it. ``BASE`` is
15
+ the printed starting table; override entries on it::
16
+
17
+ out[0, pred].lrp(config=LRPConfig(rule={**BASE, 'BmmBackward': 'uniform'}))
18
+ out[0, pred].lrp(config=LRPConfig(attn='attnlrp'))
19
+
20
+ Anything that is not a node name or a registered fact, a rule the
21
+ node's family cannot run, or a node no entry addresses, is an error.
22
+ """
23
+ __version__ = '0.1.0'
24
+
25
+ from .compat import run_selfcheck
26
+ run_selfcheck()
27
+
28
+ from .config import LRPConfig, BASE
29
+ from .tensor import LRPTensor, tensor
30
+ from .forward.intercept import (
31
+ register_rewrite, REWRITES,
32
+ set_decompose_attention, get_decompose_attention, decompose_attention,
33
+ )
34
+ from .backward.strategies import (
35
+ register_installer, installer, merge, match_installer, is_shape_node,
36
+ EXPLICIT_STRATEGY, INSTALLERS,
37
+ )
38
+ from .backward.analysis import (
39
+ register_analyzer, ANALYZERS, node_facts,
40
+ )
41
+ from .backward.engine import graph_lrp, walk, plan_report, explain, explain_summary
42
+ from .recipes import bilrp, clrp
43
+ from . import eval
44
+
45
+ __all__ = [
46
+ 'LRPConfig', 'BASE', 'LRPTensor', 'tensor',
47
+ 'register_rewrite', 'REWRITES',
48
+ 'set_decompose_attention', 'get_decompose_attention',
49
+ 'decompose_attention',
50
+ 'register_installer', 'installer', 'merge', 'match_installer',
51
+ 'is_shape_node', 'EXPLICIT_STRATEGY', 'INSTALLERS',
52
+ 'register_analyzer', 'ANALYZERS', 'node_facts',
53
+ 'graph_lrp', 'walk', 'plan_report', 'explain', 'explain_summary',
54
+ 'bilrp', 'clrp', 'eval',
55
+ ]
@@ -0,0 +1,5 @@
1
+ r"""Backward side: ``rules`` (tables), ``lrp_utils`` (machinery),
2
+ ``install`` (hook installers), ``strategies`` (node name to
3
+ installer), ``analysis`` (facts), ``resolve`` (config key to node),
4
+ ``engine`` (walk and run).
5
+ """
@@ -0,0 +1,423 @@
1
+ r"""Graph analysis: facts about autograd nodes.
2
+
3
+ An analyzer is ``fn(nodes) -> {node: fact}`` over the whole plan, run
4
+ between :func:`~autoLRP.backward.engine.walk` and
5
+ :func:`~autoLRP.backward.engine.execute`. It reads the graph and the
6
+ saved tensors, never module names, and writes ``node.metadata['lrp']``.
7
+ It writes only the fact it is registered as, so a config key can be
8
+ checked against :data:`ANALYZERS` at construction; the value may carry
9
+ data (``statistic_operand`` writes the slot to detach). Facts are
10
+ config keys: ``rule={**BASE, 'my_fact': ...}`` reaches the nodes an
11
+ analyzer registered as ``my_fact`` tagged, before the name entry.
12
+ """
13
+ from typing import Callable, Dict
14
+
15
+ import torch
16
+
17
+
18
+ # ---------------------------------------------------------------------------
19
+ # Registry
20
+ # ---------------------------------------------------------------------------
21
+
22
+ ANALYZERS: Dict[str, Callable] = {}
23
+
24
+
25
+ def register_analyzer(name_or_fn=None):
26
+ r"""Register an analyzer, ``@register_analyzer`` or
27
+ ``@register_analyzer('name')``. It returns ``{node: spec}`` for the
28
+ nodes that carry its fact; ``spec`` is the fact's value (``True`` for a
29
+ plain tag, a slot number for a side), the registered name itself
30
+ (shorthand for a ``True`` tag), or the dict ``{'name': value}`` with
31
+ the registered name. Registering a name again overwrites it.
32
+ """
33
+ def _register(fn, name):
34
+ ANALYZERS[name] = fn
35
+ return fn
36
+ if callable(name_or_fn):
37
+ return _register(name_or_fn, name_or_fn.__name__)
38
+ def _decorator(fn):
39
+ return _register(fn, name_or_fn if name_or_fn is not None
40
+ else fn.__name__)
41
+ return _decorator
42
+
43
+ def run(plan) -> None:
44
+ r"""Execute every registered analyzer once over the walked graph and
45
+ write the returned facts onto ``node.metadata['lrp']`` per the one
46
+ contract documented on :func:`register_analyzer`."""
47
+ if not plan or not ANALYZERS:
48
+ return
49
+ nodes = [n for n, _ in plan]
50
+ for _id, fn in ANALYZERS.items():
51
+ facts = fn(nodes)
52
+ if not facts:
53
+ continue
54
+ for node, spec in facts.items():
55
+ if isinstance(spec, str):
56
+ spec = {spec: True}
57
+ elif not isinstance(spec, dict):
58
+ spec = {_id: spec} # a bare value is the fact's value
59
+ try:
60
+ md = node.metadata.setdefault('lrp', {})
61
+ except (AttributeError, TypeError):
62
+ continue # nodes without a metadata dict (test stand-ins)
63
+ for k, v in spec.items():
64
+ if k != _id:
65
+ raise ValueError(
66
+ f"analyzer '{_id}' wrote fact {k!r}; an analyzer "
67
+ f"writes only the fact it is registered as, so a "
68
+ f"config key can be checked against the registry")
69
+ if k in md and md[k] != v:
70
+ raise ValueError(
71
+ f"fact '{k}' written twice with different values "
72
+ f"({md[k]!r} vs {v!r}); analyzers must not "
73
+ f"conflict on a fact name")
74
+ md[k] = v
75
+
76
+ def node_facts(node) -> dict:
77
+ r"""Return ``node.metadata['lrp']`` or an empty dict. Safe on
78
+ stand-in nodes that lack ``metadata``."""
79
+ md = getattr(node, 'metadata', None)
80
+ if isinstance(md, dict):
81
+ return md.get('lrp', {})
82
+ return {}
83
+
84
+
85
+ # ---------------------------------------------------------------------------
86
+ # Which subgraphs reach a wrapped input
87
+ # ---------------------------------------------------------------------------
88
+
89
+
90
+ def is_weight_leaf(var) -> bool:
91
+ """A leaf that carries no relevance of its own. Relevance flows to
92
+ what the user wrapped with :func:`autoLRP.tensor`; every other leaf,
93
+ an ``nn.Parameter``, a constant the forward intercept made live, a
94
+ plain tensor with ``requires_grad``, is a weight."""
95
+ return not getattr(var, '_lrp_init', False)
96
+
97
+
98
+ def leaf_reach(nodes) -> Dict[int, bool]:
99
+ r"""``reach[id(fn)] = True`` iff ``fn``'s subgraph contains a wrapped
100
+ input leaf (:func:`is_weight_leaf`). One post-order pass over the
101
+ union of the given subgraphs.
102
+ """
103
+ reach: Dict[int, bool] = {}
104
+ for root in nodes:
105
+ if root is None or id(root) in reach:
106
+ continue
107
+ stack = [(root, False)]
108
+ while stack:
109
+ fn, post = stack.pop()
110
+ if fn is None:
111
+ continue
112
+ fid = id(fn)
113
+ if post:
114
+ reach[fid] = any(
115
+ reach.get(id(p), False)
116
+ for p in parents(fn, skip_aliases=False)
117
+ if p is not None)
118
+ continue
119
+ if fid in reach:
120
+ continue
121
+ if 'AccumulateGrad' in fn.name():
122
+ var = getattr(fn, 'variable', None)
123
+ reach[fid] = not is_weight_leaf(var)
124
+ continue
125
+ reach[fid] = False # placeholder; fixed on post-visit
126
+ stack.append((fn, True))
127
+ for p in parents(fn, skip_aliases=False):
128
+ if p is not None and id(p) not in reach:
129
+ stack.append((p, False))
130
+ return reach
131
+
132
+
133
+ def reaches_input(fn) -> bool:
134
+ """``True`` iff ``fn`` reaches a wrapped input leaf. A parameter or a
135
+ constant (including a constant made live by the forward intercept)
136
+ does not; only the path from the user's ``tensor(...)`` does."""
137
+ return _reaches_input_avoiding(fn, None)
138
+
139
+
140
+ def _reaches_input_avoiding(start_fn, forbidden_fn) -> bool:
141
+ r"""``True`` iff ``start_fn`` reaches an input leaf without passing
142
+ through ``forbidden_fn``.
143
+ """
144
+ if start_fn is None:
145
+ return False
146
+ seen = set()
147
+ stack = [start_fn]
148
+ while stack:
149
+ fn = stack.pop()
150
+ if fn is None or fn is forbidden_fn:
151
+ continue
152
+ fid = id(fn)
153
+ if fid in seen:
154
+ continue
155
+ seen.add(fid)
156
+ if 'AccumulateGrad' in fn.name():
157
+ var = getattr(fn, 'variable', None)
158
+ if not is_weight_leaf(var):
159
+ return True
160
+ continue # parameter leaf: keep searching
161
+ for parent in parents(fn, skip_aliases=False):
162
+ if parent is not None and parent is not forbidden_fn:
163
+ stack.append(parent)
164
+ return False
165
+
166
+
167
+ # ---------------------------------------------------------------------------
168
+ # statistic_operand
169
+ # ---------------------------------------------------------------------------
170
+
171
+ _CANDIDATE_FAMILIES = ('MulBackward', 'DivBackward', 'SubBackward')
172
+
173
+
174
+ def _skip_aliases(fn):
175
+ r"""Collapse a chain of ``AliasBackward`` nodes to the first real op;
176
+ the subclass inserts an alias at every op boundary, and anchoring a
177
+ path test on the alias lets a sibling path slip past it.
178
+ """
179
+ while fn is not None and 'AliasBackward' in fn.name():
180
+ nfs = getattr(fn, 'next_functions', ())
181
+ fn = nfs[0][0] if nfs else None
182
+ return fn
183
+
184
+
185
+ def parents(node, skip_aliases: bool = True):
186
+ r"""Producing nodes of ``node``'s operand slots, in slot order; ``None``
187
+ for a slot with no producer. With ``skip_aliases`` (default) chains
188
+ of ``AliasBackward`` are collapsed to the first real op.
189
+ """
190
+ ps = [q for q, _ in getattr(node, 'next_functions', ())]
191
+ return [_skip_aliases(q) for q in ps] if skip_aliases else ps
192
+
193
+
194
+ def operands(node):
195
+ r"""Saved operand tensors ``(a, b)`` of a two-operand node, or
196
+ ``(None, None)``; native ops save ``_saved_self``/``_saved_other``,
197
+ our wrapped ops save through ``saved_tensors``.
198
+ """
199
+ a = getattr(node, '_saved_self', None)
200
+ b = getattr(node, '_saved_other', None)
201
+ if isinstance(a, torch.Tensor) and isinstance(b, torch.Tensor):
202
+ return a, b
203
+ saved = getattr(node, 'saved_tensors', None)
204
+ if saved and len(saved) >= 2 and all(isinstance(t, torch.Tensor) for t in saved[:2]):
205
+ return saved[0], saved[1]
206
+ return None, None
207
+
208
+
209
+ _CANCEL_TOL = 1e-4 # fractional change in z per fractional change in src;
210
+ # true cancellations read ~1e-13, everything else >= ~0.4
211
+
212
+
213
+ def _recompute(node_name, a, b):
214
+ """The node's output, rebuilt from its saved operands."""
215
+ if 'Div' in node_name:
216
+ return a / b
217
+ if 'Sub' in node_name:
218
+ return a - b
219
+ return a * b
220
+
221
+
222
+ def cancels(z, src, tol: float = _CANCEL_TOL) -> bool:
223
+ r"""``True`` when ``z`` does not move as ``src`` is scaled, or as a
224
+ constant is added to every element of ``src`` (RMSNorm cancels the
225
+ factor, mean-subtraction the constant).
226
+
227
+ The movement of every element of ``z`` is read out with two calls:
228
+ a first ``grad`` with an unfixed ``w`` records ``sum_i w[i] dz[i]/dsrc``,
229
+ and differentiating ``(g * v).sum()`` with respect to ``w`` returns all
230
+ ``(v * dz[i]/dsrc).sum()`` at once. ``tol`` is a relative rate: 1e-4
231
+ means a 1 percent change of ``src`` moves ``z`` by under 0.0001
232
+ percent.
233
+ """
234
+ w = torch.zeros_like(z, requires_grad=True)
235
+ g = torch.autograd.grad(z, src, grad_outputs=w, create_graph=True,
236
+ allow_unused=True)[0]
237
+ if g is None:
238
+ return True # z does not depend on src: cannot move
239
+ zn = float(z.detach().norm())
240
+ sn = float(src.detach().norm())
241
+ if zn == 0.0 or sn == 0.0:
242
+ return True # nothing to measure against
243
+ for v in (src, torch.ones_like(src)):
244
+ dz = torch.autograd.grad((g * v).sum(), w, retain_graph=True)[0]
245
+ if float(dz.detach().norm()) / zn / (float(v.detach().norm()) / sn) < tol:
246
+ return True
247
+ return False
248
+
249
+
250
+ def _parameter_slot(a_p, b_p, reach):
251
+ r"""Slot of an operand reaching no input, or ``None`` when both or
252
+ neither does. Relevance sent to such a side lands nowhere.
253
+ """
254
+ a_live = reach.get(id(a_p), False)
255
+ b_live = reach.get(id(b_p), False)
256
+ if a_live == b_live:
257
+ return None
258
+ return 1 if a_live else 0
259
+
260
+
261
+ def _dominated_slot(a_p, b_p):
262
+ r"""Slot of the operand that reaches a model input only through its
263
+ sibling, or ``None`` when both or neither does.
264
+
265
+ Nomination, not verdict: a dominated operand was computed from the
266
+ other one, which says where the zeros would go, not whether zeros
267
+ are right. ``x * x.mean()`` is dominated and cancels nothing.
268
+ """
269
+ a_alone = _reaches_input_avoiding(a_p, b_p)
270
+ b_alone = _reaches_input_avoiding(b_p, a_p)
271
+ if a_alone == b_alone:
272
+ return None
273
+ return 1 if a_alone else 0
274
+
275
+
276
+ @register_analyzer('statistic_operand')
277
+ def statistic_operand(nodes) -> Dict[object, dict]:
278
+ r"""For each two-operand ``Mul``/``Div``/``Sub`` node, name the operand
279
+ that is a statistic of the other; fact value is its slot, 0 or 1.
280
+
281
+ Two grounds, cheapest first. A side that reaches no wrapped input at
282
+ all (parameters, constants): relevance sent there lands nowhere. A
283
+ confirmed cancellation: with both sides from the input, the side that
284
+ still reaches the input when the other is removed is the source, and
285
+ :func:`cancels` measures whether the product discards a property of
286
+ it (scale or level). Two independent operands emit nothing.
287
+
288
+ Only naming: ``BASE`` carries ``{'statistic_operand': ('detach',
289
+ {'by': 'statistic_operand'})}``, which detaches the named side;
290
+ overriding that key changes what runs there.
291
+ """
292
+ reach = leaf_reach(nodes)
293
+ out: Dict[object, dict] = {}
294
+ for n in nodes:
295
+ name = n.name()
296
+ if not any(k in name for k in _CANDIDATE_FAMILIES):
297
+ continue
298
+ ps = parents(n) # aliases collapsed
299
+ if len(ps) < 2 or ps[0] is None or ps[1] is None or ps[0] is ps[1]:
300
+ continue # scalar edge, or x * x
301
+ a, b = operands(n)
302
+ if a is None or b is None:
303
+ continue # native Sub saves nothing
304
+ a_p, b_p = ps[0], ps[1]
305
+
306
+ slot = _parameter_slot(a_p, b_p, reach) # ground 1: role
307
+ if slot is not None:
308
+ out[n] = {'statistic_operand': slot}
309
+ continue
310
+
311
+ slot = _dominated_slot(a_p, b_p) # nomination
312
+ if slot is None:
313
+ continue
314
+ src = a if slot == 1 else b
315
+ if cancels(_recompute(name, a, b), src): # ground 2: measurement
316
+ out[n] = {'statistic_operand': slot}
317
+ return out
318
+
319
+
320
+ # ---------------------------------------------------------------------------
321
+ # input_conv
322
+ # ---------------------------------------------------------------------------
323
+
324
+
325
+ @register_analyzer('input_conv')
326
+ def input_conv(nodes) -> Dict[object, str]:
327
+ r"""Tag a ``ConvolutionBackward`` that reads the model input: its saved
328
+ input has at most four channels and a leaf producer. Used as the key
329
+ for the z-box input rule, ``rule={**BASE, 'input_conv': ('zbox',
330
+ {'low': lo, 'high': hi})}``.
331
+ """
332
+ out: Dict[object, str] = {}
333
+ for n in nodes:
334
+ if 'ConvolutionBackward' not in n.name():
335
+ continue
336
+ saved_inp = getattr(n, '_saved_input', None)
337
+ if saved_inp is None or saved_inp.ndim < 4 or saved_inp.shape[1] > 4:
338
+ continue
339
+ for parent_fn in parents(n, skip_aliases=False):
340
+ if parent_fn is not None and 'AccumulateGrad' in parent_fn.name():
341
+ out[n] = 'input_conv'
342
+ break
343
+ return out
344
+
345
+
346
+ # ---------------------------------------------------------------------------
347
+ # weights_operand
348
+ # ---------------------------------------------------------------------------
349
+
350
+ _AVERAGE_TOL = 1e-4
351
+
352
+
353
+ def is_weighted_average(m, tol: float = _AVERAGE_TOL) -> bool:
354
+ r"""``True`` when each row of ``m`` holds the weights of a weighted
355
+ average: no weight negative, each row totalling 1. In ``m @ b`` such an
356
+ ``m`` only picks points among the rows of ``b``, so everything in the
357
+ product came from ``b``. The row axis is the one the multiply
358
+ contracts; :func:`weights_operand` asks the question per operand and
359
+ transposes for the second one. Reading ``m`` assumes it holds still
360
+ while ``b`` moves, which fails when ``m`` is computed from ``b``
361
+ (``softmax(V @ V.mT) @ V``); :func:`weights_operand` closes that case
362
+ with an independence probe.
363
+ """
364
+ m = m.detach()
365
+ return (float(m.min()) >= -tol
366
+ and float((m.sum(-1) - 1.0).abs().max()) < tol)
367
+
368
+
369
+ @register_analyzer('weights_operand')
370
+ def weights_operand(nodes) -> Dict[object, dict]:
371
+ r"""For each ``BmmBackward`` with both operands from the input, name the
372
+ operand holding the weights of a weighted average
373
+ (:func:`is_weighted_average`, asked per operand); fact value is its
374
+ slot. A row-stochastic operand that the other operand depends on
375
+ emits no fact. Only naming: ``('detach', {'by': 'weights_operand'})``
376
+ in the config is what detaches it, so ``bmm(A, V)`` and
377
+ ``bmm(V.mT, A.mT)`` receive the same attribution.
378
+ """
379
+ out: Dict[object, dict] = {}
380
+ reach = leaf_reach(nodes)
381
+ for n in nodes:
382
+ if 'BmmBackward' not in n.name():
383
+ continue
384
+ ps = parents(n, skip_aliases=False)
385
+ if len(ps) < 2 or ps[0] is None or ps[1] is None or ps[0] is ps[1]:
386
+ continue
387
+ # Both operands must come from the input: with one, the node is
388
+ # a linear layer whose weight is the other operand, and there is
389
+ # no role to name.
390
+ if not (reach.get(id(ps[0]), False) and reach.get(id(ps[1]), False)):
391
+ continue
392
+ a = getattr(n, '_saved_self', None)
393
+ b = getattr(n, '_saved_mat2', None)
394
+ if not isinstance(a, torch.Tensor) or not isinstance(b, torch.Tensor) \
395
+ or a is b:
396
+ continue
397
+ first = is_weighted_average(a)
398
+ second = is_weighted_average(b.transpose(-2, -1))
399
+ if first == second:
400
+ continue
401
+ # A row-stochastic operand the other operand depends on
402
+ # (softmax(V @ V.mT) @ V) is not "the weights"; emit nothing.
403
+ cand, other = (a, b) if first else (b, a)
404
+ if not _independent(cand, other):
405
+ continue
406
+ out[n] = {'weights_operand': 0 if first else 1}
407
+ return out
408
+
409
+
410
+ def _independent(y, x):
411
+ r"""``True`` when ``y`` does not depend on ``x`` through the graph
412
+ (``torch.autograd.grad`` returns ``None`` under ``allow_unused``).
413
+ Tensors outside each other's graphs, or non-grad tensors, count as
414
+ independent."""
415
+ if not (isinstance(y, torch.Tensor) and isinstance(x, torch.Tensor)
416
+ and y.requires_grad and x.requires_grad):
417
+ return True
418
+ try:
419
+ g = torch.autograd.grad(y.sum(), x, retain_graph=True,
420
+ allow_unused=True)[0]
421
+ except RuntimeError:
422
+ return True
423
+ return g is None