onnx-ad 0.1.0__tar.gz
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.
- onnx_ad-0.1.0/LICENSE +21 -0
- onnx_ad-0.1.0/PKG-INFO +215 -0
- onnx_ad-0.1.0/README.md +192 -0
- onnx_ad-0.1.0/pyproject.toml +35 -0
- onnx_ad-0.1.0/setup.cfg +4 -0
- onnx_ad-0.1.0/src/onnx_ad/__init__.py +20 -0
- onnx_ad-0.1.0/src/onnx_ad/__main__.py +52 -0
- onnx_ad-0.1.0/src/onnx_ad/_build.py +496 -0
- onnx_ad-0.1.0/src/onnx_ad/family.py +63 -0
- onnx_ad-0.1.0/src/onnx_ad/forward.py +65 -0
- onnx_ad-0.1.0/src/onnx_ad/reverse.py +93 -0
- onnx_ad-0.1.0/src/onnx_ad/rules.py +1632 -0
- onnx_ad-0.1.0/src/onnx_ad.egg-info/PKG-INFO +215 -0
- onnx_ad-0.1.0/src/onnx_ad.egg-info/SOURCES.txt +17 -0
- onnx_ad-0.1.0/src/onnx_ad.egg-info/dependency_links.txt +1 -0
- onnx_ad-0.1.0/src/onnx_ad.egg-info/entry_points.txt +2 -0
- onnx_ad-0.1.0/src/onnx_ad.egg-info/requires.txt +5 -0
- onnx_ad-0.1.0/src/onnx_ad.egg-info/top_level.txt +1 -0
- onnx_ad-0.1.0/tests/test_ad.py +1141 -0
onnx_ad-0.1.0/LICENSE
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2026 Joris Gillis
|
|
4
|
+
|
|
5
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
6
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
7
|
+
in the Software without restriction, including without limitation the rights
|
|
8
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
9
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
10
|
+
furnished to do so, subject to the following conditions:
|
|
11
|
+
|
|
12
|
+
The above copyright notice and this permission notice shall be included in all
|
|
13
|
+
copies or substantial portions of the Software.
|
|
14
|
+
|
|
15
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
16
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
17
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
18
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
19
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
20
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
21
|
+
SOFTWARE.
|
onnx_ad-0.1.0/PKG-INFO
ADDED
|
@@ -0,0 +1,215 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: onnx-ad
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: Source-code-transforming automatic differentiation on ONNX graphs
|
|
5
|
+
Author: Joris Gillis
|
|
6
|
+
License-Expression: MIT
|
|
7
|
+
Project-URL: Homepage, https://github.com/yacoda/onnx-ad
|
|
8
|
+
Project-URL: Repository, https://github.com/yacoda/onnx-ad
|
|
9
|
+
Project-URL: Issues, https://github.com/yacoda/onnx-ad/issues
|
|
10
|
+
Keywords: onnx,automatic-differentiation,jacobian,adjoint,graph-transform
|
|
11
|
+
Classifier: Development Status :: 3 - Alpha
|
|
12
|
+
Classifier: Programming Language :: Python :: 3
|
|
13
|
+
Classifier: Programming Language :: Python :: 3 :: Only
|
|
14
|
+
Classifier: Topic :: Scientific/Engineering :: Mathematics
|
|
15
|
+
Requires-Python: >=3.9
|
|
16
|
+
Description-Content-Type: text/markdown
|
|
17
|
+
License-File: LICENSE
|
|
18
|
+
Requires-Dist: onnx>=1.14
|
|
19
|
+
Requires-Dist: numpy>=1.21
|
|
20
|
+
Provides-Extra: test
|
|
21
|
+
Requires-Dist: onnxruntime>=1.16; extra == "test"
|
|
22
|
+
Dynamic: license-file
|
|
23
|
+
|
|
24
|
+
# onnx-ad
|
|
25
|
+
|
|
26
|
+
Automatic differentiation **of ONNX graphs**, by source-code transformation: given a model,
|
|
27
|
+
produce new ONNX models that compute its Jacobian-vector and vector-Jacobian products. Pure
|
|
28
|
+
Python over the `onnx` protobuf, no runtime dependency, no framework in the loop.
|
|
29
|
+
|
|
30
|
+
```sh
|
|
31
|
+
python -m pip install onnx-ad
|
|
32
|
+
```
|
|
33
|
+
|
|
34
|
+
```python
|
|
35
|
+
import onnx
|
|
36
|
+
from onnx_ad import forward, reverse, family
|
|
37
|
+
|
|
38
|
+
model = onnx.load("f.onnx") # inputs x -> outputs y
|
|
39
|
+
onnx.save(forward(model), "fwd_f.onnx") # + fwd_x -> + fwd_y = J . fwd_x
|
|
40
|
+
onnx.save(reverse(model), "adj_f.onnx") # + adj_y -> + adj_x = J^T . adj_y
|
|
41
|
+
|
|
42
|
+
family(model, "generated/f.onnx") # the whole set CasADi discovers
|
|
43
|
+
```
|
|
44
|
+
|
|
45
|
+
A derivative model keeps the original signature as a prefix and appends the seeds, so the
|
|
46
|
+
primal outputs stay available. Any number of seed directions is evaluated in a single pass.
|
|
47
|
+
|
|
48
|
+
## Why differentiate ONNX
|
|
49
|
+
|
|
50
|
+
| Route | What it costs |
|
|
51
|
+
| --- | --- |
|
|
52
|
+
| Differentiate in PyTorch, then export | derivative graphs only for models that came from PyTorch; forward mode goes through `jvp`/`vmap`, which is where export breaks; every derivative order needs another trace |
|
|
53
|
+
| Complex step (`Im f(x + i h v)/h`) | forward mode only, one evaluation per direction, and a convention rather than an identity at piecewise operations |
|
|
54
|
+
| **This** | needs a rule per ONNX operation — but then any ONNX model has derivatives, from any producer, at any order |
|
|
55
|
+
|
|
56
|
+
The third route is the one with no ceiling. A Jacobian-vector-product graph is itself an
|
|
57
|
+
ONNX model, so it can be differentiated again, and consumers need no new capability.
|
|
58
|
+
|
|
59
|
+
## Composition: second derivatives for free
|
|
60
|
+
|
|
61
|
+
The passes keep the primal graph and emit only ordinary ONNX operations, so their own output
|
|
62
|
+
is differentiable. Forward-over-adjoint — the exact-Hessian building block — is the two
|
|
63
|
+
passes composed, with no special casing:
|
|
64
|
+
|
|
65
|
+
```python
|
|
66
|
+
adjoint = reverse(model) # x, adj_y -> adj_x
|
|
67
|
+
hessian = forward(adjoint) # + fwd_x, fwd_adj_y -> + fwd_adj_x
|
|
68
|
+
```
|
|
69
|
+
|
|
70
|
+
Repeated differentiation names itself the way CasADi's `diff_prefix` does: a model that
|
|
71
|
+
already carries `fwd_x` gets `fwd2_`/`nfwd2` next, so `forward(forward(model))` needs no
|
|
72
|
+
arguments.
|
|
73
|
+
|
|
74
|
+
## CasADi
|
|
75
|
+
|
|
76
|
+
The conventions are CasADi's, by default, so an emitted family drops into its ONNX backend
|
|
77
|
+
with nothing to adapt:
|
|
78
|
+
|
|
79
|
+
* **names** — `fwd_<x>`, `adj_<y>`, then `fwd2_`, `adj2_`, from the same rule
|
|
80
|
+
`FunctionInternal::diff_prefix` applies, and seed dimensions `nfwd`, `nadj`, `nfwd2`;
|
|
81
|
+
* **layout** — CasADi reads an ONNX tensor as a matrix (rank 0/1 as a column, rank 2
|
|
82
|
+
directly, higher ranks flattened) and wants the seeds of an `r`-by-`c` value as one
|
|
83
|
+
`r`-by-`(nseed*c)` matrix. That is what `layout="casadi"` emits. Pass `layout="onnx"` for
|
|
84
|
+
the internal form instead, where the seed count is a trailing axis on the primal's own
|
|
85
|
+
shape;
|
|
86
|
+
* **files** — `family` writes `f.onnx`, `adj_f.onnx` and `fwd_adj_f.onnx`, the
|
|
87
|
+
`<kind>_<filename>` siblings the backend looks for beside a model.
|
|
88
|
+
|
|
89
|
+
```python
|
|
90
|
+
f = casadi.GraphBuilder("generated/f.onnx").create("f")
|
|
91
|
+
f.reverse(1)(x, f(x), w) # from adj_f.onnx
|
|
92
|
+
casadi.hessian(casadi.dot(f(v), w), v) # from fwd_adj_f.onnx
|
|
93
|
+
```
|
|
94
|
+
|
|
95
|
+
`examples/torch_to_casadi.py` exports a PyTorch model's **primal only** and generates the
|
|
96
|
+
rest here; `examples/casadi_side.py` consumes it and checks gradients, Jacobians and exact
|
|
97
|
+
Hessians against PyTorch. Verified against a CasADi build with `WITH_ONNX=ON` and
|
|
98
|
+
`WITH_ONNX_RUNTIME=ON`, with `CASADI_ONNXRUNTIME_LIB` pointing at `libonnxruntime.so`.
|
|
99
|
+
|
|
100
|
+
Two traps when the primal comes from `torch.onnx.export`: pass `external_data=False`, or the
|
|
101
|
+
weights land in a sidecar `f.onnx.data` that CasADi cannot follow (it hands the model to ONNX
|
|
102
|
+
Runtime as bytes); and CasADi needs the *complete* forward signature of the adjoint, which is
|
|
103
|
+
why `family` seeds `adj_y` as well as `x`.
|
|
104
|
+
|
|
105
|
+
Deliberately **not** offered: a `jacobian` pass. CasADi builds dense Jacobians from the
|
|
106
|
+
adjoint itself, and a `jac_` sibling would only duplicate that.
|
|
107
|
+
|
|
108
|
+
## What it emits
|
|
109
|
+
|
|
110
|
+
Two walks over one rule table. Forward carries a tangent per value in graph order; reverse
|
|
111
|
+
walks backwards, accumulating a contribution per value and summing where a value has several
|
|
112
|
+
consumers. A value with no derivative is *absent* rather than zero, which is what keeps the
|
|
113
|
+
emitted graph the size of the primal one — every weight in a network is such a value, and
|
|
114
|
+
costs nothing.
|
|
115
|
+
|
|
116
|
+
Nonlinear rules read the primal tensors they need straight from the primal graph rather than
|
|
117
|
+
recomputing them: the tangent of `Tanh` is `(1 - y*y) * t`, with `y` the tensor the primal
|
|
118
|
+
`Tanh` already produced. The reverse model therefore stays a plain function of `(x, adj_y)` —
|
|
119
|
+
no "uses output" convention, nothing to wire up.
|
|
120
|
+
|
|
121
|
+
Reverse mode's sharp edge is broadcasting: a contribution arrives shaped like the *result*
|
|
122
|
+
and must be summed back over the axes the operand was broadcast along. Where the operand's
|
|
123
|
+
shape is declared this is a static axis list; where it is symbolic the axes are computed at
|
|
124
|
+
run time, so a dynamic batch dimension survives.
|
|
125
|
+
|
|
126
|
+
### Operations with rules
|
|
127
|
+
|
|
128
|
+
**Arithmetic** `Add` `Sub` `Mul` `Div` `Neg` `Pow` `Sum` `Mean` `Identity`
|
|
129
|
+
|
|
130
|
+
**Elementwise** `Exp` `Log` `Sqrt` `Reciprocal` `Abs` `Sign` `Sin` `Cos` `Tan` `Sinh`
|
|
131
|
+
`Cosh` `Asin` `Acos` `Atan` `Asinh` `Acosh` `Atanh` `Erf` `Tanh` `Sigmoid` `Relu`
|
|
132
|
+
`LeakyRelu` `Elu` `Selu` `Celu` `PRelu` `ThresholdedRelu` `Softplus` `Softsign` `Shrink`
|
|
133
|
+
`HardSigmoid` `HardSwish` `Mish` `Gelu` (exact and `tanh`)
|
|
134
|
+
|
|
135
|
+
**Linear algebra** `MatMul` `Gemm` `Conv` (strided, dilated, grouped, 1-D and up; weight and
|
|
136
|
+
bias too)
|
|
137
|
+
|
|
138
|
+
**Shape** `Reshape` `Flatten` `Transpose` `Squeeze` `Unsqueeze` `Expand` `Concat` `Split`
|
|
139
|
+
`Slice` `Pad` `Tile` `Gather` `GatherND` `CumSum`
|
|
140
|
+
|
|
141
|
+
**Reductions** `ReduceSum` `ReduceMean` `ReduceMax` `ReduceMin` `ReduceProd`
|
|
142
|
+
`ReduceLogSumExp` `ReduceL1` `ReduceL2` `ReduceSumSquare`
|
|
143
|
+
|
|
144
|
+
**Selection** `Where` `Clip` `Min` `Max`
|
|
145
|
+
|
|
146
|
+
**Networks** `Softmax` `LogSoftmax` `LayerNormalization` `BatchNormalization` (inference)
|
|
147
|
+
`Dropout` (inference) `Cast` `CastLike`
|
|
148
|
+
|
|
149
|
+
**Zero derivative, and allowed to consume differentiated values** `Shape` `Size` `NonZero`
|
|
150
|
+
`Equal` `Greater` `Less` `GreaterOrEqual` `LessOrEqual` `And` `Or` `Xor` `Not` the bitwise
|
|
151
|
+
family `ArgMax` `ArgMin` `IsNaN` `IsInf` `Floor` `Ceil` `Round` `Hardmax` `OneHot` `Det`
|
|
152
|
+
|
|
153
|
+
An operation without a rule is an error **only when a differentiated value reaches it** — a
|
|
154
|
+
`Resize` on a constant branch is fine.
|
|
155
|
+
|
|
156
|
+
Still missing: pooling (`MaxPool`, `AveragePool`, `GlobalAveragePool`), `ConvTranspose` as a
|
|
157
|
+
primal, `Einsum`, `Resize`, the scatter operations, `LpNormalization`,
|
|
158
|
+
`InstanceNormalization` and `GroupNormalization`. Out of scope for a first release: control
|
|
159
|
+
flow (`Loop`, `Scan`, `If`), sparsity, and training-mode operations.
|
|
160
|
+
|
|
161
|
+
### Against PyTorch
|
|
162
|
+
|
|
163
|
+
Every one of these exports (`torch.onnx.export`, `dynamo=True`, opset 18) differentiates in
|
|
164
|
+
both modes, and the resulting Jacobian matches `torch.autograd.functional.jacobian` to
|
|
165
|
+
float32 precision:
|
|
166
|
+
|
|
167
|
+
| Model | Operations | forward | reverse |
|
|
168
|
+
| --- | --- | ---: | ---: |
|
|
169
|
+
| MLP | `Gemm` `Tanh` | 6e-8 | 6e-8 |
|
|
170
|
+
| GELU MLP | `Gemm` `Erf` `Mul` `Div` `Add` | 2e-7 | 1e-7 |
|
|
171
|
+
| SiLU MLP | `Gemm` `Sigmoid` `Mul` | 8e-8 | 8e-8 |
|
|
172
|
+
| Attention block | `Gemm` `MatMul` `Softmax` `LayerNormalization` `Reshape` `Transpose` | 2e-7 | 2e-7 |
|
|
173
|
+
| Convolutional net | `Conv` `Relu` `Gemm` `Reshape` | 1e-7 | 8e-8 |
|
|
174
|
+
| Indexing and reductions | `GatherND` `Slice` `Clip` `ReduceMax` `ReduceProd` `ReduceL2` `ReduceLogSumExp` | 1e-7 | 1e-7 |
|
|
175
|
+
| GRU cell | `Gemm` `Split` `Sigmoid` `Tanh` | 6e-8 | 6e-8 |
|
|
176
|
+
|
|
177
|
+
## Running the result
|
|
178
|
+
|
|
179
|
+
Reverse mode emits `Transpose` feeding `MatMul`, which ONNX Runtime's extended optimizer
|
|
180
|
+
fuses into `com.microsoft.FusedMatMul` — a kernel registered for `float` only. On a
|
|
181
|
+
double-precision model, load with the fusions off:
|
|
182
|
+
|
|
183
|
+
```python
|
|
184
|
+
options = ort.SessionOptions()
|
|
185
|
+
options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_BASIC
|
|
186
|
+
session = ort.InferenceSession("adj_f.onnx", options)
|
|
187
|
+
```
|
|
188
|
+
|
|
189
|
+
## Command line
|
|
190
|
+
|
|
191
|
+
```sh
|
|
192
|
+
onnx-ad forward f.onnx fwd_f.onnx
|
|
193
|
+
onnx-ad reverse f.onnx adj_f.onnx --inputs x --outputs y
|
|
194
|
+
onnx-ad family f.onnx generated/f.onnx
|
|
195
|
+
```
|
|
196
|
+
|
|
197
|
+
## Testing
|
|
198
|
+
|
|
199
|
+
Every rule is executed through ONNX Runtime and compared against references that know nothing
|
|
200
|
+
of the rule table: central finite differences of the primal model, an analytic Jacobian
|
|
201
|
+
written in numpy, and forward against reverse — `J` and `J^T` come from separate walks over
|
|
202
|
+
separate rules, so their agreement to machine precision is a real check.
|
|
203
|
+
|
|
204
|
+
No PyTorch is involved in the unit tests; comparisons against it belong in an integration
|
|
205
|
+
suite, not here.
|
|
206
|
+
|
|
207
|
+
```sh
|
|
208
|
+
python -m pip install -e ".[test]"
|
|
209
|
+
python -m unittest discover -s tests -v
|
|
210
|
+
```
|
|
211
|
+
|
|
212
|
+
## License
|
|
213
|
+
|
|
214
|
+
MIT. Releasing is documented in [RELEASE.md](RELEASE.md): PyPI Trusted Publishing, so no API
|
|
215
|
+
token exists, and every artifact carries a PEP 740 provenance attestation.
|
onnx_ad-0.1.0/README.md
ADDED
|
@@ -0,0 +1,192 @@
|
|
|
1
|
+
# onnx-ad
|
|
2
|
+
|
|
3
|
+
Automatic differentiation **of ONNX graphs**, by source-code transformation: given a model,
|
|
4
|
+
produce new ONNX models that compute its Jacobian-vector and vector-Jacobian products. Pure
|
|
5
|
+
Python over the `onnx` protobuf, no runtime dependency, no framework in the loop.
|
|
6
|
+
|
|
7
|
+
```sh
|
|
8
|
+
python -m pip install onnx-ad
|
|
9
|
+
```
|
|
10
|
+
|
|
11
|
+
```python
|
|
12
|
+
import onnx
|
|
13
|
+
from onnx_ad import forward, reverse, family
|
|
14
|
+
|
|
15
|
+
model = onnx.load("f.onnx") # inputs x -> outputs y
|
|
16
|
+
onnx.save(forward(model), "fwd_f.onnx") # + fwd_x -> + fwd_y = J . fwd_x
|
|
17
|
+
onnx.save(reverse(model), "adj_f.onnx") # + adj_y -> + adj_x = J^T . adj_y
|
|
18
|
+
|
|
19
|
+
family(model, "generated/f.onnx") # the whole set CasADi discovers
|
|
20
|
+
```
|
|
21
|
+
|
|
22
|
+
A derivative model keeps the original signature as a prefix and appends the seeds, so the
|
|
23
|
+
primal outputs stay available. Any number of seed directions is evaluated in a single pass.
|
|
24
|
+
|
|
25
|
+
## Why differentiate ONNX
|
|
26
|
+
|
|
27
|
+
| Route | What it costs |
|
|
28
|
+
| --- | --- |
|
|
29
|
+
| Differentiate in PyTorch, then export | derivative graphs only for models that came from PyTorch; forward mode goes through `jvp`/`vmap`, which is where export breaks; every derivative order needs another trace |
|
|
30
|
+
| Complex step (`Im f(x + i h v)/h`) | forward mode only, one evaluation per direction, and a convention rather than an identity at piecewise operations |
|
|
31
|
+
| **This** | needs a rule per ONNX operation — but then any ONNX model has derivatives, from any producer, at any order |
|
|
32
|
+
|
|
33
|
+
The third route is the one with no ceiling. A Jacobian-vector-product graph is itself an
|
|
34
|
+
ONNX model, so it can be differentiated again, and consumers need no new capability.
|
|
35
|
+
|
|
36
|
+
## Composition: second derivatives for free
|
|
37
|
+
|
|
38
|
+
The passes keep the primal graph and emit only ordinary ONNX operations, so their own output
|
|
39
|
+
is differentiable. Forward-over-adjoint — the exact-Hessian building block — is the two
|
|
40
|
+
passes composed, with no special casing:
|
|
41
|
+
|
|
42
|
+
```python
|
|
43
|
+
adjoint = reverse(model) # x, adj_y -> adj_x
|
|
44
|
+
hessian = forward(adjoint) # + fwd_x, fwd_adj_y -> + fwd_adj_x
|
|
45
|
+
```
|
|
46
|
+
|
|
47
|
+
Repeated differentiation names itself the way CasADi's `diff_prefix` does: a model that
|
|
48
|
+
already carries `fwd_x` gets `fwd2_`/`nfwd2` next, so `forward(forward(model))` needs no
|
|
49
|
+
arguments.
|
|
50
|
+
|
|
51
|
+
## CasADi
|
|
52
|
+
|
|
53
|
+
The conventions are CasADi's, by default, so an emitted family drops into its ONNX backend
|
|
54
|
+
with nothing to adapt:
|
|
55
|
+
|
|
56
|
+
* **names** — `fwd_<x>`, `adj_<y>`, then `fwd2_`, `adj2_`, from the same rule
|
|
57
|
+
`FunctionInternal::diff_prefix` applies, and seed dimensions `nfwd`, `nadj`, `nfwd2`;
|
|
58
|
+
* **layout** — CasADi reads an ONNX tensor as a matrix (rank 0/1 as a column, rank 2
|
|
59
|
+
directly, higher ranks flattened) and wants the seeds of an `r`-by-`c` value as one
|
|
60
|
+
`r`-by-`(nseed*c)` matrix. That is what `layout="casadi"` emits. Pass `layout="onnx"` for
|
|
61
|
+
the internal form instead, where the seed count is a trailing axis on the primal's own
|
|
62
|
+
shape;
|
|
63
|
+
* **files** — `family` writes `f.onnx`, `adj_f.onnx` and `fwd_adj_f.onnx`, the
|
|
64
|
+
`<kind>_<filename>` siblings the backend looks for beside a model.
|
|
65
|
+
|
|
66
|
+
```python
|
|
67
|
+
f = casadi.GraphBuilder("generated/f.onnx").create("f")
|
|
68
|
+
f.reverse(1)(x, f(x), w) # from adj_f.onnx
|
|
69
|
+
casadi.hessian(casadi.dot(f(v), w), v) # from fwd_adj_f.onnx
|
|
70
|
+
```
|
|
71
|
+
|
|
72
|
+
`examples/torch_to_casadi.py` exports a PyTorch model's **primal only** and generates the
|
|
73
|
+
rest here; `examples/casadi_side.py` consumes it and checks gradients, Jacobians and exact
|
|
74
|
+
Hessians against PyTorch. Verified against a CasADi build with `WITH_ONNX=ON` and
|
|
75
|
+
`WITH_ONNX_RUNTIME=ON`, with `CASADI_ONNXRUNTIME_LIB` pointing at `libonnxruntime.so`.
|
|
76
|
+
|
|
77
|
+
Two traps when the primal comes from `torch.onnx.export`: pass `external_data=False`, or the
|
|
78
|
+
weights land in a sidecar `f.onnx.data` that CasADi cannot follow (it hands the model to ONNX
|
|
79
|
+
Runtime as bytes); and CasADi needs the *complete* forward signature of the adjoint, which is
|
|
80
|
+
why `family` seeds `adj_y` as well as `x`.
|
|
81
|
+
|
|
82
|
+
Deliberately **not** offered: a `jacobian` pass. CasADi builds dense Jacobians from the
|
|
83
|
+
adjoint itself, and a `jac_` sibling would only duplicate that.
|
|
84
|
+
|
|
85
|
+
## What it emits
|
|
86
|
+
|
|
87
|
+
Two walks over one rule table. Forward carries a tangent per value in graph order; reverse
|
|
88
|
+
walks backwards, accumulating a contribution per value and summing where a value has several
|
|
89
|
+
consumers. A value with no derivative is *absent* rather than zero, which is what keeps the
|
|
90
|
+
emitted graph the size of the primal one — every weight in a network is such a value, and
|
|
91
|
+
costs nothing.
|
|
92
|
+
|
|
93
|
+
Nonlinear rules read the primal tensors they need straight from the primal graph rather than
|
|
94
|
+
recomputing them: the tangent of `Tanh` is `(1 - y*y) * t`, with `y` the tensor the primal
|
|
95
|
+
`Tanh` already produced. The reverse model therefore stays a plain function of `(x, adj_y)` —
|
|
96
|
+
no "uses output" convention, nothing to wire up.
|
|
97
|
+
|
|
98
|
+
Reverse mode's sharp edge is broadcasting: a contribution arrives shaped like the *result*
|
|
99
|
+
and must be summed back over the axes the operand was broadcast along. Where the operand's
|
|
100
|
+
shape is declared this is a static axis list; where it is symbolic the axes are computed at
|
|
101
|
+
run time, so a dynamic batch dimension survives.
|
|
102
|
+
|
|
103
|
+
### Operations with rules
|
|
104
|
+
|
|
105
|
+
**Arithmetic** `Add` `Sub` `Mul` `Div` `Neg` `Pow` `Sum` `Mean` `Identity`
|
|
106
|
+
|
|
107
|
+
**Elementwise** `Exp` `Log` `Sqrt` `Reciprocal` `Abs` `Sign` `Sin` `Cos` `Tan` `Sinh`
|
|
108
|
+
`Cosh` `Asin` `Acos` `Atan` `Asinh` `Acosh` `Atanh` `Erf` `Tanh` `Sigmoid` `Relu`
|
|
109
|
+
`LeakyRelu` `Elu` `Selu` `Celu` `PRelu` `ThresholdedRelu` `Softplus` `Softsign` `Shrink`
|
|
110
|
+
`HardSigmoid` `HardSwish` `Mish` `Gelu` (exact and `tanh`)
|
|
111
|
+
|
|
112
|
+
**Linear algebra** `MatMul` `Gemm` `Conv` (strided, dilated, grouped, 1-D and up; weight and
|
|
113
|
+
bias too)
|
|
114
|
+
|
|
115
|
+
**Shape** `Reshape` `Flatten` `Transpose` `Squeeze` `Unsqueeze` `Expand` `Concat` `Split`
|
|
116
|
+
`Slice` `Pad` `Tile` `Gather` `GatherND` `CumSum`
|
|
117
|
+
|
|
118
|
+
**Reductions** `ReduceSum` `ReduceMean` `ReduceMax` `ReduceMin` `ReduceProd`
|
|
119
|
+
`ReduceLogSumExp` `ReduceL1` `ReduceL2` `ReduceSumSquare`
|
|
120
|
+
|
|
121
|
+
**Selection** `Where` `Clip` `Min` `Max`
|
|
122
|
+
|
|
123
|
+
**Networks** `Softmax` `LogSoftmax` `LayerNormalization` `BatchNormalization` (inference)
|
|
124
|
+
`Dropout` (inference) `Cast` `CastLike`
|
|
125
|
+
|
|
126
|
+
**Zero derivative, and allowed to consume differentiated values** `Shape` `Size` `NonZero`
|
|
127
|
+
`Equal` `Greater` `Less` `GreaterOrEqual` `LessOrEqual` `And` `Or` `Xor` `Not` the bitwise
|
|
128
|
+
family `ArgMax` `ArgMin` `IsNaN` `IsInf` `Floor` `Ceil` `Round` `Hardmax` `OneHot` `Det`
|
|
129
|
+
|
|
130
|
+
An operation without a rule is an error **only when a differentiated value reaches it** — a
|
|
131
|
+
`Resize` on a constant branch is fine.
|
|
132
|
+
|
|
133
|
+
Still missing: pooling (`MaxPool`, `AveragePool`, `GlobalAveragePool`), `ConvTranspose` as a
|
|
134
|
+
primal, `Einsum`, `Resize`, the scatter operations, `LpNormalization`,
|
|
135
|
+
`InstanceNormalization` and `GroupNormalization`. Out of scope for a first release: control
|
|
136
|
+
flow (`Loop`, `Scan`, `If`), sparsity, and training-mode operations.
|
|
137
|
+
|
|
138
|
+
### Against PyTorch
|
|
139
|
+
|
|
140
|
+
Every one of these exports (`torch.onnx.export`, `dynamo=True`, opset 18) differentiates in
|
|
141
|
+
both modes, and the resulting Jacobian matches `torch.autograd.functional.jacobian` to
|
|
142
|
+
float32 precision:
|
|
143
|
+
|
|
144
|
+
| Model | Operations | forward | reverse |
|
|
145
|
+
| --- | --- | ---: | ---: |
|
|
146
|
+
| MLP | `Gemm` `Tanh` | 6e-8 | 6e-8 |
|
|
147
|
+
| GELU MLP | `Gemm` `Erf` `Mul` `Div` `Add` | 2e-7 | 1e-7 |
|
|
148
|
+
| SiLU MLP | `Gemm` `Sigmoid` `Mul` | 8e-8 | 8e-8 |
|
|
149
|
+
| Attention block | `Gemm` `MatMul` `Softmax` `LayerNormalization` `Reshape` `Transpose` | 2e-7 | 2e-7 |
|
|
150
|
+
| Convolutional net | `Conv` `Relu` `Gemm` `Reshape` | 1e-7 | 8e-8 |
|
|
151
|
+
| Indexing and reductions | `GatherND` `Slice` `Clip` `ReduceMax` `ReduceProd` `ReduceL2` `ReduceLogSumExp` | 1e-7 | 1e-7 |
|
|
152
|
+
| GRU cell | `Gemm` `Split` `Sigmoid` `Tanh` | 6e-8 | 6e-8 |
|
|
153
|
+
|
|
154
|
+
## Running the result
|
|
155
|
+
|
|
156
|
+
Reverse mode emits `Transpose` feeding `MatMul`, which ONNX Runtime's extended optimizer
|
|
157
|
+
fuses into `com.microsoft.FusedMatMul` — a kernel registered for `float` only. On a
|
|
158
|
+
double-precision model, load with the fusions off:
|
|
159
|
+
|
|
160
|
+
```python
|
|
161
|
+
options = ort.SessionOptions()
|
|
162
|
+
options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_BASIC
|
|
163
|
+
session = ort.InferenceSession("adj_f.onnx", options)
|
|
164
|
+
```
|
|
165
|
+
|
|
166
|
+
## Command line
|
|
167
|
+
|
|
168
|
+
```sh
|
|
169
|
+
onnx-ad forward f.onnx fwd_f.onnx
|
|
170
|
+
onnx-ad reverse f.onnx adj_f.onnx --inputs x --outputs y
|
|
171
|
+
onnx-ad family f.onnx generated/f.onnx
|
|
172
|
+
```
|
|
173
|
+
|
|
174
|
+
## Testing
|
|
175
|
+
|
|
176
|
+
Every rule is executed through ONNX Runtime and compared against references that know nothing
|
|
177
|
+
of the rule table: central finite differences of the primal model, an analytic Jacobian
|
|
178
|
+
written in numpy, and forward against reverse — `J` and `J^T` come from separate walks over
|
|
179
|
+
separate rules, so their agreement to machine precision is a real check.
|
|
180
|
+
|
|
181
|
+
No PyTorch is involved in the unit tests; comparisons against it belong in an integration
|
|
182
|
+
suite, not here.
|
|
183
|
+
|
|
184
|
+
```sh
|
|
185
|
+
python -m pip install -e ".[test]"
|
|
186
|
+
python -m unittest discover -s tests -v
|
|
187
|
+
```
|
|
188
|
+
|
|
189
|
+
## License
|
|
190
|
+
|
|
191
|
+
MIT. Releasing is documented in [RELEASE.md](RELEASE.md): PyPI Trusted Publishing, so no API
|
|
192
|
+
token exists, and every artifact carries a PEP 740 provenance attestation.
|
|
@@ -0,0 +1,35 @@
|
|
|
1
|
+
[build-system]
|
|
2
|
+
requires = ["setuptools>=77"]
|
|
3
|
+
build-backend = "setuptools.build_meta"
|
|
4
|
+
|
|
5
|
+
[project]
|
|
6
|
+
name = "onnx-ad"
|
|
7
|
+
version = "0.1.0"
|
|
8
|
+
description = "Source-code-transforming automatic differentiation on ONNX graphs"
|
|
9
|
+
readme = "README.md"
|
|
10
|
+
requires-python = ">=3.9"
|
|
11
|
+
license = "MIT"
|
|
12
|
+
license-files = ["LICENSE"]
|
|
13
|
+
authors = [{name = "Joris Gillis"}]
|
|
14
|
+
keywords = ["onnx", "automatic-differentiation", "jacobian", "adjoint", "graph-transform"]
|
|
15
|
+
classifiers = [
|
|
16
|
+
"Development Status :: 3 - Alpha",
|
|
17
|
+
"Programming Language :: Python :: 3",
|
|
18
|
+
"Programming Language :: Python :: 3 :: Only",
|
|
19
|
+
"Topic :: Scientific/Engineering :: Mathematics",
|
|
20
|
+
]
|
|
21
|
+
dependencies = ["onnx>=1.14", "numpy>=1.21"]
|
|
22
|
+
|
|
23
|
+
[project.optional-dependencies]
|
|
24
|
+
test = ["onnxruntime>=1.16"]
|
|
25
|
+
|
|
26
|
+
[project.scripts]
|
|
27
|
+
onnx-ad = "onnx_ad.__main__:main"
|
|
28
|
+
|
|
29
|
+
[project.urls]
|
|
30
|
+
Homepage = "https://github.com/yacoda/onnx-ad"
|
|
31
|
+
Repository = "https://github.com/yacoda/onnx-ad"
|
|
32
|
+
Issues = "https://github.com/yacoda/onnx-ad/issues"
|
|
33
|
+
|
|
34
|
+
[tool.setuptools.packages.find]
|
|
35
|
+
where = ["src"]
|
onnx_ad-0.1.0/setup.cfg
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
"""Source-code-transforming automatic differentiation on ONNX graphs.
|
|
2
|
+
|
|
3
|
+
`forward` and `reverse` rewrite an ONNX model into another ONNX model that computes its
|
|
4
|
+
Jacobian-vector or vector-Jacobian products, by applying a rule per operation over the
|
|
5
|
+
protobuf. No framework is in the loop: any ONNX model has derivatives, from any producer,
|
|
6
|
+
and since a derivative model is itself an ONNX model, `forward(reverse(model))` gives
|
|
7
|
+
forward-over-adjoint -- the exact-Hessian building block -- with no special casing.
|
|
8
|
+
|
|
9
|
+
Several seed directions ride in one evaluation, and the conventions are CasADi's own: the
|
|
10
|
+
derivative prefixes and seed dimensions its `diff_prefix` rule would pick (`fwd_`/`nfwd`,
|
|
11
|
+
then `fwd2_`/`nfwd2`), the seed layout its 2-D reading of an ONNX tensor expects, and the
|
|
12
|
+
`<kind>_<filename>` sibling names its ONNX backend discovers. `family` writes the whole set.
|
|
13
|
+
"""
|
|
14
|
+
from ._build import UnsupportedOperator
|
|
15
|
+
from .family import family, sibling
|
|
16
|
+
from .forward import forward
|
|
17
|
+
from .reverse import reverse
|
|
18
|
+
|
|
19
|
+
__all__ = ["forward", "reverse", "family", "sibling", "UnsupportedOperator", "__version__"]
|
|
20
|
+
__version__ = "0.1.0"
|
|
@@ -0,0 +1,52 @@
|
|
|
1
|
+
"""Command line: write the derivative model, or the whole sibling family, of an ONNX model."""
|
|
2
|
+
import argparse
|
|
3
|
+
|
|
4
|
+
import onnx
|
|
5
|
+
|
|
6
|
+
from . import __version__, family, forward, reverse
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def main(argv=None):
|
|
10
|
+
parser = argparse.ArgumentParser(prog="onnx-ad", description=__doc__)
|
|
11
|
+
parser.add_argument("--version", action="version", version=__version__)
|
|
12
|
+
sub = parser.add_subparsers(dest="command", required=True)
|
|
13
|
+
for name, help_text in [
|
|
14
|
+
("forward", "emit a model computing fwd_y = J . fwd_x"),
|
|
15
|
+
("reverse", "emit a model computing adj_x = J^T . adj_y"),
|
|
16
|
+
("family", "write the primal and the sibling derivative models CasADi discovers")]:
|
|
17
|
+
command = sub.add_parser(name, help=help_text)
|
|
18
|
+
command.add_argument("input")
|
|
19
|
+
command.add_argument("output")
|
|
20
|
+
command.add_argument("--inputs", nargs="+",
|
|
21
|
+
help="graph inputs to differentiate [every floating-point one]")
|
|
22
|
+
command.add_argument("--outputs", nargs="+",
|
|
23
|
+
help="graph outputs to differentiate [every floating-point one]")
|
|
24
|
+
if name == "family":
|
|
25
|
+
command.add_argument("--forward-sibling", action="store_true",
|
|
26
|
+
help="also write fwd_<name>, worth it when inputs < outputs")
|
|
27
|
+
command.add_argument("--no-second-order", dest="second_order",
|
|
28
|
+
action="store_false",
|
|
29
|
+
help="skip fwd_adj_<name>, the exact-Hessian sibling")
|
|
30
|
+
else:
|
|
31
|
+
command.add_argument("--prefix", help="derivative-tensor prefix [CasADi's rule]")
|
|
32
|
+
command.add_argument("--dim", help="symbolic seed dimension [CasADi's rule]")
|
|
33
|
+
command.add_argument("--layout", choices=["casadi", "onnx"], default="casadi",
|
|
34
|
+
help="seed layout: CasADi's packed matrix, or a trailing "
|
|
35
|
+
"axis on the primal's own shape [casadi]")
|
|
36
|
+
args = parser.parse_args(argv)
|
|
37
|
+
model = onnx.load(args.input)
|
|
38
|
+
if args.command == "family":
|
|
39
|
+
for path in family(model, args.output, forward_sibling=args.forward_sibling,
|
|
40
|
+
second_order=args.second_order, inputs=args.inputs,
|
|
41
|
+
outputs=args.outputs):
|
|
42
|
+
print(path)
|
|
43
|
+
return
|
|
44
|
+
pass_ = forward if args.command == "forward" else reverse
|
|
45
|
+
result = pass_(model, inputs=args.inputs, outputs=args.outputs, prefix=args.prefix,
|
|
46
|
+
dim=args.dim, layout=args.layout)
|
|
47
|
+
onnx.checker.check_model(result)
|
|
48
|
+
onnx.save(result, args.output)
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
if __name__ == "__main__":
|
|
52
|
+
main()
|