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 +55 -0
- autoLRP/backward/__init__.py +5 -0
- autoLRP/backward/analysis.py +423 -0
- autoLRP/backward/engine.py +299 -0
- autoLRP/backward/install.py +850 -0
- autoLRP/backward/lrp_utils.py +353 -0
- autoLRP/backward/resolve.py +51 -0
- autoLRP/backward/rules.py +294 -0
- autoLRP/backward/strategies.py +288 -0
- autoLRP/compat.py +54 -0
- autoLRP/config.py +284 -0
- autoLRP/eval.py +357 -0
- autoLRP/forward/__init__.py +3 -0
- autoLRP/forward/intercept.py +511 -0
- autoLRP/forward/ops.py +176 -0
- autoLRP/recipes.py +150 -0
- autoLRP/tensor.py +111 -0
- autoLRP/utils.py +21 -0
- autolrp-0.1.0.dist-info/METADATA +207 -0
- autolrp-0.1.0.dist-info/RECORD +23 -0
- autolrp-0.1.0.dist-info/WHEEL +5 -0
- autolrp-0.1.0.dist-info/licenses/LICENSE +21 -0
- autolrp-0.1.0.dist-info/top_level.txt +1 -0
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,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
|