cuthbertlib 0.0.10__tar.gz → 0.0.12__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.
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/PKG-INFO +1 -1
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/enkf/filtering.py +1 -1
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linalg/tria.py +2 -2
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/__init__.py +1 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/autodiff.py +0 -1
- cuthbertlib-0.0.12/cuthbertlib/resampling/no_resampling.py +57 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/pyproject.toml +1 -1
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/.gitignore +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/README.md +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/README.md +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/__init__.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/discrete/__init__.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/discrete/filtering.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/discrete/smoothing.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/enkf/README.md +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/enkf/__init__.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/kalman/README.md +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/kalman/__init__.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/kalman/filtering.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/kalman/generate.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/kalman/sampling.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/kalman/smoothing.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linalg/README.md +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linalg/__init__.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linalg/collect_nans_chol.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linalg/marginal_sqrt_cov.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linalg/symmetric_inv_sqrt.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linearize/README.md +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linearize/__init__.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linearize/log_density.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linearize/moments.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linearize/taylor.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/README.md +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/__init__.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/common.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/cubature.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/gauss_hermite.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/linearize.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/unscented.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/utils.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/README.md +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/adaptive.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/killing.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/multinomial.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/protocols.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/systematic.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/utils.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/README.md +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/__init__.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/ess.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/smoothing/__init__.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/smoothing/exact_sampling.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/smoothing/mcmc.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/smoothing/protocols.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/smoothing/tracing.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/stats/README.md +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/stats/__init__.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/stats/multivariate_normal.py +0 -0
- {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/types.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: cuthbertlib
|
|
3
|
-
Version: 0.0.
|
|
3
|
+
Version: 0.0.12
|
|
4
4
|
Summary: Atomic building blocks for state-space model inference with JAX
|
|
5
5
|
Author-email: Sam Duffield <s@mduffield.com>, Sahel Iqbal <sahel13miqbal@proton.me>, Adrien Corenflos <adrien.corenflos.stats@gmail.com>
|
|
6
6
|
License: Apache-2.0
|
|
@@ -13,7 +13,7 @@ from jax.scipy.linalg import cho_solve
|
|
|
13
13
|
|
|
14
14
|
from cuthbertlib.linalg import collect_nans_chol, tria
|
|
15
15
|
from cuthbertlib.stats import multivariate_normal
|
|
16
|
-
from cuthbertlib.types import Array,
|
|
16
|
+
from cuthbertlib.types import Array, KeyArray, ScalarArray
|
|
17
17
|
|
|
18
18
|
ObservationFn = Callable[[Array], Array]
|
|
19
19
|
DynamicsFn = Callable[[Array, KeyArray], Array]
|
|
@@ -85,8 +85,8 @@ def _tria_jvp(primals, tangents):
|
|
|
85
85
|
K_T = jnp.swapaxes(K, -1, -2)
|
|
86
86
|
|
|
87
87
|
# Solve for lower triangular perturbation dM + dM^T = K + K^T
|
|
88
|
-
|
|
89
|
-
dM = jnp.tril(K + K_T) - K *
|
|
88
|
+
Id = jnp.eye(K.shape[-1], dtype=K.dtype)
|
|
89
|
+
dM = jnp.tril(K + K_T) - K * Id
|
|
90
90
|
|
|
91
91
|
# Compute the null-space part
|
|
92
92
|
dR_null = (jnp.eye(R.shape[-2], dtype=R.dtype) - R @ R_pinv) @ dA @ Q
|
|
@@ -15,7 +15,6 @@ import jax.numpy as jnp
|
|
|
15
15
|
|
|
16
16
|
from cuthbertlib.resampling.protocols import Resampling
|
|
17
17
|
from cuthbertlib.resampling.utils import apply_resampling_indices
|
|
18
|
-
from cuthbertlib.smc.ess import log_ess
|
|
19
18
|
from cuthbertlib.types import Array, ArrayLike, ArrayTree, ArrayTreeLike
|
|
20
19
|
|
|
21
20
|
|
|
@@ -0,0 +1,57 @@
|
|
|
1
|
+
"""No resampling dummy implementation."""
|
|
2
|
+
|
|
3
|
+
from functools import partial
|
|
4
|
+
|
|
5
|
+
from jax import numpy as jnp
|
|
6
|
+
|
|
7
|
+
from cuthbertlib.resampling.protocols import (
|
|
8
|
+
conditional_resampling_decorator,
|
|
9
|
+
resampling_decorator,
|
|
10
|
+
)
|
|
11
|
+
from cuthbertlib.resampling.utils import apply_resampling_indices
|
|
12
|
+
from cuthbertlib.types import (
|
|
13
|
+
Array,
|
|
14
|
+
ArrayLike,
|
|
15
|
+
ArrayTree,
|
|
16
|
+
ArrayTreeLike,
|
|
17
|
+
ScalarArrayLike,
|
|
18
|
+
)
|
|
19
|
+
|
|
20
|
+
_DESCRIPTION = """
|
|
21
|
+
No resampling is performed.
|
|
22
|
+
Useful for factorial SMC where resampling is applied during `join` rather than
|
|
23
|
+
`filter_combine`."""
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
@partial(resampling_decorator, name="No Resampling", desc=_DESCRIPTION)
|
|
27
|
+
def resampling(
|
|
28
|
+
key: Array, logits: ArrayLike, positions: ArrayTreeLike, n: int
|
|
29
|
+
) -> tuple[Array, Array, ArrayTree]:
|
|
30
|
+
logits = jnp.asarray(logits)
|
|
31
|
+
if n != logits.shape[0]:
|
|
32
|
+
raise AssertionError(
|
|
33
|
+
"The number of sampled indices must be equal to the number of "
|
|
34
|
+
"output particles for `No Resampling` resampling."
|
|
35
|
+
)
|
|
36
|
+
return jnp.arange(n), logits, positions
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
@partial(conditional_resampling_decorator, name="No Resampling", desc=_DESCRIPTION)
|
|
40
|
+
def conditional_resampling(
|
|
41
|
+
key: Array,
|
|
42
|
+
logits: ArrayLike,
|
|
43
|
+
positions: ArrayTreeLike,
|
|
44
|
+
n: int,
|
|
45
|
+
pivot_in: ScalarArrayLike,
|
|
46
|
+
pivot_out: ScalarArrayLike,
|
|
47
|
+
) -> tuple[Array, Array, ArrayTree]:
|
|
48
|
+
logits = jnp.asarray(logits)
|
|
49
|
+
if n != logits.shape[0]:
|
|
50
|
+
raise AssertionError(
|
|
51
|
+
"The number of sampled indices must be equal to the number of "
|
|
52
|
+
"output particles for `No Resampling` resampling."
|
|
53
|
+
)
|
|
54
|
+
pivot_in = jnp.asarray(pivot_in, dtype=jnp.int32)
|
|
55
|
+
pivot_out = jnp.asarray(pivot_out, dtype=jnp.int32)
|
|
56
|
+
idx = jnp.arange(n).at[pivot_in].set(pivot_out)
|
|
57
|
+
return idx, logits, apply_resampling_indices(positions, idx)
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|