Axisymmetric synthetic analysis#

This example samples scale radius, density, cylindrical anisotropy, inclination and systemic velocity, and computes finite-cone J/D factors from a deterministic subset of saved posterior draws. Priors are explicit in log10_rhos_Msunpc3, log10_rs_pc, log10_one_minus_beta_z, cos_inclination and vmem_kms.

Apply the six tutorial steps to a flattened system. Unlike the spherical Quickstart, this complete script also handles projected flattening, inclination and posterior J/D-factor calculation.

1. Define the geometry and observations#

The tracer is an oblate Plummer component, and the halo is a spheroidal Zhao profile. The example fixes the projected tracer axis ratio and converts it to an intrinsic ratio for each inclination. Positions are signed x_pc, y_pc coordinates in pc. The velocity likelihood conditions on those positions; it does not infer the spatial selection function. The geometry reference defines inclination, deprojection and the physical parameter domains.

2. Choose priors and a sampler#

The prior table explicitly distinguishes logarithmic scales, transformed cylindrical anisotropy and cosine inclination. A uniform prior in cosine inclination is not uniform in inclination. The two commands below use the same physical configuration and mock-data recipe.

python examples/axisymmetric_inference.py \
    --sampler emcee --output-dir /tmp/axisymmetric-emcee
JEANSPY_JAX_ENABLE_X64=true python examples/axisymmetric_inference.py \
    --sampler numpyro --output-dir /tmp/axisymmetric-nuts

3. Save, resume and inspect#

Repeat the same command to resume. Keep the same warmup setting when resuming the emcee example so that its export discards the original warmup steps. Outputs include observations.csv, prior.csv, posterior.csv, derived_factors.csv, summary.json and persisted sampler state.

Default settings use eight stars and short chains. They exercise the complete storage path; the summary explicitly records that convergence and calibration are unestablished.

Inspect the saved chain as described in the MCMC tutorial before interpreting derived factors. The J/D factors are NumPy postprocessing of a deterministic subset of posterior draws, not differentiable outputs of the NUTS model. Their cone aperture and distance are part of the analysis. See the J/D-factor guide for cutoff, cusp and quadrature requirements.

Complete executable script
"""Reproducible axisymmetric inference, restart and posterior J/D example.

The defaults deliberately run short chains on eight synthetic stars to exercise
the workflow. They do not establish posterior convergence or calibration.
Repeat the same command/output directory to append to the saved chain.
"""
from __future__ import annotations

import argparse
from functools import partial
import json
from pathlib import Path

import numpy as np
import pandas as pd

from jeanspy.model import (
    AxisymmetricDSphModel, AxisymmetricPlummerModel,
    AxisymmetricZhaoModel, AxisymmetricConstantAnisotropyModel,
)
from jeanspy.parameters import SamplingParameter
from jeanspy.axisymmetric import intrinsic_axis_ratio
from jeanspy.axisymmetric_inference import AxisymmetricDSphEstimationModel, AxisymmetricKinematicData


FIXED = dict(re_pc=300., q_projected=.8, Q=.8, alpha=2., beta=3., gamma=.5, r_t_pc=3000.)
TRUTH = dict(**FIXED, rs_pc=500., rhos_Msunpc3=.1, beta_z=-.3,
             inclination=float(np.arccos(.3)), vmem_kms=0.)
PRIOR = pd.DataFrame(dict(lower=[-1.3, 2.5, .07, .1, -20.],
                          upper=[-.7, 2.85, .2, .6, 20.]),
                     index=["log10_rhos_Msunpc3", "log10_rs_pc", "log10_one_minus_beta_z",
                            "cos_inclination", "vmem_kms"])

PARAMETER_SPECS = [
    SamplingParameter("log10_rhos_Msunpc3", "rhos_Msunpc3", "pow10"),
    SamplingParameter("log10_rs_pc", "rs_pc", "pow10"),
    SamplingParameter("log10_one_minus_beta_z", "beta_z", "one_minus_pow10"),
    SamplingParameter("cos_inclination", "inclination", "arccos"),
    SamplingParameter("vmem_kms", "vmem_kms"),
]


