Commit 8849837
Add Rational Quadratic Spline Transforms to Normalizing Flows (#291)
* Splines draft
* update keras requirement
* small improvements to error messages
* add rq spline function
* add spline transform
* update searchsorted utils for jax
also add padd util
* update tests
* add assert_allclose util for improved messages
* parametrize transform for flow tests
* update jacobian, jacobian trace, vjp, jvp, and corresponding usages and tests
* fix imports, remove old jacobian and jvp, fix application in free form flow
* improve logdet computation in free form flows
* Fix comparison for symbolic tensors under tf
* Add splines to twomoons notebook
* improve pad utility
* fix missing left edge in spline
* fix inside mask edge case
* explicitly set bias initializer
* add better expand utility
* small clean up, renaming
* fix indexing, fix inside check
* dump
* fix sign of log jacobian for inverse pass in rq spline
* fix parameter splitting for spline transform
* improve readability
* fix scale and shift trailing dimension
* fix inverse pass return value
* correctly choose bins once for each dimension, even for multi-dimensional inputs
* run formatter
* reduce searchsorted log spam
* log backend used at setup
* remove maximum message cache size
* Improve warning message for jax searchsorted
* Fix spline parameter binning for compiled contexts
* update inverse transform same as forward
* Update TwoMoons notebook with splines WIP [skip ci]
* fix spline inverse call for out of bounds values
* Add working splines
---------
Co-authored-by: stefanradev93 <[email protected]>1 parent ef3892e commit 8849837
File tree
33 files changed
+1033
-762
lines changed- bayesflow
- diagnostics/plots
- networks
- consistency_models
- coupling_flow
- couplings
- transforms
- flow_matching/integrators
- free_form_flow
- simulators
- utils
- jacobian_trace
- jacobian
- docsrc/source
- examples
- tests
- test_networks/test_coupling_flow
- utils
33 files changed
+1033
-762
lines changed| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
36 | 36 | | |
37 | 37 | | |
38 | 38 | | |
| 39 | + | |
| 40 | + | |
| 41 | + | |
| 42 | + | |
39 | 43 | | |
40 | 44 | | |
41 | 45 | | |
| |||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
176 | 176 | | |
177 | 177 | | |
178 | 178 | | |
179 | | - | |
| 179 | + | |
180 | 180 | | |
181 | 181 | | |
182 | 182 | | |
| |||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
79 | 79 | | |
80 | 80 | | |
81 | 81 | | |
82 | | - | |
| 82 | + | |
83 | 83 | | |
84 | 84 | | |
85 | 85 | | |
| |||
Lines changed: 1 addition & 1 deletion
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
249 | 249 | | |
250 | 250 | | |
251 | 251 | | |
252 | | - | |
| 252 | + | |
253 | 253 | | |
254 | 254 | | |
255 | 255 | | |
| |||
Lines changed: 1 addition & 1 deletion
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
1 | 1 | | |
2 | | - | |
3 | 2 | | |
4 | 3 | | |
5 | 4 | | |
| |||
24 | 23 | | |
25 | 24 | | |
26 | 25 | | |
| 26 | + | |
27 | 27 | | |
28 | 28 | | |
29 | 29 | | |
| |||
Lines changed: 81 additions & 0 deletions
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
| 1 | + | |
| 2 | + | |
| 3 | + | |
| 4 | + | |
| 5 | + | |
| 6 | + | |
| 7 | + | |
| 8 | + | |
| 9 | + | |
| 10 | + | |
| 11 | + | |
| 12 | + | |
| 13 | + | |
| 14 | + | |
| 15 | + | |
| 16 | + | |
| 17 | + | |
| 18 | + | |
| 19 | + | |
| 20 | + | |
| 21 | + | |
| 22 | + | |
| 23 | + | |
| 24 | + | |
| 25 | + | |
| 26 | + | |
| 27 | + | |
| 28 | + | |
| 29 | + | |
| 30 | + | |
| 31 | + | |
| 32 | + | |
| 33 | + | |
| 34 | + | |
| 35 | + | |
| 36 | + | |
| 37 | + | |
| 38 | + | |
| 39 | + | |
| 40 | + | |
| 41 | + | |
| 42 | + | |
| 43 | + | |
| 44 | + | |
| 45 | + | |
| 46 | + | |
| 47 | + | |
| 48 | + | |
| 49 | + | |
| 50 | + | |
| 51 | + | |
| 52 | + | |
| 53 | + | |
| 54 | + | |
| 55 | + | |
| 56 | + | |
| 57 | + | |
| 58 | + | |
| 59 | + | |
| 60 | + | |
| 61 | + | |
| 62 | + | |
| 63 | + | |
| 64 | + | |
| 65 | + | |
| 66 | + | |
| 67 | + | |
| 68 | + | |
| 69 | + | |
| 70 | + | |
| 71 | + | |
| 72 | + | |
| 73 | + | |
| 74 | + | |
| 75 | + | |
| 76 | + | |
| 77 | + | |
| 78 | + | |
| 79 | + | |
| 80 | + | |
| 81 | + | |
0 commit comments