From 485dfae93876e80c1e1c22da2c6a62b4bbb1568f Mon Sep 17 00:00:00 2001 From: Intron7 Date: Wed, 23 Sep 2026 12:21:31 +0200 Subject: [PATCH] Preserve dispersion regularization in optimizer fallback --- src/pydeseq2/dispersions.py | 12 ++++++++- tests/test_dispersions.py | 51 +++++++++++++++++++++++++++++++++++++ 2 files changed, 62 insertions(+), 1 deletion(-) create mode 100644 tests/test_dispersions.py diff --git a/src/pydeseq2/dispersions.py b/src/pydeseq2/dispersions.py index 9a69da73..e7e94bbc 100644 --- a/src/pydeseq2/dispersions.py +++ b/src/pydeseq2/dispersions.py @@ -223,7 +223,17 @@ def dloss(log_alpha: float) -> float: else: return ( np.exp( - grid_fit_alpha(counts, design_matrix, mu, alpha_hat, min_disp, max_disp) + grid_fit_alpha( + counts, + design_matrix, + mu, + alpha_hat, + min_disp, + max_disp, + prior_disp_var=prior_disp_var, + cr_reg=cr_reg, + prior_reg=prior_reg, + ) ), res.success, ) diff --git a/tests/test_dispersions.py b/tests/test_dispersions.py new file mode 100644 index 00000000..700bf66b --- /dev/null +++ b/tests/test_dispersions.py @@ -0,0 +1,51 @@ +import numpy as np +import pytest +from scipy.optimize import OptimizeResult +from scipy.optimize import minimize_scalar +from scipy.stats import nbinom + +from pydeseq2 import dispersions + + +@pytest.mark.parametrize("optimizer", ["BFGS", "L-BFGS-B"]) +@pytest.mark.parametrize("cr_reg", [False, True]) +@pytest.mark.parametrize("prior_reg", [False, True]) +def test_dispersion_fallback_preserves_objective( + monkeypatch, optimizer, cr_reg, prior_reg +): + counts = np.array([1, 2, 4, 8, 2, 10, 15, 30]) + design = np.column_stack([np.ones(8), np.repeat([0, 1], 4)]) + mu = np.repeat([5.0, 15.0], 4) + alpha_hat, prior_var = 0.05, 0.03 + bounds = np.log([1e-4, 10.0]) + monkeypatch.setattr( + dispersions, "minimize", lambda *args, **kwargs: OptimizeResult(success=False) + ) + + def objective(log_alpha): + alpha = np.exp(log_alpha) + loss = -nbinom.logpmf(counts, 1 / alpha, 1 / (1 + mu * alpha)).sum() + if cr_reg: + weights = mu / (1 + mu * alpha) + loss += 0.5 * np.linalg.slogdet((design.T * weights) @ design)[1] + if prior_reg: + loss += (log_alpha - np.log(alpha_hat)) ** 2 / (2 * prior_var) + return loss + + expected = minimize_scalar(objective, bounds=bounds, method="bounded") + actual, converged = dispersions.fit_alpha_mle( + counts, + design, + mu, + alpha_hat, + 1e-4, + 10.0, + prior_disp_var=prior_var, + cr_reg=cr_reg, + prior_reg=prior_reg, + optimizer=optimizer, + ) + + assert expected.success and not converged + assert np.log(actual) == pytest.approx(expected.x, abs=0.002) + assert objective(np.log(actual)) <= expected.fun + 1e-4