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.
Files changed (59) hide show
  1. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/PKG-INFO +1 -1
  2. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/enkf/filtering.py +1 -1
  3. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linalg/tria.py +2 -2
  4. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/__init__.py +1 -0
  5. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/autodiff.py +0 -1
  6. cuthbertlib-0.0.12/cuthbertlib/resampling/no_resampling.py +57 -0
  7. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/pyproject.toml +1 -1
  8. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/.gitignore +0 -0
  9. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/README.md +0 -0
  10. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/README.md +0 -0
  11. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/__init__.py +0 -0
  12. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/discrete/__init__.py +0 -0
  13. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/discrete/filtering.py +0 -0
  14. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/discrete/smoothing.py +0 -0
  15. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/enkf/README.md +0 -0
  16. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/enkf/__init__.py +0 -0
  17. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/kalman/README.md +0 -0
  18. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/kalman/__init__.py +0 -0
  19. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/kalman/filtering.py +0 -0
  20. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/kalman/generate.py +0 -0
  21. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/kalman/sampling.py +0 -0
  22. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/kalman/smoothing.py +0 -0
  23. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linalg/README.md +0 -0
  24. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linalg/__init__.py +0 -0
  25. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linalg/collect_nans_chol.py +0 -0
  26. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linalg/marginal_sqrt_cov.py +0 -0
  27. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linalg/symmetric_inv_sqrt.py +0 -0
  28. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linearize/README.md +0 -0
  29. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linearize/__init__.py +0 -0
  30. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linearize/log_density.py +0 -0
  31. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linearize/moments.py +0 -0
  32. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/linearize/taylor.py +0 -0
  33. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/README.md +0 -0
  34. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/__init__.py +0 -0
  35. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/common.py +0 -0
  36. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/cubature.py +0 -0
  37. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/gauss_hermite.py +0 -0
  38. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/linearize.py +0 -0
  39. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/unscented.py +0 -0
  40. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/quadrature/utils.py +0 -0
  41. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/README.md +0 -0
  42. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/adaptive.py +0 -0
  43. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/killing.py +0 -0
  44. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/multinomial.py +0 -0
  45. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/protocols.py +0 -0
  46. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/systematic.py +0 -0
  47. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/resampling/utils.py +0 -0
  48. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/README.md +0 -0
  49. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/__init__.py +0 -0
  50. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/ess.py +0 -0
  51. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/smoothing/__init__.py +0 -0
  52. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/smoothing/exact_sampling.py +0 -0
  53. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/smoothing/mcmc.py +0 -0
  54. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/smoothing/protocols.py +0 -0
  55. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/smc/smoothing/tracing.py +0 -0
  56. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/stats/README.md +0 -0
  57. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/stats/__init__.py +0 -0
  58. {cuthbertlib-0.0.10 → cuthbertlib-0.0.12}/cuthbertlib/stats/multivariate_normal.py +0 -0
  59. {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.10
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, ArrayTreeLike, KeyArray, ScalarArray
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
- 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
@@ -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)
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
4
4
 
5
5
  [project]
6
6
  name = "cuthbertlib"
7
- version = "0.0.10"
7
+ version = "0.0.12"
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