cuthbertlib 0.0.9__tar.gz → 0.0.11__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.
Files changed (59) hide show
  1. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/PKG-INFO +1 -1
  2. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/__init__.py +1 -0
  3. cuthbertlib-0.0.11/cuthbertlib/enkf/README.md +14 -0
  4. cuthbertlib-0.0.11/cuthbertlib/enkf/__init__.py +1 -0
  5. cuthbertlib-0.0.11/cuthbertlib/enkf/filtering.py +119 -0
  6. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linalg/tria.py +2 -2
  7. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linearize/log_density.py +2 -0
  8. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/__init__.py +1 -0
  9. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/autodiff.py +0 -1
  10. cuthbertlib-0.0.11/cuthbertlib/resampling/no_resampling.py +57 -0
  11. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/smoothing/exact_sampling.py +2 -0
  12. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/smoothing/mcmc.py +2 -1
  13. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/smoothing/protocols.py +1 -0
  14. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/types.py +3 -1
  15. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/pyproject.toml +1 -1
  16. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/.gitignore +0 -0
  17. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/README.md +0 -0
  18. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/README.md +0 -0
  19. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/discrete/__init__.py +0 -0
  20. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/discrete/filtering.py +0 -0
  21. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/discrete/smoothing.py +0 -0
  22. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/kalman/README.md +0 -0
  23. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/kalman/__init__.py +0 -0
  24. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/kalman/filtering.py +0 -0
  25. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/kalman/generate.py +0 -0
  26. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/kalman/sampling.py +0 -0
  27. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/kalman/smoothing.py +0 -0
  28. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linalg/README.md +0 -0
  29. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linalg/__init__.py +0 -0
  30. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linalg/collect_nans_chol.py +0 -0
  31. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linalg/marginal_sqrt_cov.py +0 -0
  32. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linalg/symmetric_inv_sqrt.py +0 -0
  33. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linearize/README.md +0 -0
  34. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linearize/__init__.py +0 -0
  35. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linearize/moments.py +0 -0
  36. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linearize/taylor.py +0 -0
  37. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/README.md +0 -0
  38. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/__init__.py +0 -0
  39. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/common.py +0 -0
  40. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/cubature.py +0 -0
  41. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/gauss_hermite.py +0 -0
  42. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/linearize.py +0 -0
  43. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/unscented.py +0 -0
  44. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/utils.py +0 -0
  45. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/README.md +0 -0
  46. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/adaptive.py +0 -0
  47. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/killing.py +0 -0
  48. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/multinomial.py +0 -0
  49. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/protocols.py +0 -0
  50. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/systematic.py +0 -0
  51. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/utils.py +0 -0
  52. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/README.md +0 -0
  53. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/__init__.py +0 -0
  54. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/ess.py +0 -0
  55. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/smoothing/__init__.py +0 -0
  56. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/smoothing/tracing.py +0 -0
  57. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/stats/README.md +0 -0
  58. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/stats/__init__.py +0 -0
  59. {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/stats/multivariate_normal.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: cuthbertlib
3
- Version: 0.0.9
3
+ Version: 0.0.11
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
@@ -5,6 +5,7 @@ del version
5
5
 
6
6
  from cuthbertlib import (
7
7
  discrete,
8
+ enkf,
8
9
  kalman,
9
10
  linalg,
10
11
  linearize,
@@ -0,0 +1,14 @@
1
+ # Ensemble Kalman Filter (EnKF)
2
+
3
+ This sub-repository provides modular functions for the Ensemble Kalman Filter.
4
+
5
+ The core functions are:
6
+
7
+ - `predict`: Propagate ensemble members through nonlinear dynamics with additive Gaussian noise.
8
+ - `update`: Update ensemble members with an observation using the EnKF update equation.
9
+
10
+ Together, `predict` and `update` can be used to perform an online EnKF filtering step.
11
+
12
+ The EnKF uses an ensemble of particles with a Kalman-style measurement update based on
13
+ empirical covariances. Unlike the EKF, it does not require Jacobians, while naturally
14
+ handling nonlinear dynamics.
@@ -0,0 +1 @@
1
+ from cuthbertlib.enkf.filtering import predict, update
@@ -0,0 +1,119 @@
1
+ """Implements the Ensemble Kalman Filter (EnKF) predict and update steps.
2
+
3
+ See Algorithm 10.2, [Sanz-Alonso et al., Inverse Problems and Data Assimilation](https://arxiv.org/abs/1810.06191).
4
+ Based in part on the [CD-Dynamax implementation](https://github.com/hd-UQ/cd_dynamax/blob/public/cd_dynamax/src/continuous_discrete_nonlinear_gaussian_ssm/inference_enkf.py).
5
+ """
6
+
7
+ from typing import Callable
8
+
9
+ import jax
10
+ import jax.numpy as jnp
11
+ from jax import random
12
+ from jax.scipy.linalg import cho_solve
13
+
14
+ from cuthbertlib.linalg import collect_nans_chol, tria
15
+ from cuthbertlib.stats import multivariate_normal
16
+ from cuthbertlib.types import Array, KeyArray, ScalarArray
17
+
18
+ ObservationFn = Callable[[Array], Array]
19
+ DynamicsFn = Callable[[Array, KeyArray], Array]
20
+
21
+
22
+ def predict(
23
+ key: KeyArray,
24
+ ensemble: Array,
25
+ dynamics_fn: DynamicsFn,
26
+ inflation: float = 0.0,
27
+ ) -> Array:
28
+ """Propagate ensemble members through an arbitrary simulator p(x_{t+1} | x_t).
29
+
30
+ Args:
31
+ key: JAX PRNG key.
32
+ ensemble: Ensemble of state vectors, shape (N, x_dim).
33
+ dynamics_fn: Dynamics function mapping (state, key) -> state.
34
+ inflation: Multiplicative inflation factor applied to ensemble deviations.
35
+
36
+ Returns:
37
+ Predicted ensemble, shape (N, x_dim).
38
+ """
39
+ N, x_dim = ensemble.shape
40
+
41
+ # Propagate each member through the dynamics
42
+ keys = random.split(key, N)
43
+ propagated = jax.vmap(dynamics_fn, (0, 0))(ensemble, keys)
44
+
45
+ # Apply multiplicative inflation
46
+ mean = jnp.mean(propagated, axis=0)
47
+ propagated = mean + (1 + inflation) * (propagated - mean)
48
+
49
+ return propagated
50
+
51
+
52
+ def update(
53
+ key: KeyArray,
54
+ predicted_ensemble: Array,
55
+ observation_fn: ObservationFn,
56
+ chol_R: Array,
57
+ y: Array,
58
+ perturbed_obs: bool = True,
59
+ ) -> tuple[Array, ScalarArray]:
60
+ """Update ensemble members with an observation using the EnKF update.
61
+
62
+ NaNs in ``y`` are treated as missing dimensions and are excluded from the
63
+ update. When ``y`` is entirely NaN, the update is a no-op: the predicted
64
+ ensemble is returned unchanged with zero log-likelihood contribution.
65
+
66
+ Args:
67
+ key: JAX PRNG key.
68
+ predicted_ensemble: Predicted ensemble, shape (N, x_dim).
69
+ observation_fn: Observation function mapping state -> obs.
70
+ chol_R: Cholesky factor of the observation noise covariance, shape (y_dim, y_dim).
71
+ y: Observation vector, shape (y_dim,). NaNs indicate missing dimensions.
72
+ perturbed_obs: If True, use perturbed observations (stochastic EnKF).
73
+ If False, use deterministic update.
74
+
75
+ Returns:
76
+ Tuple of (updated_ensemble, log_likelihood).
77
+ """
78
+ N, x_dim = predicted_ensemble.shape
79
+
80
+ # Map ensemble to observation space
81
+ y_pred = jax.vmap(observation_fn, (0,))(predicted_ensemble)
82
+
83
+ # Handle partially-missing observations by reordering and zeroing missing dims.
84
+ # Use y_pred.T because y_pred is (N, y_dim) and we want to reorder along axis 0.
85
+ flag = jnp.isnan(y)
86
+ flag, chol_R, y, y_pred = collect_nans_chol(flag, chol_R, y, y_pred.T)
87
+ y_pred = y_pred.T
88
+ y_dim = y.shape[0]
89
+
90
+ # Ensemble means
91
+ x_mean = jnp.mean(predicted_ensemble, axis=0)
92
+ y_mean = jnp.mean(y_pred, axis=0)
93
+
94
+ # Deviations from ensemble mean
95
+ x_dev = predicted_ensemble - x_mean
96
+ y_dev = y_pred - y_mean
97
+
98
+ # Square-root innovation covariance via tria
99
+ chol_S = tria(jnp.concatenate([y_dev.T / jnp.sqrt(N - 1), chol_R], axis=1))
100
+
101
+ # Cross-covariance
102
+ C_xy = x_dev.T @ y_dev / (N - 1)
103
+
104
+ # Kalman gain: K = C_xy @ S^{-1} = C_xy @ cho_solve(chol_S, I)
105
+ K = cho_solve((chol_S, True), C_xy.T).T
106
+
107
+ # Innovation per member
108
+ if perturbed_obs:
109
+ y_n = y[None, :] + (chol_R @ random.normal(key, (y_dim, N))).T
110
+ else:
111
+ y_n = jnp.broadcast_to(y[None, :], (N, y_dim))
112
+
113
+ # Update ensemble
114
+ updated = predicted_ensemble + (y_n - y_pred) @ K.T
115
+
116
+ # Log-likelihood
117
+ ll = multivariate_normal.logpdf(y, y_mean, chol_S, nan_support=False)
118
+
119
+ return updated, jnp.asarray(ll)
@@ -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
- I = jnp.eye(K.shape[-1], dtype=K.dtype)
89
- dM = jnp.tril(K + K_T) - K * I
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
@@ -63,6 +63,7 @@ def linearize_log_density(
63
63
 
64
64
  Args:
65
65
  log_density: A conditional log density of y given x. Returns a scalar.
66
+ x must be the first argument and y the second.
66
67
  x: The input points.
67
68
  y: The output points.
68
69
  has_aux: Whether `log_density` returns an auxiliary value.
@@ -137,6 +138,7 @@ def linearize_log_density_given_chol_cov(
137
138
 
138
139
  Args:
139
140
  log_density: A conditional log density of y given x. Returns a scalar.
141
+ x must be the first argument and y the second.
140
142
  x: The input points.
141
143
  y: The output points.
142
144
  chol_cov: The Cholesky factor of the covariance matrix of the Gaussian.
@@ -3,6 +3,7 @@ from cuthbertlib.resampling import (
3
3
  autodiff,
4
4
  killing,
5
5
  multinomial,
6
+ no_resampling,
6
7
  systematic,
7
8
  )
8
9
  from cuthbertlib.resampling.adaptive import ess_decorator
@@ -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)
@@ -28,6 +28,7 @@ def log_weights_single(
28
28
  x1: The current state.
29
29
  log_weight_x0: The log weights of the previous state.
30
30
  log_density: The log density function of x1 given x0.
31
+ x0 must be the first argument and x1 the second.
31
32
 
32
33
  Returns:
33
34
  The smoothing weight for sample x0 given a single sample x1.
@@ -97,6 +98,7 @@ def simulate(
97
98
  x1_all: A collection of current states $x_1$.
98
99
  log_weight_x0_all: The log weights of $x_0$.
99
100
  log_density: The log density function of $x_1$ given $x_0$.
101
+ $x_0$ must be the first argument and $x_1$ the second.
100
102
  x1_ancestor_indices: The ancestor indices of $x_1$. Not used.
101
103
 
102
104
  Returns:
@@ -33,6 +33,7 @@ def simulate(
33
33
  x1_all: A collection of current states $x_1$.
34
34
  log_weight_x0_all: The log weights of $x_0$.
35
35
  log_density: The log density function of $x_1$ given $x_0$.
36
+ $x_0$ must be the first argument and $x_1$ the second.
36
37
  x1_ancestor_indices: The ancestor indices of $x_1$.
37
38
  n_steps: Number of MCMC steps to perform.
38
39
 
@@ -71,7 +72,7 @@ def simulate(
71
72
  return (idx, x0_res, idx_log_p), None
72
73
 
73
74
  x0_init = jax.tree.map(lambda z: z[x1_ancestor_indices], x0_all)
74
- init_log_p = jax.vmap(log_density)(x1_all, x0_init)
75
+ init_log_p = jax.vmap(log_density)(x0_init, x1_all)
75
76
  init = (x1_ancestor_indices, x0_init, init_log_p)
76
77
  (out_index, out_samples, _), _ = jax.lax.scan(body, init, keys)
77
78
  return out_samples, out_index
@@ -36,6 +36,7 @@ class BackwardSampling(Protocol):
36
36
  x1_all: A collection of current states $x_1$.
37
37
  log_weight_x0_all: The log weights of $x_0$.
38
38
  log_density: The log density function of $x_1$ given $x_0$.
39
+ $x_0$ must be the first argument and $x_1$ the second.
39
40
  x1_ancestor_indices: The ancestor indices of $x_1$.
40
41
 
41
42
  Returns:
@@ -14,7 +14,9 @@ ScalarArray: TypeAlias = (
14
14
  ScalarArrayLike: TypeAlias = ArrayLike # Object that will be cast to a ScalarArray
15
15
 
16
16
  LogDensity: TypeAlias = Callable[[ArrayTreeLike], ScalarArray]
17
- LogConditionalDensity: TypeAlias = Callable[[ArrayTreeLike, ArrayTreeLike], ScalarArray]
17
+ LogConditionalDensity: TypeAlias = Callable[
18
+ [ArrayTreeLike, ArrayTreeLike], ScalarArray
19
+ ] # p(x_1 | x_0), where x_0 is the first argument and x_1 the second.
18
20
  LogConditionalDensityAux: TypeAlias = Callable[
19
21
  [ArrayTreeLike, ArrayTreeLike], tuple[ScalarArray, ArrayTree]
20
22
  ]
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
4
4
 
5
5
  [project]
6
6
  name = "cuthbertlib"
7
- version = "0.0.9"
7
+ version = "0.0.11"
8
8
  description = "Atomic building blocks for state-space model inference with JAX"
9
9
  requires-python = ">=3.10"
10
10
  readme = "README.md"
File without changes
File without changes