def numpy_model(nodes):
    return AxisymmetricDSphModel(
        submodels={
            "StellarModel": AxisymmetricPlummerModel(
                re_pc=TRUTH["re_pc"],
                q=intrinsic_axis_ratio(TRUTH["q_projected"], TRUTH["inclination"])),
            "DMModel": AxisymmetricZhaoModel(**{name: TRUTH[name] for name in
                ("rs_pc", "rhos_Msunpc3", "Q", "alpha", "beta", "gamma", "r_t_pc")}),
            "AnisotropyModel": AxisymmetricConstantAnisotropyModel(TRUTH["beta_z"]),
        }, inclination=TRUTH["inclination"], n_force=nodes, n_vertical=nodes, n_los=nodes,
    )


def mock_data(stars, seed):
    rng = np.random.default_rng(seed)
    # Positions selected inside a finite photometric window. The inference
    # conditions on these observed positions rather than fitting their counts.
    probability = rng.uniform(.03, .8, stars)
    radius = TRUTH["re_pc"]*np.sqrt(probability/(1-probability))
    azimuth = rng.uniform(0., 2*np.pi, stars)
    x, y = radius*np.cos(azimuth), radius*np.sin(azimuth)*TRUTH["q_projected"]
    error = np.full(stars, 2.)
    variance = numpy_model(96).sigmalos2(x, y)
    velocity = rng.normal(TRUTH["vmem_kms"], np.sqrt(variance+error**2))
    return AxisymmetricKinematicData(x, y, velocity, error)


def run_emcee(args, estimation):
    from jeanspy.sampler import Sampler
    generator = partial(estimation.sample, rng=np.random.default_rng(args.seed+1))
    sampler = Sampler(estimation, generator, nwalkers=2*estimation.ndim+2,
                       prefix=str(args.output_dir/"chain_"))
    resumed = bool(sampler.backend.iteration > 0)
    if not resumed:
        sampler.sampler.random_state = np.random.RandomState(args.seed+2).get_state()
        sampler.burn_in(args.warmup, generator)
    sampler.run_mcmc(args.draws, 1, enable_convergence_check=False)
    frame = sampler.get_dataframe(discard=args.warmup)
    return frame[estimation.p_names_lnprob], dict(resumed=resumed,
        stored_steps=int(sampler.backend.iteration), walkers=sampler.nwalkers,
        mean_acceptance=float(np.mean(sampler.sampler.acceptance_fraction)),
        warmup_steps_to_discard=args.warmup)


