Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 5 additions & 4 deletions docs/api/hypothesis/variant.md
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
61 changes: 46 additions & 15 deletions src/jaxqtl/hypothesis/_spa.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
Expand All @@ -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,
Expand All @@ -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

Expand All @@ -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
Expand Down
103 changes: 103 additions & 0 deletions tests/test_spa_bisection.py
Original file line number Diff line number Diff line change
@@ -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)