{
"cells": [
{
"cell_type": "markdown",
"id": "b9100fca",
"metadata": {},
"source": [
"# 6. Save and resume results\n",
"\n",
"A posterior export is useful for plots and summaries. Continuing MCMC also\n",
"requires sampler state and the same analysis configuration. This chapter\n",
"extends the [Quickstart](../quickstart.ipynb) with both kinds of persistence.\n",
"\n",
"\n",
"This notebook is self-contained. Install JeansPy with the `numpyro_cpu` and\n",
"`plotting` extras as described in the installation guide, then select that\n",
"environment as your Jupyter kernel and run cells from top to bottom.\n",
"Saved outputs are an example run; timings and short-chain results can vary."
]
},
{
"cell_type": "markdown",
"id": "9ec481e5",
"metadata": {},
"source": [
"## Setup\n",
"\n",
"Run this setup in a fresh kernel. It defines all objects used below.\n",
"\n",
"The generalized NFW halo uses `ZhaoModel` with fixed `alpha=1`, `beta=3`\n",
"and an untruncated cutoff. Its inner slope `gamma` is sampled in the\n",
"inference examples; `truth[\"gamma\"]=1` generates standard-NFW mocks."
]
},
{
"cell_type": "code",
"execution_count": 1,
"id": "425ca5e6",
"metadata": {},
"outputs": [],
"source": [
"import os\n",
"os.environ.setdefault(\"JEANSPY_JAX_PLATFORM\", \"cpu\")\n",
"os.environ.setdefault(\"JEANSPY_JAX_ENABLE_X64\", \"true\")\n",
"\n",
"# The example uses no progress widgets; ignore only their optional-import warning.\n",
"import warnings\n",
"warnings.filterwarnings(\"ignore\", message=\"IProgress not found.*\")\n",
"\n",
"# Import JeansPy before JAX so its runtime settings take effect.\n",
"from jeanspy.model_jax import (\n",
" DSphModel, PlummerModel, ZhaoModel, ConstantAnisotropyModel,\n",
")\n",
"from jeanspy.sampler_numpyro import JeansLikelihoodModel, ParameterSpec\n",
"import jax\n",
"import numpyro.distributions as dist\n",
"from numpyro.infer import MCMC, NUTS, init_to_value\n",
"import numpy as np\n",
"import matplotlib.pyplot as plt\n",
"\n",
"model = DSphModel(submodels={\n",
" \"StellarModel\": PlummerModel(),\n",
" \"DMModel\": ZhaoModel(),\n",
" \"AnisotropyModel\": ConstantAnisotropyModel(),\n",
"})\n",
"fixed = dict(re_pc=200., alpha=1., beta=3., r_t_pc=np.inf)\n",
"truth = dict(**fixed, rs_pc=500., rhos_Msunpc3=.1, gamma=1., beta_ani=0.)\n",
"options = dict(solver=\"kernel\", n_u=128, n_kernel=32,\n",
" dm_mass_method=\"numeric\", dm_mass_n_steps=128)"
]
},
{
"cell_type": "code",
"execution_count": 2,
"id": "d754180a",
"metadata": {},
"outputs": [],
"source": [
"rng = np.random.default_rng(123)\n",
"u = rng.uniform(size=32)\n",
"R_pc = 200. * np.sqrt(u / (1. - u)) # Projected Plummer radii.\n",
"e_vlos_kms = np.full(32, 2.)\n",
"variance = np.asarray(model.sigmalos2(R_pc, params=truth, **options))\n",
"vlos_kms = rng.normal(0., np.sqrt(variance + e_vlos_kms**2))\n",
"data = dict(R_pc=R_pc, vlos_kms=vlos_kms, e_vlos_kms=e_vlos_kms)"
]
},
{
"cell_type": "code",
"execution_count": 3,
"id": "e5d3b986",
"metadata": {},
"outputs": [],
"source": [
"def physical_parameters(sampled):\n",
" return {**fixed, **sampled}\n",
"\n",
"likelihood = JeansLikelihoodModel(\n",
" model,\n",
" [ParameterSpec.pow10(\"log10_rs_pc\", dist.Uniform(1.5, 4.), param_name=\"rs_pc\"),\n",
" ParameterSpec.pow10(\"log10_rhos_Msunpc3\", dist.Uniform(-3., 1.),\n",
" param_name=\"rhos_Msunpc3\"),\n",
" ParameterSpec(\"gamma\", dist.Uniform(0., 2.)),\n",
" ParameterSpec(\"beta_ani\", dist.Uniform(-1., .75)),\n",
" ParameterSpec(\"vmem_kms\", dist.Uniform(-50., 50.))],\n",
" parameter_postprocess=physical_parameters,\n",
" sigmalos2_kwargs=options,\n",
")"
]
},