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.
Files changed (41) hide show
  1. bayesmith/__init__.py +173 -0
  2. bayesmith/bridge/__init__.py +1 -0
  3. bayesmith/bridge/numpyro_bridge.py +205 -0
  4. bayesmith/diagnose/__init__.py +50 -0
  5. bayesmith/diagnose/identifiability.py +398 -0
  6. bayesmith/diagnose/local.py +332 -0
  7. bayesmith/diagnose/priors.py +422 -0
  8. bayesmith/diagnose/sensitivity.py +816 -0
  9. bayesmith/dispatch/__init__.py +24 -0
  10. bayesmith/dispatch/classify.py +623 -0
  11. bayesmith/dispatch/execute.py +834 -0
  12. bayesmith/dispatch/plan.py +843 -0
  13. bayesmith/dispatch/streaming.py +237 -0
  14. bayesmith/errors.py +85 -0
  15. bayesmith/evidence/__init__.py +56 -0
  16. bayesmith/evidence/campaign.py +408 -0
  17. bayesmith/evidence/compress.py +384 -0
  18. bayesmith/evidence/diagnostics.py +177 -0
  19. bayesmith/evidence/factorize.py +281 -0
  20. bayesmith/evidence/sqrtinfo.py +403 -0
  21. bayesmith/exact/__init__.py +53 -0
  22. bayesmith/exact/block.py +456 -0
  23. bayesmith/exact/conditioning.py +118 -0
  24. bayesmith/exact/correct.py +344 -0
  25. bayesmith/exact/discrete.py +246 -0
  26. bayesmith/exact/fisher.py +500 -0
  27. bayesmith/exact/gaussian.py +548 -0
  28. bayesmith/exact/gibbs.py +416 -0
  29. bayesmith/exact/gls.py +646 -0
  30. bayesmith/exact/linearity.py +809 -0
  31. bayesmith/exact/precision.py +451 -0
  32. bayesmith/exact/solve.py +537 -0
  33. bayesmith/graph/__init__.py +1 -0
  34. bayesmith/graph/evaluate.py +177 -0
  35. bayesmith/graph/graph.py +151 -0
  36. bayesmith/graph/nodes.py +180 -0
  37. bayesmith/graph/trace.py +271 -0
  38. bayesmith-0.1.0.dist-info/METADATA +87 -0
  39. bayesmith-0.1.0.dist-info/RECORD +41 -0
  40. bayesmith-0.1.0.dist-info/WHEEL +4 -0
  41. 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
+ ]