{ "cells": [ { "cell_type": "markdown", "id": "e89fe867", "metadata": {}, "source": [ "# 3. Define and extend a model\n", "\n", "A spherical Jeans model combines three components: a tracer number density,\n", "a gravitating halo, and velocity anisotropy. A sampler and its priors are\n", "separate objects. Start by constructing a forward model and checking its\n", "predictions at fixed physical parameters.\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.\n", "\n", "## Compose the JAX model\n", "\n", "The JAX interface used by the Quickstart separates the components from their\n", "physical values. The keys in `submodels` are role names, so preserve their\n", "spelling even when replacing the concrete classes." ] }, { "cell_type": "code", "execution_count": 1, "id": "879b9e0b", "metadata": {}, "outputs": [], "source": [ "import os\n", "os.environ.setdefault(\"JEANSPY_JAX_PLATFORM\", \"cpu\")\n", "os.environ.setdefault(\"JEANSPY_JAX_ENABLE_X64\", \"true\")\n", "from jeanspy.model_jax import (\n", " DSphModel, PlummerModel, NFWModel, ConstantAnisotropyModel, get_runtime_config,\n", ")\n", "import jax\n", "import jax.numpy as jnp\n", "import numpy as np\n", "import matplotlib.pyplot as plt" ] }, { "cell_type": "code", "execution_count": 2, "id": "1f585465", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "{'jax_platform_requested': 'cpu', 'jax_platform_effective_env': 'cpu', 'jax_backend_active': 'cpu', 'jax_enable_x64': True, 'kernel_backend_default': 'jax', 'constant_kernel_n_quad_default': 32, 'sigmalos2_solver_default': 'auto', 'sigmalos2_jit_default': True, 'sigmalos2_kernel_outer_transform_default': 'sqrtlog', 'baes_kernel_n_quad_default': 32, 'baes_eta_recommended_max': 10.0, 'sigmalos2_n_u_default': 128, 'sigmalos2_n_r_default': 768, 'sigmalos2_u_max_default': 10000.0} [91.44464827 89.19126387 74.77863921]\n" ] } ], "source": [ "tracer, halo, anisotropy = PlummerModel(), NFWModel(), ConstantAnisotropyModel()\n", "model = DSphModel(submodels={\"StellarModel\": tracer, \"DMModel\": halo,\n", " \"AnisotropyModel\": anisotropy})\n", "params = dict(re_pc=300., rs_pc=500., rhos_Msunpc3=.1,\n", " r_t_pc=5000., beta_ani=-.2)\n", "R_pc = jnp.array([30., 100., 300.])\n", "variance = jax.block_until_ready(model.sigmalos2(\n", " R_pc, params=params, solver=\"kernel\", n_u=128, n_kernel=64))\n", "assert variance.shape == (3,) and np.all(np.asarray(variance) > 0)\n", "sampled_R = tracer.sample_R(jax.random.PRNGKey(20260913), 8, re_pc=300.)\n", "assert sampled_R.shape == (8,)\n", "print(get_runtime_config(), variance)\n", "#" ] }, { "cell_type": "markdown", "id": "b7bc3ca7", "metadata": {}, "source": [ "This notebook includes its own imports and runtime setup above.\n", "`params` holds scalar physical values. A new dictionary changes the prediction\n", "without mutating the model:" ] }, { "cell_type": "code", "execution_count": 3, "id": "8df1f278", "metadata": {}, "outputs": [], "source": [ "denser = {**params, \"rhos_Msunpc3\": 2 * params[\"rhos_Msunpc3\"]}\n", "variance_denser = model.sigmalos2(R_pc, params=denser, solver=\"kernel\")" ] }, { "cell_type": "markdown", "id": "10f789dc", "metadata": {}, "source": [ "Logarithms and other sampling transforms belong in `SamplingParameter` (NumPy/SciPy) or `ParameterSpec` (NumPyro), not in\n", "the physical dictionary. For instance, `rs_pc` is a radius in pc even when\n", "the sampler uses a coordinate called `log10_rs_pc`.\n", "\n", "## Generalize the NFW inner slope\n", "\n", "`NFWModel` fixes the inner and outer density slopes to 1 and 3. The\n", "Quickstart instead uses `ZhaoModel` with `alpha=1`, `beta=3` and a variable `gamma`:\n", "\n", "$$\n", "\\rho(r)=\\frac{\\rho_s}{(r/r_s)^gamma(1+r/r_s)^{3-gamma}}.\n", "$$\n", "\n", "Here `gamma=0` has a finite central density, `gamma=1` recovers NFW and `gamma=2`\n", "has a steeper cusp. `rhos_Msunpc3` is the normalization $\\rho_s$; the density\n", "at $r_s$ is $\\rho_s/2^{3-gamma}$. Keep `a` and `b` in the fixed physical\n", "dictionary, and sample `gamma` in `SamplingParameter` (NumPy/SciPy) or `ParameterSpec` (NumPyro). Use the differentiable\n", "numerical Zhao mass when fitting `gamma`; switching to `NFWModel` fixes this\n", "slope instead. The [inference notebook](inference.ipynb) demonstrates this fit.\n", "\n", "\n", "\n", "## Use the NumPy/SciPy interface\n", "\n", "NumPy/SciPy spherical components store their parameters. All required physical\n", "values must be supplied; omitted values can remain NaN. This example supplies\n", "every parameter and evaluates the assembled model:" ] }, { "cell_type": "code", "execution_count": 4, "id": "8e1a8cf9", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "[4.42046919 4.21841733 4.22716713]\n" ] } ], "source": [ "import numpy as np\n", "from jeanspy.model import ConstantAnisotropyModel, DSphModel, NFWModel, PlummerModel\n", "\n", "model = DSphModel(\n", " vmem_kms=0.,\n", " submodels={\n", " \"StellarModel\": PlummerModel(re_pc=200.),\n", " \"DMModel\": NFWModel(rs_pc=1000., rhos_Msunpc3=.01, r_t_pc=10000.),\n", " \"AnisotropyModel\": ConstantAnisotropyModel(beta_ani=0.),\n", " },\n", ")\n", "R_pc = np.array([50., 100., 300.])\n", "variance = model.sigmalos2(R_pc)\n", "sigma_kms = np.sqrt(variance)\n", "assert variance.shape == R_pc.shape\n", "assert np.all(np.isfinite(sigma_kms) & (sigma_kms > 0))\n", "print(sigma_kms)\n", "#" ] }, { "cell_type": "markdown", "id": "b0e1889b", "metadata": {}, "source": [ "Use `model[\"DMModel\"]` to access a component, and\n", "`model.update(rhos_Msunpc3=.02)` to change its density scale through the\n", "composite model. This mutation affects subsequent predictions. Use separate\n", "model instances when comparing parameter choices concurrently.\n", "\n", "The [profile guide](../guides/profiles.md) describes tracer scales,\n", "deprojection domains, anisotropy families and halo cutoffs. The [API catalogue](../api/spherical.rst)\n", "separates the NumPy and JAX implementations.\n", "\n", "## Implement your own tracer\n", "\n", "The following minimal NumPy/SciPy tracer implements the Plummer formula again\n", "so that it can be checked against a built-in reference. Replace these two\n", "consistent projected and spatial densities with your own profile once their\n", "normalization and projection relation are established." ] }, { "cell_type": "code", "execution_count": 5, "id": "be9c8d40", "metadata": {}, "outputs": [], "source": [ "import numpy as np\n", "from jeanspy.model import StellarModel, PlummerModel, DSphModel, NFWModel, ConstantAnisotropyModel\n", "\n", "class CustomPlummer(StellarModel):\n", " \"\"\"Unit-normalized Plummer tracer with projected half-light radius re_pc.\"\"\"\n", " required_param_names = [\"re_pc\"]\n", " required_models = {}\n", "\n", " def density_2d(self, R_pc):\n", " \"\"\"Projected number density in pc^-2; accepts scalar or array radii.\"\"\"\n", " re = self.params.re_pc\n", " return (1. + (np.asarray(R_pc) / re)**2)**-2 / (np.pi * re**2)\n", "\n", " def density_3d(self, r_pc):\n", " \"\"\"Spatial number density in pc^-3; accepts scalar or array radii.\"\"\"\n", " re = self.params.re_pc\n", " return 3. * (1. + (np.asarray(r_pc) / re)**2)**-2.5 / (4. * np.pi * re**3)\n", "\n", "custom = CustomPlummer(re_pc=200.)\n", "model = DSphModel(vmem_kms=0., submodels={\n", " \"StellarModel\": custom,\n", " \"DMModel\": NFWModel(rs_pc=500., rhos_Msunpc3=.1, r_t_pc=5000.),\n", " \"AnisotropyModel\": ConstantAnisotropyModel(beta_ani=0.),\n", "})\n", "R_pc = np.geomspace(10., 2000., 20)\n", "variance = model.sigmalos2(R_pc, n=256, n_kernel=64)\n", "#" ] }, { "cell_type": "markdown", "id": "2ac778ee", "metadata": {}, "source": [ "`required_param_names` declares the scalar parameters maintained by the\n", "NumPy/SciPy `Model` container. `required_models={}` means that the component has\n", "no nested components. The LOS solver needs both `density_2d` and `density_3d`;\n", "implement them for scalar and broadcast array inputs. Extra helpers such as a\n", "radial CDF or sampling method are needed only when your own workflow uses them.\n", "\n", "## Check a custom implementation before inference" ] }, { "cell_type": "code", "execution_count": 6, "id": "065ae4b4", "metadata": {}, "outputs": [], "source": [ "from scipy.integrate import quad\n", "builtin = PlummerModel(re_pc=200.)\n", "np.testing.assert_allclose(custom.density_2d(R_pc), builtin.density_2d(R_pc))\n", "np.testing.assert_allclose(custom.density_3d(R_pc), builtin.density_3d(R_pc))\n", "projected_total = quad(lambda R: 2*np.pi*R*custom.density_2d(R), 0., np.inf)[0]\n", "spatial_total = quad(lambda r: 4*np.pi*r**2*custom.density_3d(r), 0., np.inf)[0]\n", "np.testing.assert_allclose([projected_total, spatial_total], 1., rtol=1e-7)\n", "assert np.all(np.isfinite(variance) & (variance > 0))\n", "#" ] }, { "cell_type": "markdown", "id": "f2b89b55", "metadata": {}, "source": [ "The checks compare the two densities against the built-in Plummer expression\n", "and integrate each to unit tracer number. For a new profile, also check that\n", "projecting the 3-D density reproduces the 2-D density, inspect central and outer\n", "limits, and refine the Jeans quadrature. These checks do not establish the\n", "existence of a physical distribution function.\n", "\n", "For a custom NumPy/SciPy halo, implement a consistent mass density and enclosed\n", "mass with an explicit cutoff convention; consult\n", "[`jeanspy.model.DMModel`](https://gomeshun.github.io/jeanspy/dev/api/all.html). A custom anisotropy must implement the\n", "`beta`, `f` and `kernel` interfaces consistently; see\n", "[`jeanspy.model.AnisotropyModel`](https://gomeshun.github.io/jeanspy/dev/api/all.html).\n", "\n", "JAX extensions use JAX operations and explicit parameters. The current spherical\n", "JAX solver calls tracer densities as `density_2d(R, re_pc=...)` and\n", "`density_3d(r, re_pc=...)`: it does **not** pass an arbitrary tracer parameter\n", "dictionary. A tracer requiring additional sampled shape parameters therefore\n", "needs a compatible solver extension. A custom JAX halo must satisfy the\n", "`DMModel` mass interface; an all-JAX anisotropy kernel is needed for gradients.\n", "Defining an arbitrary Python class alone does not make it NUTS-compatible.\n", "\n", "## Extend the geometry\n", "\n", "Axisymmetric components store their physical values in immutable objects;\n", "use `dataclasses.replace` or the supported per-call parameter mapping to\n", "change them. Their flattening and cylindrical anisotropy differ from the\n", "spherical roles above. Follow the [axisymmetric guide](../guides/axisymmetric.md)\n", "and [worked analysis](axisymmetric.md) when changing geometry.\n", "\n", "Next: [predict observables and generate mocks](predictions.ipynb)." ] } ], "metadata": { "jeanspy": { "execution": { "code_sha256": "de6678e58e6a7d43dbdea2c75a76ba8a5ad4bfd367d1f2c53d3ce1c303a36753", "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:37.971031+00:00" }, "mcmc": false }, "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 }