def run_numpyro(args, data):
    import jax
    import jax.numpy as jnp
    import numpyro.distributions as dist
    from numpyro.infer import MCMC, NUTS, init_to_value
    from jeanspy.model_jax import (
        AxisymmetricDSphModel as JaxModel,
        AxisymmetricPlummerModel as JaxPlummerModel,
        AxisymmetricZhaoModel as JaxZhaoModel,
        AxisymmetricConstantAnisotropyModel as JaxAnisotropyModel,
    )
    from jeanspy.sampler_numpyro import AxisymmetricJeansLikelihoodModel, ParameterSpec, NumPyroSampler

    transforms = {"identity": None, "pow10": lambda x: 10.**x,
                  "one_minus_pow10": lambda x: 1-10.**x, "arccos": jnp.arccos}
    specifications = [
        ParameterSpec(spec.sample_name,
                      dist.Uniform(PRIOR.loc[spec.sample_name, "lower"],
                                   PRIOR.loc[spec.sample_name, "upper"]),
                      param_name=spec.param_name, transform=transforms[spec.transform])
        for spec in PARAMETER_SPECS
    ]
    forward = JaxModel(
        submodels={"StellarModel": JaxPlummerModel(), "DMModel": JaxZhaoModel(),
                   "AnisotropyModel": JaxAnisotropyModel()},
        n_force=args.nodes, n_vertical=args.nodes, n_los=args.nodes,
    )
    likelihood = AxisymmetricJeansLikelihoodModel(forward, specifications, fixed_params=FIXED)
    initial = dict(log10_rhos_Msunpc3=-1., log10_rs_pc=float(np.log10(500.)),
                   log10_one_minus_beta_z=float(np.log10(1.3)), cos_inclination=.3, vmem_kms=0.)
    kernel = NUTS(likelihood, max_tree_depth=4, init_strategy=init_to_value(values=initial))
    mcmc = MCMC(kernel, num_warmup=args.warmup, num_samples=args.draws,
                num_chains=1, progress_bar=False)
    with NumPyroSampler(mcmc, output_dir=args.output_dir/"numpyro",
                        storage_backend=args.storage, async_writes=False) as sampler:
        result = sampler.run(jax.random.PRNGKey(args.seed), **data.as_kwargs())
        posterior = sampler.load_samples()["posterior"].dataset
        frame = pd.DataFrame({name: posterior[name].values.reshape(-1) for name in PRIOR.index})
        diagnostics = dict(resumed=result.resumed, stored_draws=len(frame),
                            divergences_this_run=int(np.sum(mcmc.get_extra_fields()["diverging"])))
    return frame, diagnostics


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--sampler", choices=["emcee", "numpyro"], default="emcee",
                        help="emcee: NumPy/SciPy forward model; numpyro: JAX forward model with NUTS")
    parser.add_argument("--output-dir", type=Path, required=True)
    parser.add_argument("--stars", type=int, default=8)
    parser.add_argument("--nodes", type=int, default=24)
    parser.add_argument("--warmup", type=int, default=16)
    parser.add_argument("--draws", type=int, default=12)
    parser.add_argument("--seed", type=int, default=52)
    parser.add_argument("--storage", choices=["h5netcdf", "zarr", "netcdf4"], default="h5netcdf")
    args = parser.parse_args()
    if args.stars < 2 or args.draws < 1 or args.warmup < 1:
        parser.error("stars >= 2 and positive warmup/draws are required")
    data = mock_data(args.stars, args.seed)
    estimation = AxisymmetricDSphEstimationModel(data, PRIOR, fixed_params=FIXED,
                                                parameter_specs=PARAMETER_SPECS,
                    dsph_model=numpy_model(args.nodes))
    args.output_dir.mkdir(parents=True, exist_ok=True)
    frame, diagnostics = (run_emcee(args, estimation) if args.sampler == "emcee"
                           else run_numpyro(args, data))
    frame.to_csv(args.output_dir/"posterior.csv", index=False)
    pd.DataFrame(data.as_kwargs()).to_csv(args.output_dir/"observations.csv", index=False)
    PRIOR.to_csv(args.output_dir/"prior.csv")
    # A small deterministic subset illustrates postprocessing a chain from
    # either backend with the same finite-cone factors and physical parameters.
    derived = []
    for index in np.unique(np.linspace(0, len(frame)-1, min(4, len(frame)), dtype=int)):
        physical = estimation.convert_params(frame.iloc[index].to_numpy())
        derived.append(dict(row=int(index),
            J_GeV2_cm_minus5=estimation.dsph_model.jfactor(80000., .5, params=physical),
            D_GeV_cm_minus2=estimation.dsph_model.dfactor(80000., .5, params=physical)))
    pd.DataFrame(derived).to_csv(args.output_dir/"derived_factors.csv", index=False)
    summary = dict(sampler=args.sampler, seed=args.seed, stars=args.stars, nodes=args.nodes,
                    truth=TRUTH, diagnostics=diagnostics,
                    note="Short workflow demonstration; convergence and scientific calibration are not established.")
    (args.output_dir/"summary.json").write_text(json.dumps(summary, indent=2)+"\n")
    print(json.dumps(summary, indent=2))


if __name__ == "__main__":
    main()

Forward predictions and gradients#

For a smaller calculation before running inference, this example evaluates an axisymmetric JAX prediction and its derivative with respect to the halo density scale. It checks the derivative against the known linear scaling. Follow the backend tutorial for runtime configuration and the numerical accuracy guide for refinement checks.

"""A synchronized JAX prediction and physical-parameter derivative."""
# example-start
# Start Python with JEANSPY_JAX_ENABLE_X64=true for reference precision.
import jax
import jax.numpy as jnp
import numpy as np
from jeanspy.axisymmetric_jax import AxisymmetricDSphModel

model = AxisymmetricDSphModel(24, 24, 24)
fixed = dict(re_pc=300., rs_pc=500., q=.7, Q=.8,
             alpha=2., beta=3., gamma=.5, beta_z=-.3, inclination=1.1)

@jax.jit
def prediction(rho):
    return model.sigmalos2(jnp.array([100., 300.]), 50.,
                          params={**fixed, "rhos_Msunpc3": rho})

rho = .1
variance = jax.block_until_ready(prediction(rho))
derivative = jax.block_until_ready(jax.jacfwd(prediction)(rho))
# Gravity and second moments scale linearly with density at fixed geometry.
np.testing.assert_allclose(derivative, variance / rho, rtol=1e-5)
print(variance, derivative)
# example-end