bayesmith 0.1.0__py3-none-any.whl
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.
- bayesmith/__init__.py +173 -0
- bayesmith/bridge/__init__.py +1 -0
- bayesmith/bridge/numpyro_bridge.py +205 -0
- bayesmith/diagnose/__init__.py +50 -0
- bayesmith/diagnose/identifiability.py +398 -0
- bayesmith/diagnose/local.py +332 -0
- bayesmith/diagnose/priors.py +422 -0
- bayesmith/diagnose/sensitivity.py +816 -0
- bayesmith/dispatch/__init__.py +24 -0
- bayesmith/dispatch/classify.py +623 -0
- bayesmith/dispatch/execute.py +834 -0
- bayesmith/dispatch/plan.py +843 -0
- bayesmith/dispatch/streaming.py +237 -0
- bayesmith/errors.py +85 -0
- bayesmith/evidence/__init__.py +56 -0
- bayesmith/evidence/campaign.py +408 -0
- bayesmith/evidence/compress.py +384 -0
- bayesmith/evidence/diagnostics.py +177 -0
- bayesmith/evidence/factorize.py +281 -0
- bayesmith/evidence/sqrtinfo.py +403 -0
- bayesmith/exact/__init__.py +53 -0
- bayesmith/exact/block.py +456 -0
- bayesmith/exact/conditioning.py +118 -0
- bayesmith/exact/correct.py +344 -0
- bayesmith/exact/discrete.py +246 -0
- bayesmith/exact/fisher.py +500 -0
- bayesmith/exact/gaussian.py +548 -0
- bayesmith/exact/gibbs.py +416 -0
- bayesmith/exact/gls.py +646 -0
- bayesmith/exact/linearity.py +809 -0
- bayesmith/exact/precision.py +451 -0
- bayesmith/exact/solve.py +537 -0
- bayesmith/graph/__init__.py +1 -0
- bayesmith/graph/evaluate.py +177 -0
- bayesmith/graph/graph.py +151 -0
- bayesmith/graph/nodes.py +180 -0
- bayesmith/graph/trace.py +271 -0
- bayesmith-0.1.0.dist-info/METADATA +87 -0
- bayesmith-0.1.0.dist-info/RECORD +41 -0
- bayesmith-0.1.0.dist-info/WHEEL +4 -0
- bayesmith-0.1.0.dist-info/licenses/LICENSE +21 -0
bayesmith/__init__.py
ADDED
|
@@ -0,0 +1,173 @@
|
|
|
1
|
+
"""bayesmith: a graph of operators is a Bayesian model.
|
|
2
|
+
|
|
3
|
+
Deterministic operators propagate dependence; probabilistic ones contribute a
|
|
4
|
+
conditional density. The graph's structure is what selects the inference
|
|
5
|
+
method -- exact where a subgraph permits one, NUTS where it does not.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import importlib
|
|
11
|
+
import importlib.metadata as _metadata
|
|
12
|
+
from typing import Any
|
|
13
|
+
|
|
14
|
+
from bayesmith.errors import (
|
|
15
|
+
BayesmithError,
|
|
16
|
+
ConvergenceError,
|
|
17
|
+
GraphError,
|
|
18
|
+
NotGaussian,
|
|
19
|
+
StructureError,
|
|
20
|
+
TraceError,
|
|
21
|
+
)
|
|
22
|
+
|
|
23
|
+
__all__ = [
|
|
24
|
+
# tracing
|
|
25
|
+
"trace",
|
|
26
|
+
"const",
|
|
27
|
+
"det",
|
|
28
|
+
"sample",
|
|
29
|
+
"observe",
|
|
30
|
+
"plate",
|
|
31
|
+
"NodeRef",
|
|
32
|
+
"PlateRef",
|
|
33
|
+
# graph
|
|
34
|
+
"Graph",
|
|
35
|
+
"Plate",
|
|
36
|
+
"Node",
|
|
37
|
+
"Const",
|
|
38
|
+
"Deterministic",
|
|
39
|
+
"Probabilistic",
|
|
40
|
+
# evaluation
|
|
41
|
+
"evaluate",
|
|
42
|
+
"log_joint",
|
|
43
|
+
# inference
|
|
44
|
+
"compile",
|
|
45
|
+
"Posterior",
|
|
46
|
+
"Estimate",
|
|
47
|
+
"to_numpyro",
|
|
48
|
+
"nuts",
|
|
49
|
+
"predict",
|
|
50
|
+
# exact
|
|
51
|
+
"linear_operator",
|
|
52
|
+
"check_linearity",
|
|
53
|
+
"wiener_solve",
|
|
54
|
+
"gcr_sample",
|
|
55
|
+
"condition_bound",
|
|
56
|
+
"iterative_gls",
|
|
57
|
+
"sigma_from_graph",
|
|
58
|
+
"noise_std_at",
|
|
59
|
+
"precision_at",
|
|
60
|
+
"fisher_information",
|
|
61
|
+
"parameter_covariance",
|
|
62
|
+
# diagnose
|
|
63
|
+
"identifiability",
|
|
64
|
+
"prior_sensitivity",
|
|
65
|
+
"JeffreysPrior",
|
|
66
|
+
# errors
|
|
67
|
+
"BayesmithError",
|
|
68
|
+
"GraphError",
|
|
69
|
+
"TraceError",
|
|
70
|
+
"StructureError",
|
|
71
|
+
"ConvergenceError",
|
|
72
|
+
"NotGaussian",
|
|
73
|
+
]
|
|
74
|
+
|
|
75
|
+
# Every public name above except the six error classes is resolved lazily,
|
|
76
|
+
# on first attribute access, rather than imported here at module scope.
|
|
77
|
+
# Importing eagerly would make `import bayesmith` load numpyro (hence jax) as
|
|
78
|
+
# a side effect -- exactly the regression this module previously had: Python
|
|
79
|
+
# always runs a package's __init__.py before any of its submodules, so even
|
|
80
|
+
# `import bayesmith.errors` was dragging in the whole bridge, which broke the
|
|
81
|
+
# stdlib-only contract errors.py documents for itself and which
|
|
82
|
+
# test_errors_module_imports_no_heavy_dependency enforces. Only errors.py is
|
|
83
|
+
# cheap and stdlib-only, so it alone is still imported eagerly above.
|
|
84
|
+
#
|
|
85
|
+
# name -> (owning submodule, attribute name within it)
|
|
86
|
+
_LAZY_ATTRS: dict[str, tuple[str, str]] = {
|
|
87
|
+
"trace": ("bayesmith.graph.trace", "trace"),
|
|
88
|
+
"const": ("bayesmith.graph.trace", "const"),
|
|
89
|
+
"det": ("bayesmith.graph.trace", "det"),
|
|
90
|
+
"sample": ("bayesmith.graph.trace", "sample"),
|
|
91
|
+
"observe": ("bayesmith.graph.trace", "observe"),
|
|
92
|
+
"plate": ("bayesmith.graph.trace", "plate"),
|
|
93
|
+
"NodeRef": ("bayesmith.graph.trace", "NodeRef"),
|
|
94
|
+
"PlateRef": ("bayesmith.graph.trace", "PlateRef"),
|
|
95
|
+
"Graph": ("bayesmith.graph.graph", "Graph"),
|
|
96
|
+
"Plate": ("bayesmith.graph.graph", "Plate"),
|
|
97
|
+
"Node": ("bayesmith.graph.nodes", "Node"),
|
|
98
|
+
"Const": ("bayesmith.graph.nodes", "Const"),
|
|
99
|
+
"Deterministic": ("bayesmith.graph.nodes", "Deterministic"),
|
|
100
|
+
"Probabilistic": ("bayesmith.graph.nodes", "Probabilistic"),
|
|
101
|
+
"evaluate": ("bayesmith.graph.evaluate", "evaluate"),
|
|
102
|
+
"log_joint": ("bayesmith.graph.evaluate", "log_joint"),
|
|
103
|
+
"compile": ("bayesmith.dispatch.plan", "compile"),
|
|
104
|
+
"Posterior": ("bayesmith.dispatch.execute", "Posterior"),
|
|
105
|
+
"Estimate": ("bayesmith.dispatch.execute", "Estimate"),
|
|
106
|
+
"to_numpyro": ("bayesmith.bridge.numpyro_bridge", "to_numpyro"),
|
|
107
|
+
"nuts": ("bayesmith.bridge.numpyro_bridge", "nuts"),
|
|
108
|
+
"predict": ("bayesmith.bridge.numpyro_bridge", "predict"),
|
|
109
|
+
"linear_operator": ("bayesmith.exact.linearity", "linear_operator"),
|
|
110
|
+
"check_linearity": ("bayesmith.exact.linearity", "check_linearity"),
|
|
111
|
+
"wiener_solve": ("bayesmith.exact.solve", "wiener_solve"),
|
|
112
|
+
"gcr_sample": ("bayesmith.exact.solve", "gcr_sample"),
|
|
113
|
+
"condition_bound": ("bayesmith.exact.solve", "condition_bound"),
|
|
114
|
+
"iterative_gls": ("bayesmith.exact.gls", "iterative_gls"),
|
|
115
|
+
"sigma_from_graph": ("bayesmith.exact.gls", "sigma_from_graph"),
|
|
116
|
+
"noise_std_at": ("bayesmith.exact.gaussian", "noise_std_at"),
|
|
117
|
+
"precision_at": ("bayesmith.exact.gaussian", "precision_at"),
|
|
118
|
+
"fisher_information": ("bayesmith.exact.fisher", "fisher_information"),
|
|
119
|
+
"parameter_covariance": ("bayesmith.exact.fisher", "parameter_covariance"),
|
|
120
|
+
"identifiability": ("bayesmith.diagnose.identifiability", "identifiability"),
|
|
121
|
+
"prior_sensitivity": ("bayesmith.diagnose.sensitivity", "prior_sensitivity"),
|
|
122
|
+
"JeffreysPrior": ("bayesmith.diagnose.priors", "JeffreysPrior"),
|
|
123
|
+
}
|
|
124
|
+
|
|
125
|
+
# Subpackages reachable as `bayesmith.<name>` after a bare `import bayesmith`,
|
|
126
|
+
# without eagerly importing any of them -- `bridge` in particular is what
|
|
127
|
+
# pulls in numpyro, and `exact` reaches numpyro through `bridge` too (see
|
|
128
|
+
# `gaussian.py`'s use of `numpyro.distributions`). `errors` is listed too for
|
|
129
|
+
# __dir__'s sake even though the eager import above already binds it as a
|
|
130
|
+
# real attribute, so __getattr__ is never actually consulted for it.
|
|
131
|
+
#
|
|
132
|
+
# `evidence` reaches jax through `sqrtinfo.py`'s module-scope import, so it is
|
|
133
|
+
# listed here and NOT imported eagerly, for the same reason `exact` is not.
|
|
134
|
+
# It was missing from this tuple for the whole of B11: the layer was complete,
|
|
135
|
+
# dense-oracled and cross-checked, and `import bayesmith; bayesmith.evidence`
|
|
136
|
+
# still raised AttributeError, so only an explicit `import bayesmith.evidence`
|
|
137
|
+
# reached it. Both halves are pinned in `tests/test_public_api.py` -- that it
|
|
138
|
+
# resolves, and that resolving it is what pulls jax in rather than importing
|
|
139
|
+
# this package.
|
|
140
|
+
_LAZY_SUBMODULES = (
|
|
141
|
+
"graph",
|
|
142
|
+
"bridge",
|
|
143
|
+
"exact",
|
|
144
|
+
"dispatch",
|
|
145
|
+
"evidence",
|
|
146
|
+
"diagnose",
|
|
147
|
+
"errors",
|
|
148
|
+
)
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
#: The installed distribution's version, READ from the installed metadata
|
|
152
|
+
#: rather than written here. Two spellings of one version is the defect this
|
|
153
|
+
#: package has spent the most effort repairing, and a release number is the
|
|
154
|
+
#: worst candidate for a second copy: it changes on exactly the commit where
|
|
155
|
+
#: everyone is busy doing something else.
|
|
156
|
+
__version__ = _metadata.version("bayesmith")
|
|
157
|
+
|
|
158
|
+
|
|
159
|
+
def __getattr__(name: str) -> Any:
|
|
160
|
+
if name in _LAZY_ATTRS:
|
|
161
|
+
module_name, attr_name = _LAZY_ATTRS[name]
|
|
162
|
+
value = getattr(importlib.import_module(module_name), attr_name)
|
|
163
|
+
globals()[name] = value # cache: later lookups skip __getattr__
|
|
164
|
+
return value
|
|
165
|
+
if name in _LAZY_SUBMODULES:
|
|
166
|
+
module = importlib.import_module(f"{__name__}.{name}")
|
|
167
|
+
globals()[name] = module
|
|
168
|
+
return module
|
|
169
|
+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
|
170
|
+
|
|
171
|
+
|
|
172
|
+
def __dir__() -> list[str]:
|
|
173
|
+
return sorted(set(globals()) | set(_LAZY_ATTRS) | set(_LAZY_SUBMODULES))
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""Bridges to external inference engines."""
|
|
@@ -0,0 +1,205 @@
|
|
|
1
|
+
"""Turning a graph into a NumPyro model, and running NUTS on it.
|
|
2
|
+
|
|
3
|
+
This is the last row of the dispatch table: whatever structure bayesmith
|
|
4
|
+
cannot solve exactly is handed to NumPyro. It is also the oracle every exact
|
|
5
|
+
path is checked against, because a graph that qualifies for an exact method
|
|
6
|
+
always also qualifies for NUTS.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
from collections.abc import Callable, Mapping
|
|
12
|
+
from typing import Any
|
|
13
|
+
|
|
14
|
+
import jax
|
|
15
|
+
import jax.numpy as jnp
|
|
16
|
+
import numpyro
|
|
17
|
+
from numpyro.infer import MCMC, NUTS
|
|
18
|
+
|
|
19
|
+
from bayesmith.graph.evaluate import apply_deterministic, apply_probabilistic
|
|
20
|
+
from bayesmith.graph.graph import Graph
|
|
21
|
+
from bayesmith.graph.nodes import Const, Deterministic, Probabilistic
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def to_numpyro(graph: Graph) -> Callable[[], dict[str, Any]]:
|
|
25
|
+
"""Build a NumPyro model that declares the same joint distribution.
|
|
26
|
+
|
|
27
|
+
Latent and observed nodes become ``numpyro.sample`` sites carrying the
|
|
28
|
+
graph's own node names, so posterior samples come back keyed by them.
|
|
29
|
+
Deterministic nodes are recorded with ``numpyro.deterministic`` so they
|
|
30
|
+
appear in traces and predictives without contributing a density.
|
|
31
|
+
"""
|
|
32
|
+
|
|
33
|
+
def model() -> dict[str, Any]:
|
|
34
|
+
env: dict[str, Any] = {}
|
|
35
|
+
for node in graph.nodes:
|
|
36
|
+
if isinstance(node, Const):
|
|
37
|
+
env[node.name] = node.value
|
|
38
|
+
elif isinstance(node, Deterministic):
|
|
39
|
+
env[node.name] = numpyro.deterministic(
|
|
40
|
+
node.name, apply_deterministic(graph, node, env)
|
|
41
|
+
)
|
|
42
|
+
elif isinstance(node, Probabilistic):
|
|
43
|
+
distribution = apply_probabilistic(graph, node, env)
|
|
44
|
+
if node.plate:
|
|
45
|
+
name = node.plate[0]
|
|
46
|
+
with numpyro.plate(name, graph.plate_size(name)):
|
|
47
|
+
env[node.name] = numpyro.sample(
|
|
48
|
+
node.name, distribution, obs=node.observed
|
|
49
|
+
)
|
|
50
|
+
else:
|
|
51
|
+
env[node.name] = numpyro.sample(
|
|
52
|
+
node.name, distribution, obs=node.observed
|
|
53
|
+
)
|
|
54
|
+
return env
|
|
55
|
+
|
|
56
|
+
return model
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def nuts(
|
|
60
|
+
graph: Graph,
|
|
61
|
+
key: jax.Array,
|
|
62
|
+
*,
|
|
63
|
+
num_warmup: int = 1000,
|
|
64
|
+
num_samples: int = 2000,
|
|
65
|
+
num_chains: int = 1,
|
|
66
|
+
chain_method: str = "sequential",
|
|
67
|
+
progress_bar: bool = False,
|
|
68
|
+
nuts_options: Mapping[str, Any] | None = None,
|
|
69
|
+
) -> dict[str, jax.Array]:
|
|
70
|
+
"""Sample the posterior of ``graph`` with NUTS.
|
|
71
|
+
|
|
72
|
+
``chain_method`` and ``nuts_options`` are here because
|
|
73
|
+
:meth:`~bayesmith.dispatch.plan.InferencePlan.sample` promises them on
|
|
74
|
+
every path, and two of its five shapes -- the graph with no exact block,
|
|
75
|
+
and the SNIS collapse -- run through this function. Until they existed
|
|
76
|
+
those two shapes silently ignored both keywords while the mixed shape,
|
|
77
|
+
which reaches ``HMCGibbs`` through
|
|
78
|
+
:func:`~bayesmith.exact.gibbs.assemble`, honoured them. The names and the
|
|
79
|
+
defaults are ``assemble``'s, so the two spellings of "run a chain" take
|
|
80
|
+
the same words.
|
|
81
|
+
|
|
82
|
+
Args:
|
|
83
|
+
graph: the model.
|
|
84
|
+
key: a PRNG key.
|
|
85
|
+
num_warmup: adaptation draws, discarded.
|
|
86
|
+
num_samples: retained draws per chain.
|
|
87
|
+
num_chains: independent chains.
|
|
88
|
+
chain_method: how ``num_chains`` are run -- ``"sequential"``,
|
|
89
|
+
``"parallel"`` or ``"vectorized"``. All three are numpyro's own
|
|
90
|
+
and all three are legal here; ``assemble`` refuses
|
|
91
|
+
``"vectorized"`` for a Gibbs sweep, but that refusal is about
|
|
92
|
+
``HMCGibbs.init`` and does not apply to a bare kernel.
|
|
93
|
+
progress_bar: whether NumPyro prints progress.
|
|
94
|
+
nuts_options: keywords for the ``NUTS`` kernel itself
|
|
95
|
+
(``target_accept_prob``, ``max_tree_depth``, ``dense_mass``, ...).
|
|
96
|
+
|
|
97
|
+
**``init_strategy`` goes here, and on a narrow posterior it is
|
|
98
|
+
not a tuning knob.** NumPyro's default is ``init_to_uniform``,
|
|
99
|
+
which draws in the UNCONSTRAINED space with no knowledge of
|
|
100
|
+
where the graph's priors sit; a posterior far narrower than its
|
|
101
|
+
prior is a needle, and warmup then adapts a step size for
|
|
102
|
+
wherever it landed in the haystack. Measured on a power law
|
|
103
|
+
whose amplitude has a prior of 1e4 and a posterior width of
|
|
104
|
+
0.4, two chains of 400:
|
|
105
|
+
|
|
106
|
+
============================================== ====== ======
|
|
107
|
+
init r_hat ESS
|
|
108
|
+
============================================== ====== ======
|
|
109
|
+
default (``init_to_uniform``) 1609 1.0
|
|
110
|
+
``init_to_value`` at the declared values 1.006 138.6
|
|
111
|
+
============================================== ====== ======
|
|
112
|
+
|
|
113
|
+
The declared point does not have to be good -- only somewhere a
|
|
114
|
+
gradient can be followed. This is rheplicant's
|
|
115
|
+
``init_to_declared`` lesson, and the lesson is what was carried:
|
|
116
|
+
the remedy already exists here as a keyword, so what was missing
|
|
117
|
+
was the sentence saying so.
|
|
118
|
+
|
|
119
|
+
Returns:
|
|
120
|
+
A mapping from latent node name to its draws.
|
|
121
|
+
"""
|
|
122
|
+
mcmc = MCMC(
|
|
123
|
+
NUTS(to_numpyro(graph), **dict(nuts_options or {})),
|
|
124
|
+
num_warmup=num_warmup,
|
|
125
|
+
num_samples=num_samples,
|
|
126
|
+
num_chains=num_chains,
|
|
127
|
+
chain_method=chain_method,
|
|
128
|
+
progress_bar=progress_bar,
|
|
129
|
+
)
|
|
130
|
+
mcmc.run(key)
|
|
131
|
+
return mcmc.get_samples()
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
def predict(
|
|
135
|
+
graph: Graph, samples: Mapping[str, Any], key: jax.Array | None = None
|
|
136
|
+
) -> dict[str, jax.Array]:
|
|
137
|
+
"""Posterior predictive: every node's value over a stack of draws.
|
|
138
|
+
|
|
139
|
+
Args:
|
|
140
|
+
graph: the model the draws came from.
|
|
141
|
+
samples: ``{latent name: (n_draws, *latent shape)}`` -- what
|
|
142
|
+
:func:`nuts` returns, unchanged.
|
|
143
|
+
key: PRNG key for any node the draws do not fix. Fixed by default,
|
|
144
|
+
so a predictive is reproducible.
|
|
145
|
+
|
|
146
|
+
Returns:
|
|
147
|
+
``{node name: (n_draws, *node shape)}``, deterministic nodes
|
|
148
|
+
included.
|
|
149
|
+
|
|
150
|
+
Raises:
|
|
151
|
+
GraphError: if a latent is missing from ``samples``; if a stack's
|
|
152
|
+
PER-SAMPLE shape is not the latent's own; or if the stacks
|
|
153
|
+
disagree about how many draws there are.
|
|
154
|
+
|
|
155
|
+
Note:
|
|
156
|
+
**Why the per-sample shape is checked and not just the name.** This
|
|
157
|
+
is rheplicant's ``predict_from_samples`` guard, ported because the
|
|
158
|
+
failure it prevents is reachable here and is silent. Measured on a
|
|
159
|
+
length-3 latent with three draws: handing the stack in TRANSPOSED
|
|
160
|
+
returns a finite, correctly-shaped ``(3, 3)`` predictive whose every
|
|
161
|
+
entry is wrong --
|
|
162
|
+
|
|
163
|
+
correct [[0, 2, 6], [3, 8, 15], [6, 14, 24]]
|
|
164
|
+
transposed [[0, 6, 18], [1, 8, 21], [2, 10, 24]]
|
|
165
|
+
|
|
166
|
+
-- because NumPyro's ``Predictive`` maps over the leading axis and
|
|
167
|
+
has no independent statement of what the latent's shape is. The
|
|
168
|
+
graph does have one: each latent's ``dist_fn``. A non-square
|
|
169
|
+
transposition raises a broadcast error from three layers down that
|
|
170
|
+
names neither the site nor the axis; a square one raises nothing at
|
|
171
|
+
all.
|
|
172
|
+
"""
|
|
173
|
+
from numpyro.infer import Predictive
|
|
174
|
+
|
|
175
|
+
from bayesmith.dispatch.classify import prior_environment
|
|
176
|
+
from bayesmith.errors import GraphError
|
|
177
|
+
|
|
178
|
+
declared = prior_environment(graph)
|
|
179
|
+
draws: set[int] = set()
|
|
180
|
+
for name in graph.latents:
|
|
181
|
+
if name not in samples:
|
|
182
|
+
raise GraphError(
|
|
183
|
+
f"samples is missing latent {name!r}; available: "
|
|
184
|
+
f"{sorted(samples)}. A predictive needs every latent the "
|
|
185
|
+
"graph declares, because the deterministic nodes read them."
|
|
186
|
+
)
|
|
187
|
+
expected = jnp.shape(declared[name])
|
|
188
|
+
got = jnp.shape(samples[name])
|
|
189
|
+
if got[1:] != expected:
|
|
190
|
+
raise GraphError(
|
|
191
|
+
f"samples[{name!r}] has per-sample shape {got[1:]}, but the "
|
|
192
|
+
f"latent is {expected} (its full stack is {got}). The LEADING "
|
|
193
|
+
"axis must be the draw axis. Checking only the name would let "
|
|
194
|
+
"a transposed stack broadcast into the prediction and return a "
|
|
195
|
+
"finite, correctly-shaped, wrong predictive -- silently, "
|
|
196
|
+
"whenever the draw count happens to equal the latent's size."
|
|
197
|
+
)
|
|
198
|
+
draws.add(got[0])
|
|
199
|
+
if len(draws) > 1:
|
|
200
|
+
raise GraphError(
|
|
201
|
+
f"the sample stacks disagree about the number of draws: "
|
|
202
|
+
f"{sorted(draws)}. They must all come from one run."
|
|
203
|
+
)
|
|
204
|
+
predictive = Predictive(to_numpyro(graph), posterior_samples=dict(samples))
|
|
205
|
+
return predictive(jax.random.key(0) if key is None else key)
|
|
@@ -0,0 +1,50 @@
|
|
|
1
|
+
"""Design-time diagnostics: what the model cannot tell apart, and what the
|
|
2
|
+
priors did to it.
|
|
3
|
+
|
|
4
|
+
Three questions, each answered about the model *at a point* rather than
|
|
5
|
+
claimed globally -- which is why this package linearizes locally and never
|
|
6
|
+
reads ``linear_in``:
|
|
7
|
+
|
|
8
|
+
* :func:`identifiability` -- the rank of the joint Jacobian: which
|
|
9
|
+
combinations of latents the data is blind to, across blocks, where every
|
|
10
|
+
per-block guard is structurally silent.
|
|
11
|
+
* :func:`prior_sensitivity` -- how far each latent's own prior moved the
|
|
12
|
+
mode, in posterior sigmas, from two deterministic routes that verify each
|
|
13
|
+
other.
|
|
14
|
+
* :class:`JeffreysPrior` -- ``sqrt(det I)`` over a named block, evaluated
|
|
15
|
+
from the graph's own noise, with the determinant taken by ``eigvalsh``
|
|
16
|
+
plus a rank floor because ``slogdet`` and ``cholesky`` both return
|
|
17
|
+
plausible finite answers on a singular block.
|
|
18
|
+
|
|
19
|
+
Ported from ``rheplicant.inference.{identifiability,sensitivity,priors}``
|
|
20
|
+
(migration spec §八 step 5); the per-module cross-check records live in
|
|
21
|
+
``docs/migration/``. Precision discipline: run everything, graph
|
|
22
|
+
construction included, inside ``with jax.enable_x64(True):`` -- these
|
|
23
|
+
verdicts live at 1e-17 of the largest singular value, and a float32 result
|
|
24
|
+
is refused by name rather than silently reported.
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
from bayesmith.diagnose.identifiability import (
|
|
28
|
+
DEFAULT_RANK_RTOL,
|
|
29
|
+
IdentifiabilityReport,
|
|
30
|
+
identifiability,
|
|
31
|
+
)
|
|
32
|
+
from bayesmith.diagnose.priors import JeffreysPrior
|
|
33
|
+
from bayesmith.diagnose.sensitivity import (
|
|
34
|
+
CRITERION_SHIFT,
|
|
35
|
+
PriorSensitivityReport,
|
|
36
|
+
prior_sensitivity,
|
|
37
|
+
)
|
|
38
|
+
|
|
39
|
+
__all__ = [
|
|
40
|
+
# identifiability
|
|
41
|
+
"identifiability",
|
|
42
|
+
"IdentifiabilityReport",
|
|
43
|
+
"DEFAULT_RANK_RTOL",
|
|
44
|
+
# sensitivity
|
|
45
|
+
"prior_sensitivity",
|
|
46
|
+
"PriorSensitivityReport",
|
|
47
|
+
"CRITERION_SHIFT",
|
|
48
|
+
# joint priors
|
|
49
|
+
"JeffreysPrior",
|
|
50
|
+
]
|