Source code for pyglmGamPoi.shrinkage

"""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])


[docs] def loc_median_fit( x: np.ndarray, y: np.ndarray, *, fraction: float = 0.1, npoints: Optional[int] = None, weighted: bool = True, ignore_zeros: bool = False, ) -> np.ndarray: """ Local median fit (glmGamPoi::loc_median_fit). Used to estimate dispersion trend across gene means. """ x = np.asarray(x, dtype=np.float64) y = np.asarray(y, dtype=np.float64) if x.shape != y.shape: raise ValueError("x and y must have the same shape") n = len(x) if n == 0: return np.array([], dtype=np.float64) if npoints is None: npoints = max(1, int(round(n * fraction))) npoints = max(1.0, min(float(n), float(npoints))) npoints_int = max(1, int(npoints)) order = np.argsort(x) x_sorted = x[order] y_sorted = y[order] half = npoints_int // 2 start = half + 1 end = n - half weights = stats.norm.pdf(np.linspace(-3, 3, half * 2 + 1)) if end < start: if weighted: w = stats.norm.pdf(np.linspace(-3, 3, n)) wm = _weighted_median(y_sorted, w) return np.full(n, wm, dtype=np.float64) return np.full(n, float(np.median(y_sorted)), dtype=np.float64) res = np.full(n, np.nan, dtype=np.float64) for idx in range(start - 1, end): lo = idx - half hi = idx + half + 1 selection = y_sorted[lo:hi] used_weights = weights if ignore_zeros: mask = selection != 0 selection = selection[mask] used_weights = weights[mask] if mask.any() else np.array([1.0]) if selection.size == 0: res[idx] = np.nan continue if weighted: res[idx] = _weighted_median(selection, used_weights) else: res[idx] = float(np.median(selection)) res[: max(0, start - 1)] = res[start - 1] res[min(n, end) :] = res[end - 1] out = np.empty(n, dtype=np.float64) out[order] = res return out
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"]), )