
{
"cell_type": "code",
"execution_count": 4,
"id": "22d694f1",
"metadata": {},
"outputs": [],
"source": [
"kernel = NUTS(likelihood, target_accept_prob=.85,\n",
" init_strategy=init_to_value(values={\n",
" \"log10_rs_pc\": np.log10(500.), \"log10_rhos_Msunpc3\": -1.,\n",
" \"gamma\": 1., \"beta_ani\": 0., \"vmem_kms\": 0.}))\n",
"mcmc = MCMC(kernel, num_warmup=24, num_samples=32,\n",
" num_chains=2, chain_method=\"sequential\", progress_bar=False)\n",
"mcmc.run(jax.random.PRNGKey(42), **data)"
]
},
{
"cell_type": "markdown",
"id": "fc5e1db0",
"metadata": {},
"source": [
"## Export a finished Quickstart run\n",
"\n",
"The setup below runs the same likelihood with two chains, 24 warmup steps\n",
"and 32 posterior draws each. These reduced settings exercise storage only.\n",
"Preserve the chain axes in a NumPy file:"
]
},
{
"cell_type": "code",
"execution_count": 5,
"id": "5af3841d",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"(2, 32)\n"
]
}
],
"source": [
"from pathlib import Path\n",
"\n",
"from uuid import uuid4\n",
"run_id = uuid4().hex[:8]\n",
"output = Path(\"jeanspy-result-\" + run_id)\n",
"output.mkdir(exist_ok=False) # Choose a new directory for a new analysis.\n",
"np.savez(output / \"posterior.npz\",\n",
" **{name: np.asarray(value) for name, value in\n",
" mcmc.get_samples(group_by_chain=True).items()})\n",
"np.savez(output / \"sample_stats.npz\",\n",
" **{name: np.asarray(value) for name, value in\n",
" mcmc.get_extra_fields(group_by_chain=True).items()})\n",
"np.savez(output / \"observations.npz\", **data)\n",
"with np.load(output / \"posterior.npz\") as saved:\n",
" print(saved[\"beta_ani\"].shape) # (2, 32) for this storage example."
]
},
{
"cell_type": "markdown",
"id": "d7eb8b1b",
"metadata": {},
"source": [
"These files allow read-back of values; they do not contain a restartable\n",
"NUTS checkpoint. Keep the script defining the model, transforms and priors,\n",
"the random seed, numerical settings, source revision and environment with\n",
"the exported data. A collection of marginal intervals alone loses the joint\n",
"posterior and cannot be used for chain diagnostics.\n",
"\n",
"## Enable checkpointing before sampling\n",
"\n",
"For automatic storage and restart, construct a fresh MCMC object with the\n",
"same likelihood and wrap it in [`jeanspy.sampler_numpyro.NumPyroSampler`](https://gomeshun.github.io/jeanspy/dev/api/all.html):"
]
},
{
"cell_type": "code",
"execution_count": 6,
"id": "c1ee0a47",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"last_state.pkl\n",
"(2, 32)\n"
]
}
],
"source": [
"from jeanspy.sampler_numpyro import NumPyroSampler\n",
"\n",
"def make_mcmc():\n",
" return MCMC(kernel, num_warmup=24, num_samples=32,\n",
" num_chains=2, chain_method=\"sequential\", progress_bar=False)\n",
"\n",
"with NumPyroSampler(make_mcmc(), output_dir=\"jeanspy-checkpoint-\" + run_id,\n",
" storage_backend=\"h5netcdf\", async_writes=False) as stored:\n",
" result = stored.run(jax.random.PRNGKey(42), **data, resume=False)\n",
" print(result.checkpoint_path.name)\n",
" tree = stored.load_samples(combine=True)\n",
" posterior = tree[\"posterior\"].ds\n",
" print(posterior[\"beta_ani\"].shape)"
]
},
{
"cell_type": "markdown",
"id": "a4d26289",
"metadata": {},
"source": [
"Use a new directory for the first run. This launches sampling with storage\n",
"enabled; it does not import the already completed Quickstart draws. `kernel`\n",
"and `data` are defined in this notebook's setup. The context manager\n",
"flushes writes before closing. Setting `async_writes=False` makes the disk\n",
"write complete before `run` returns.\n",
"\n",
"## Reconstruct and resume the same analysis\n",
"\n",
"Recreate the same model, likelihood, kernel and observations, then open the\n",
"existing output directory:"
]
},