"""Overdispersion shrinkage (glmGamPoi::overdispersion_shrinkage)."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Optional, Union
import numpy as np
from scipy import optimize, stats
def _weighted_median(values: np.ndarray, weights: np.ndarray) -> float:
order = np.argsort(values)
values = values[order]
weights = weights[order]
cum = np.cumsum(weights)
if cum[-1] <= 0:
return float(np.median(values))
idx = np.searchsorted(cum, 0.5 * cum[-1])
idx = min(int(idx), len(values) - 1)
return float(values[idx])
def variance_prior(
s2: np.ndarray,
df: Union[float, np.ndarray],
*,
covariate: Optional[np.ndarray] = None,
abundance_trend: Optional[bool] = None,
) -> dict[str, np.ndarray | float]:
"""
Empirical Bayes variance prior (glmGamPoi::variance_prior / limma squeezeVar).
"""
s2 = np.asarray(s2, dtype=np.float64)
df_arr = np.broadcast_to(np.asarray(df, dtype=np.float64), s2.shape).copy()
if np.all(np.isnan(s2)):
return {
"variance0": np.full_like(s2, np.nan),
"df0": np.nan,
"var_post": s2.copy(),
}
valid = np.isfinite(s2) & (s2 > 0) & np.isfinite(df_arr) & (df_arr > 0)
s2_sub = s2[valid]
df_sub = df_arr[valid] if df_arr.size == s2.size else float(df_arr.flat[0])
if s2_sub.size == 0:
return {"variance0": np.full_like(s2, np.nan), "df0": np.nan, "var_post": s2.copy()}
if np.allclose(s2_sub, 1.0):
return {"variance0": np.ones_like(s2), "df0": np.inf, "var_post": s2.copy()}
def nll(params: np.ndarray) -> float:
log_v0, log_df0 = params
v0 = np.exp(log_v0)
df0 = np.exp(log_df0)
d = df_sub if np.ndim(df_sub) else np.full_like(s2_sub, df_sub)
return -float(
np.sum(stats.f.logpdf(s2_sub / v0, dfn=d, dfd=df0) - log_v0)
)
opt = optimize.minimize(nll, x0=np.array([0.0, 0.0]), method="L-BFGS-B")
variance0 = np.full(s2_sub.shape, float(np.exp(opt.x[0])), dtype=np.float64)
df0 = float(np.exp(opt.x[1]))
use_trend = abundance_trend
if use_trend is None and covariate is not None:
use_trend = s2_sub.size >= 10 and np.sum(np.asarray(covariate)[valid] > 1e-8) >= 10
if use_trend and covariate is not None and s2_sub.size >= 10:
cov_sub = np.asarray(covariate, dtype=np.float64)[valid]
not_zero = cov_sub > 1e-8
if not_zero.sum() >= 10:
log_cov = np.log(cov_sub[not_zero])
lo, hi = log_cov.min(), log_cov.max()
knots = lo + np.array([1 / 3, 2 / 3]) * (hi - lo)
# Natural cubic spline basis (df=4) via truncated power basis approximation
design = np.column_stack(
[log_cov, np.maximum(0, log_cov - knots[0]) ** 3, np.maximum(0, log_cov - knots[1]) ** 3]
)
design = np.column_stack([np.ones(len(log_cov)), design])
try:
coef = np.linalg.lstsq(design, np.log(s2_sub[not_zero]), rcond=None)[0]
def nll_trend(par: np.ndarray) -> float:
betas = par[:-1]
df0_t = np.exp(par[-1])
v0 = np.exp(design @ betas)
d = df_sub if np.ndim(df_sub) else np.full_like(s2_sub[not_zero], df_sub)
return -float(
np.sum(stats.f.logpdf(s2_sub[not_zero] / v0, dfn=d, dfd=df0_t) - np.log(v0))
)
init = np.concatenate([coef, [np.log(df0)]])
opt2 = optimize.minimize(nll_trend, x0=init, method="L-BFGS-B", options={"maxiter": 5000})
variance0 = np.exp(design @ opt2.x[:-1])
df0 = float(np.exp(opt2.x[-1]))
except (np.linalg.LinAlgError, ValueError):
pass
var_post = np.full_like(s2, np.nan, dtype=np.float64)
var_post[valid] = (df0 * variance0 + df_sub * s2_sub) / (df0 + df_sub)
variance0_full = np.full_like(s2, np.nan, dtype=np.float64)
variance0_full[valid] = variance0 if variance0.size == s2_sub.size else float(variance0.flat[0])
return {"variance0": variance0_full, "df0": df0, "var_post": var_post}
@dataclass
class OverdispersionShrinkageResult:
dispersion_trend: np.ndarray
ql_disp_estimate: np.ndarray
ql_disp_trend: np.ndarray
ql_disp_shrunken: np.ndarray
ql_df0: float
[docs]
def overdispersion_shrinkage(
disp_est: np.ndarray,
gene_means: np.ndarray,
df: Union[float, np.ndarray],
*,
disp_trend: Union[bool, np.ndarray, None] = True,
ql_disp_trend: Optional[bool] = None,
npoints: Optional[int] = None,
) -> OverdispersionShrinkageResult:
"""
Shrink overdispersion estimates (glmGamPoi::overdispersion_shrinkage).
``glm_gp`` uses ``dispersion_trend`` (local median of raw MLE) for the second
beta fit pass; ``ql_disp_shrunken`` is returned for downstream DE workflows.
"""
disp_est = np.asarray(disp_est, dtype=np.float64)
gene_means = np.asarray(gene_means, dtype=np.float64)
if disp_est.shape != gene_means.shape:
raise ValueError("disp_est and gene_means must have the same shape")
est_value = np.isfinite(disp_est) & np.isfinite(gene_means) & np.isfinite(df)
trend = np.full_like(disp_est, np.nan, dtype=np.float64)
if disp_trend is None or disp_trend is True:
if npoints is None:
npoints = max(0.1 * int(est_value.sum()), 100)
trend[est_value] = loc_median_fit(
gene_means[est_value],
disp_est[est_value],
npoints=npoints,
)
elif disp_trend is False:
mean_disp = float(np.nanmean(disp_est[est_value])) if est_value.any() else np.nan
trend[est_value] = mean_disp
else:
trend = np.asarray(disp_trend, dtype=np.float64)
with np.errstate(divide="ignore", invalid="ignore"):
ql_disp = (1.0 + gene_means * disp_est) / (1.0 + gene_means * trend)
if ql_disp_trend is None:
ql_disp_trend = int(est_value.sum()) >= 100
var_pr = variance_prior(
ql_disp,
df,
covariate=gene_means,
abundance_trend=ql_disp_trend,
)
return OverdispersionShrinkageResult(
dispersion_trend=trend,
ql_disp_estimate=ql_disp,
ql_disp_trend=np.asarray(var_pr["variance0"], dtype=np.float64),
ql_disp_shrunken=np.asarray(var_pr["var_post"], dtype=np.float64),
ql_df0=float(var_pr["df0"]),
)