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.
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/PKG-INFO +1 -1
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/__init__.py +1 -0
- cuthbertlib-0.0.11/cuthbertlib/enkf/README.md +14 -0
- cuthbertlib-0.0.11/cuthbertlib/enkf/__init__.py +1 -0
- cuthbertlib-0.0.11/cuthbertlib/enkf/filtering.py +119 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linalg/tria.py +2 -2
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linearize/log_density.py +2 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/__init__.py +1 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/autodiff.py +0 -1
- cuthbertlib-0.0.11/cuthbertlib/resampling/no_resampling.py +57 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/smoothing/exact_sampling.py +2 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/smoothing/mcmc.py +2 -1
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/smoothing/protocols.py +1 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/types.py +3 -1
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/pyproject.toml +1 -1
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/.gitignore +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/README.md +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/README.md +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/discrete/__init__.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/discrete/filtering.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/discrete/smoothing.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/kalman/README.md +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/kalman/__init__.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/kalman/filtering.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/kalman/generate.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/kalman/sampling.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/kalman/smoothing.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linalg/README.md +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linalg/__init__.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linalg/collect_nans_chol.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linalg/marginal_sqrt_cov.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linalg/symmetric_inv_sqrt.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linearize/README.md +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linearize/__init__.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linearize/moments.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/linearize/taylor.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/README.md +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/__init__.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/common.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/cubature.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/gauss_hermite.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/linearize.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/unscented.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/quadrature/utils.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/README.md +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/adaptive.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/killing.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/multinomial.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/protocols.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/systematic.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/resampling/utils.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/README.md +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/__init__.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/ess.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/smoothing/__init__.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/smc/smoothing/tracing.py +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/stats/README.md +0 -0
- {cuthbertlib-0.0.9 → cuthbertlib-0.0.11}/cuthbertlib/stats/__init__.py +0 -0
- {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.
|
|
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
|
|
@@ -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
|
-
|
|
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
|
|
@@ -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.
|
|
@@ -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)(
|
|
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[
|
|
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
|
]
|
|
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
|