
{
"cell_type": "code",
"execution_count": 7,
"id": "7bbb9f16",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Resumed: True\n",
"(2, 64)\n"
]
}
],
"source": [
"with NumPyroSampler(make_mcmc(), output_dir=\"jeanspy-checkpoint-\" + run_id,\n",
" storage_backend=\"h5netcdf\", async_writes=False) as stored:\n",
" result = stored.run(jax.random.PRNGKey(43), **data, resume=True)\n",
" print(\"Resumed:\", result.resumed)\n",
" tree = stored.load_samples(combine=True)\n",
" print(tree[\"posterior\"].ds[\"beta_ani\"].shape) # (2, 64) after two chunks."
]
},
{
"cell_type": "markdown",
"id": "09a367ac",
"metadata": {},
"source": [
"The checkpoint restores the continuation state, including its random state,\n",
"and skips completed warmup. A new seed argument is not a request for an\n",
"independent chain when resuming. `load_samples(combine=True)` concatenates\n",
"chunks along `draw`, preserving `chain`. In the locked environment it returns\n",
"an Xarray `DataTree`; use `tree[\"posterior\"].ds` to obtain the dataset.\n",
"\n",
"The wrapper checks data, model/prior configuration, computational source/dependencies and\n",
"effective JAX backend/precision against the stored analysis identity. A\n",
"changed analysis needs a new output directory. If an existing analysis lacks\n",
"`metadata.json`, restore the original metadata or start separately; do not\n",
"assign current inputs as a replacement identity for old samples.\n",
"\n",
"Identity format 2 allows documentation-only edits while retaining checks on\n",
"calculation code, data and numerical settings. Full source/data byte hashes are\n",
"kept separately in `metadata.json` under `source_provenance`. Format-1 chains\n",
"require their original code/environment to resume; see the\n",
"[API migration guide](../guides/api-migration.md#resume-compatibility-and-source-provenance).\n",
"\n",
"## Understand the stored files\n",
"\n",
"| File | Purpose |\n",
"| --- | --- |\n",
"| `metadata.json` | Analysis identity, source provenance and storage configuration |\n",
"| `last_state.pkl` | Sampler continuation state; load only trusted local checkpoints |\n",
"| `chunks/*.nc` with `h5netcdf` | Posterior draws and available sampler statistics |\n",
"| `observations.csv` in the worked example | Input data saved by the example, not by the wrapper automatically |\n",
"\n",
"The [NumPyroSampler](https://gomeshun.github.io/jeanspy/dev/api/all.html) API\n",
"describes the other storage backends and asynchronous-write options. To run\n",
"the complete save/reconstruct/resume example:\n",
"\n",
"```bash\n",
"python examples/docs_quickstart_numpyro.py --output-dir /tmp/jeanspy-numpyro-example\n",
"```\n",
"\n",
"That script deliberately requires a fresh directory, then performs both runs\n",
"itself. Its [recorded results](spherical.md) include a check that the first\n",
"chunk is unchanged after resuming.\n",
"\n",
"\n",
"\n",
"## Save emcee results\n",
"\n",
"`emcee.backends.HDFBackend` stores the chain and random state. Reopen the\n",
"backend and pass `None` as the initial state to append. The direct emcee route\n",
"requires you to preserve the model, priors and data yourself; JeansPy's\n",
"[`jeanspy.sampler.Sampler`](https://gomeshun.github.io/jeanspy/dev/api/all.html) adds identity checks for its estimation\n",
"models. See the {ref}`NumPy/SciPy wrapper example ` and\n",
"[complete spherical analysis](spherical.md).\n",
"\n",
"Continue with a [worked axisymmetric analysis](axisymmetric.md), or use the\n",
"[API dictionary](../api/index.md) for an individual method's arguments."
]
}
],
"metadata": {
"jeanspy": {
"execution": {
"code_sha256": "648d245f766d4ddb1d6ce10e2b142e89a3ef56fce03d1af4a92192350d429ab9",
"jax_enable_x64": true,
"lock_sha256": "a30f366620e23a49d6d7466055223a761824454169d8fc0364c24ac25b626dc7",
"packages": {
"arviz": "1.0.0",
"corner": "2.2.3",
"jax": "0.9.1",
"jeanspy": "0.1.0",
"matplotlib": "3.10.8",
"numpy": "2.4.3",
"numpyro": "0.20.0",
"scipy": "1.17.1"
},
"platform": "cpu",
"source_sha256": {
"src/jeanspy/__init__.py": "61671d182af04c0056cfb8d273db499d8774694711506698dbe7dc4a63d3b403",
"src/jeanspy/_axisymmetric_components.py": "4b203c616c87e55182d105b37f9fa49e44fe1df6bc430ccccb09edc5cca42530",
"src/jeanspy/_axisymmetric_params.py": "85da8de45bf5ba8cd9a40a137f7ccf971f7ce9534f312ef49fcd33539a230802",
"src/jeanspy/_jax_env.py": "2ef46eaf5cbd8e72a0c17351490207c29586979082cabc87128c9bd16fc689a4",
"src/jeanspy/_numpy/__init__.py": "de2d19d61e00702f13996582369f6acbe92b74628e6f569f91c3215a53580117",
"src/jeanspy/_numpy/core.py": "d073c6409d5f8cf527947c1ddf3916cb39cef3f6d4e43173cad59d148e463296",
"src/jeanspy/_numpy/inference.py": "ab58728400b6e56189ffba9ee805d9edbf59ca13dfebe05498801592314552db",
"src/jeanspy/_numpy/jfactor.py": "823f58a66a615ba59e057022ce51e2adfab577c25a4ced2f2033304c2cbd6cfe",
"src/jeanspy/_numpy/profiles.py": "e77fda6fb068392a63f0479e4c33bf80a22de3d2ab2cd101ea0c1f33e579491d",
"src/jeanspy/_numpy/solver.py": "b419796e2a912a46f65404cb5026ccb73799e1dbb8f80052f404148dd1f5ef3c",
"src/jeanspy/_sampling_identity.py": "b0e70b77e38f25e2148ab1f108a714512a8b6054405c846c3a41730b07081cec",
"src/jeanspy/_sersic_deprojection.py": "004c5566448698ecfe536096c3fc8d5c348a8d4e1610042763c77568c0f3d821",
"src/jeanspy/_zhao.py": "df83c7427dc583123c7464a03d2d45272e90f5a6724c9729eee64fad92e0c41a",
"src/jeanspy/axisymmetric.py": "683d67b91f8eb67e8d898eb23e6c8ca0e5c78020a65a7d05e183c19b076bd506",
"src/jeanspy/axisymmetric_factors.py": "754b256f7d3d43cc72571cde2b22bc9419917cf13b17342cd5f52406c4111336",
"src/jeanspy/axisymmetric_inference.py": "57078226b3643f8babda4d942b96ae546622da28802119d3d99ff6a0c5b45695",
"src/jeanspy/axisymmetric_jax.py": "de5f78937104ce7008242d55c321c8a3c6d9491f6aa25d4be3bdab8de076ed54",
"src/jeanspy/baes_eta2.py": "c274a16ca8c4b4d8d76ec30a2e705b5a5f6ce7659b31e31b5c685f92ab7a7ece",
"src/jeanspy/dequad.py": "6260e7ee2d8da12265573e86ff81ba0a3e8912df5fec8588a84bfda9d3b1133d",
"src/jeanspy/hyp2f1_jax.py": "03d8df3b48fcd9d418e289ea2ffaf174ffbdf14dffe483a150403c092bbaeea8",
"src/jeanspy/model.py": "72af25df8f34816f27d373fdecf41408c79eb505cb4cb52c830822d5d8f12f2b",
"src/jeanspy/model_jax.py": "355ba42ad16b2cadd00906679dd06ee12fd85fe4f1fc48630c779e68e4fabd64",
"src/jeanspy/parameters.py": "7760f4c85063105eb1328910867f3c57516f885dc577533bf9b91ac08a648cb2",
"src/jeanspy/sampler.py": "42ca0bbcef84b7c1620fd977d71810284a5d6abd70ff6047ae3175822d63f10e",
"src/jeanspy/sampler_numpyro.py": "ccd367d576bc55173d85939ad4bf7eeb8a44130d693f40ee56f4eade278c2195",
"src/jeanspy/sersic.py": "487f473351f9a309ad7d8dad1b1405e496196413e7606d0b0a18c5f71bfcf4a3"
},
"utc": "2026-09-18T00:32:26.413303+00:00"
},
"mcmc": true
},
"kernelspec": {
"display_name": "Python 3 (JeansPy)",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.12.13"
}
},
"nbformat": 4,
"nbformat_minor": 5
}