From d0237e5dd92f7501e244be8072568fab164ea124 Mon Sep 17 00:00:00 2001 From: quattro Date: Thu, 10 Sep 2026 09:58:48 -0700 Subject: [PATCH] Stabilize SPA root finding and tail evaluation --- docs/api/hypothesis/variant.md | 9 +-- src/jaxqtl/hypothesis/_spa.py | 61 ++++++++++++++----- tests/test_spa_bisection.py | 103 +++++++++++++++++++++++++++++++++ 3 files changed, 154 insertions(+), 19 deletions(-) create mode 100644 tests/test_spa_bisection.py diff --git a/docs/api/hypothesis/variant.md b/docs/api/hypothesis/variant.md index c494d45..a673855 100644 --- a/docs/api/hypothesis/variant.md +++ b/docs/api/hypothesis/variant.md @@ -92,10 +92,11 @@ methods must support JAX transformations; file loading, block packing, and outpu ## Saddlepoint approximation -SPA starts from the score test's null fit. Choose a CGF matching the model family. The implementation attempts -SPA above the normal-score cutoff and within its score-support checks; otherwise it uses a normal tail. -An unsuccessful SPA root solve or invalid correction also returns the normal approximation. The returned -`converged` field describes model fitting, not whether SPA was applied successfully. +SPA starts from the score test's null fit. Choose a CGF matching the model family. It uses bisection with +finite, sign-changing brackets constructed inside the CGF domain. The normal approximation is used when SPA +is not attempted under the score cutoff and support checks. An attempted SPA calculation that does not +converge or yields an invalid correction returns NaN; ACAT propagates NaN inputs. The returned `converged` +field describes model fitting, not whether SPA was applied successfully. ::: jaxqtl.hypothesis.SpaTest options: diff --git a/src/jaxqtl/hypothesis/_spa.py b/src/jaxqtl/hypothesis/_spa.py index b70bc0a..5b908bc 100644 --- a/src/jaxqtl/hypothesis/_spa.py +++ b/src/jaxqtl/hypothesis/_spa.py @@ -335,14 +335,17 @@ def saddlepoint_pvalue( **Returns:** A scalar p-value (or log p-value) for the observed score statistic. + Returns NaN when an attempted SPA calculation fails to converge or is invalid. """ is_discrete = isinstance(cgf, PoissonCGF | NegativeBinomialCGF) g_resid = jnp.asarray(g_resid, dtype=float) score = jnp.asarray(score, dtype=float) - solver = optx.Newton(rtol=1e-8, atol=1e-8) - t_bounds = cgf.get_t_bounds(g_resid, state) + tolerance = max(1e-8, 8 * float(jnp.finfo(g_resid.dtype).eps)) + # Bisection supplies its scalar norm as a ClassVar, not an init argument. + solver = optx.Bisection(rtol=tolerance, atol=tolerance, flip=False) # ty: ignore[missing-argument] + t_bounds = cgf.get_t_bounds(g_resid * scale, state) score_bounds = cgf.get_score_bounds(g_resid, state) offset = g_resid.T @ state.pred_mean @@ -355,7 +358,7 @@ def _closure(t): def _fn(t, args): (current_score,) = args _val, deriv = _closure(t) - return deriv - current_score + return (deriv - current_score) / jnp.maximum(1.0, jnp.abs(current_score)) _, (_, score_var) = jax.jvp(_closure, (0.0,), (1.0,)) zscore = score / jnp.sqrt(score_var) @@ -366,11 +369,38 @@ def _fn(t, args): should_attempt_spa = (jnp.fabs(zscore) > cutoff) & is_valid def _spa(current_score): - lower, upper = t_bounds + # K' is increasing. Search from zero toward the target, expanding on + # unbounded domains and approaching finite domain limits from inside. + # Nonfinite evaluations tighten the search limit instead of becoming + # bisection endpoints. A bounded loop also handles unreachable targets. + direction = jnp.sign(current_score) + limit = jnp.where(direction > 0, t_bounds[1], -t_bounds[0]) + distance = jnp.minimum(1.0, 0.5 * limit) + error = direction * _fn(direction * distance, (current_score,)) + + def searching(carry): + _, _, _, error, steps = carry + return (~jnp.isfinite(error) | (error < 0)) & (steps < 64) + + def expand(carry): + inner, limit, distance, error, steps = carry + finite = jnp.isfinite(error) + inner = jnp.where(finite, distance, inner) + limit = jnp.where(finite, limit, distance) + distance = jnp.minimum(2 * distance, inner + 0.5 * (limit - inner)) + error = direction * _fn(direction * distance, (current_score,)) + return inner, limit, distance, error, steps + 1 + + _, _, distance, error, _ = lax.while_loop( + searching, expand, (jnp.zeros_like(distance), limit, distance, error, jnp.array(0)) + ) + bracket_valid = jnp.isfinite(error) & (error >= 0) & (distance > 0) + endpoint = lax.stop_gradient(direction * distance) + lower, upper = jnp.minimum(0.0, endpoint), jnp.maximum(0.0, endpoint) sol = optx.root_find( _fn, solver, - 0.0, + 0.5 * (lower + upper), args=(current_score,), options={"lower": lower, "upper": upper}, has_aux=False, @@ -385,21 +415,22 @@ def _spa(current_score): under_radical = 2 * (t_bar * current_score - K_val) w = jnp.sign(t_bar) * jnp.sqrt(under_radical) - scale_factor = -jnp.expm1(-t_bar) if is_discrete else t_bar - v = scale_factor * jnp.sqrt(K_pp) - - ratio = v / w - r = w + jnp.log(ratio) / w - r = jnp.where(ratio <= 0, jnp.nan, r) + # Both v and w have the sign of t. Evaluate log(|v|/|w|) + # directly: exp(-t) overflows for valid, large negative roots. + abs_t = jnp.abs(t_bar) + log_factor = jnp.maximum(-t_bar, 0.0) + jax.nn.log1mexp(abs_t) if is_discrete else jnp.log(abs_t) + log_ratio = log_factor + 0.5 * jnp.log(K_pp) - jnp.log(jnp.abs(w)) + r = w + log_ratio / w t_result_lower = stats.norm.logcdf(r) t_result_upper = stats.norm.logsf(r) t_result_symm = jnp.log(2.0) + jnp.where(r <= 0.0, t_result_lower, t_result_upper) - w_is_valid = ~jnp.isnan(w) - r_is_valid = ~jnp.isnan(r) + w_is_valid = jnp.isfinite(w) + r_is_valid = jnp.isfinite(r) solver_success = sol.result == optx.RESULTS.successful - is_successful = w_is_valid & r_is_valid & solver_success + # Bisection checks both bracket width and residual for convergence. + is_successful = bracket_valid & w_is_valid & r_is_valid & solver_success return t_result_lower, t_result_upper, t_result_symm, is_successful @@ -417,7 +448,7 @@ def compute_spa_p_value(_): else: spa_result = jnp.log(2.0) + jnp.min(log_tails) - return jnp.where(is_successful, spa_result, log_p_normal_two_sided) + return jnp.where(is_successful, spa_result, jnp.nan) def compute_normal_p_value(_): return log_p_normal_two_sided diff --git a/tests/test_spa_bisection.py b/tests/test_spa_bisection.py new file mode 100644 index 0000000..e3a2131 --- /dev/null +++ b/tests/test_spa_bisection.py @@ -0,0 +1,103 @@ +# pattern: Functional Core +"""Regression checks for SPA roots in sparse and unbounded-domain models.""" + +import math + +import pytest + +from scipy.optimize import brentq +from scipy.special import ndtr + +import jax +import jax.numpy as jnp + +from jaxqtl.hypothesis._spa import ( + BasicCGFState, + GaussianCGF, + GaussianCGFState, + NegativeBinomialCGF, + NegBinCGFState, + PoissonCGF, + saddlepoint_pvalue, +) + + +@pytest.mark.parametrize("alpha", [1e-9, 0.0001911442384]) +@pytest.mark.parametrize("negative_weight", [-0.02, -0.005]) +@pytest.mark.parametrize("x64", [False, True]) +def test_sparse_nb_spa_matches_independently_bracketed_roots(alpha, negative_weight, x64): + with jax.enable_x64(x64): + mu = [0.0127, 0.05] + g = [1.0, negative_weight] + score = 0.9873 + r = 1 / alpha + center = sum(a * b for a, b in zip(mu, g)) + + def cumulants(t): + k, kp, kpp = -t * center, -center, 0.0 + for m, a in zip(mu, g): + e = math.exp(t * a) + d = 1 - m / r * math.expm1(t * a) + k -= r * math.log1p(-m / r * math.expm1(t * a)) + kp += a * m * e / d + kpp += a * a * m * e * (1 + m / r) / d**2 + return k, kp, kpp + + bounds = [0.9 * math.log1p(r / mu[1]) / negative_weight, 0.9 * math.log1p(r / mu[0])] + tails = [] + for target in [score, -score]: + t = brentq(lambda t: cumulants(t)[1] - target, *bounds) + k, _, kpp = cumulants(t) + w = math.copysign(math.sqrt(2 * (t * target - k)), t) + log_factor = math.log(-math.expm1(-t)) if t > 0 else -t + math.log1p(-math.exp(t)) + corrected = w + (log_factor + 0.5 * math.log(kpp) - math.log(abs(w))) / w + tails.append(ndtr(-corrected if target > 0 else corrected)) + observed = saddlepoint_pvalue( + score, + jnp.array(g), + NegativeBinomialCGF(), + NegBinCGFState(jnp.array(mu), jnp.array(r)), + two_sided_mode="abs", + ) + assert float(observed) == pytest.approx(sum(tails), rel=1e-6 if x64 else 2e-5) + + +def test_failed_spa_does_not_silently_return_normal_tail(): + with jax.enable_x64(True): + p = saddlepoint_pvalue( + 0.9873, + jnp.array([1.0, -0.02]), + NegativeBinomialCGF(), + NegBinCGFState(jnp.array([0.0127, 0.05]), jnp.array(1e9)), + two_sided_mode="abs", + max_iter=1, + ) + assert jnp.isnan(p) + + +@pytest.mark.parametrize("x64", [False, True]) +def test_unbounded_gaussian_spa_vmap_and_gradient(x64): + with jax.enable_x64(x64): + state = GaussianCGFState(jnp.zeros(2), jnp.full(2, 0.5)) + + def pvalue(score): + return saddlepoint_pvalue(score, jnp.array([1.0, -1.0]), GaussianCGF(), state) + + scores = jnp.array([0.5, 3.0, -4.0, 10.0]) + actual = jax.jit(jax.vmap(pvalue))(scores) + expected = 2 * jax.scipy.stats.norm.sf(jnp.abs(scores)) + assert jnp.allclose(actual, expected, rtol=2e-5, atol=0) + derivative = jax.jit(jax.grad(pvalue))(jnp.array(3.0)) + assert derivative == pytest.approx(float(-2 * jax.scipy.stats.norm.pdf(3.0)), rel=2e-5) + + +def test_sparse_poisson_spa_unbounded_domain(): + with jax.enable_x64(True): + p = saddlepoint_pvalue( + 0.9873, + jnp.array([1.0, -0.02]), + PoissonCGF(), + BasicCGFState(jnp.array([0.0127, 0.05])), + two_sided_mode="abs", + ) + assert float(p) == pytest.approx(0.01297173849793784, rel=1e-6)