Numerical accuracy#

The solvers use fixed quadrature rules without an automatic error estimate. Compare increasing node counts and an independent reference when using a new parameter domain. See the API reference for each solver’s numerical options.

Axisymmetric quadrature#

Force, vertical-pressure and LOS integrals use fixed Gauss–Legendre rules. Rational maps cover infinite vertical and LOS ranges. LOS quadrature is centered at maximum tracer density. Defaults are 96 nodes per integral; numerical settings must be refined for the actual parameter regime, especially sharp truncation, extreme scale ratios/cusps, flattening and outer positions. There is no global adaptive accuracy guarantee. Cost grows approximately as N_stars*n_force*n_vertical*n_los. NumPy processes one sky position at a time; JAX uses lax.map and rematerialization to bound intermediate memory during parameter differentiation.

The three node counts control the force and Jeans integrals. J/D factors have separate quadrature settings.

Precision and physical-parameter gradients#

Use the backend tutorial to configure JAX, inspect the effective device/precision and measure synchronized execution. CPU float64, CPU float32 and actual GPU conditions need separate accuracy assessments. Check derivatives against finite differences at several step sizes and quadrature orders over the intended parameter domain. The model contract describes differentiability boundaries.

Specialized spherical kernels#

The specialized hypergeometric and fixed-\(\eta=2\) Baes kernels support the spherical JAX solver. They are real-valued approximations with restricted domains; they do not replace a general complex special-function library. The following check uses SciPy and mpmath references in a small regular domain.

import jax
import jax.numpy as jnp
from scipy.special import hyp2f1
from jeanspy.hyp2f1_jax import hyp2f1_1b_3half
from jeanspy.baes_eta2 import (baes_eta2_kernel_jax,
                               baes_eta2_kernel_appell_reference)

w = jnp.array([.05, .2, .5])
values = jax.block_until_ready(hyp2f1_1b_3half(.8, w))
np.testing.assert_allclose(values, hyp2f1(1., .8, 1.5, np.asarray(w)), rtol=1e-5)
u, R_pc = np.array([1.1, 1.5]), np.array([.7, 2.3])
kernel = baes_eta2_kernel_jax(u, R_pc, -.5, .7, 1.4, n_kernel=128)
reference = baes_eta2_kernel_appell_reference(u, R_pc, -.5, .7, 1.4, dps=25)
np.testing.assert_allclose(kernel, reference, rtol=8e-4, atol=2e-6)
# This small reference domain is not a universal prior-domain accuracy claim.

baes_eta2_kernel_jax clamps some intermediate arguments and replaces nonfinite values. Validate physical parameters separately: a finite low-level kernel output does not certify a valid anisotropy model.