OSCR

Uncovering causal relationships in single-cell omic studies with causarray.

Code ↔ Paper

17 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 17 matches · 2 of them tie a paragraph to a whole file, not to given lines: weak matches, whose lines are not tinted
  1. [1] § Materials and methods › Semiparametric estimation ↔ causarray/DR_estimation.py, lines 945–980 · score 0.92 · max_features, max_samples, ccp_alpha, cross validation, class weights, leaf
  2. [2] § Materials and methods › Semiparametric estimation ↔ paper/methods/causarray/DR_estimation.py, lines 207–242 · score 0.92 · max_features, max_samples, ccp_alpha, cross validation, class weights, leaf
  3. [3] § Results › Simulation study demonstrates the advantages of causarray ↔ paper/simu_nb/Plot.ipynb, lines 105–174 · score 0.72 · RUV III NB, CINEMA OT, CoCoA, DESeq2, TPRs, Mixscape
  4. [4] § Results › Simulation study demonstrates the advantages of causarray ↔ paper/simu_poi/Plot.ipynb, lines 151–226 · score 0.71 · RUV III NB, CINEMA OT, CoCoA, DESeq2, Mixscape, Wilcoxon
  5. [5] § Results › Simulation study demonstrates the advantages of causarray ↔ paper/simu_nb/simu_nb_plot.py, lines 141–201 · score 0.70 · RUV III NB, CINEMA OT, CoCoA, DESeq2, Mixscape, Wilcoxon
  6. [6] § Results › Simulation study demonstrates the advantages of causarray ↔ paper/simu_nb/simu_nb_plot.py, lines 141–201 · score 0.60 · RUV III NB, CINEMA OT, CoCoA, ARI, ASW, scores
  7. [7] § Results › Simulation study demonstrates the advantages of causarray ↔ paper/simu_nb/Plot.ipynb, lines 105–174 · score 0.60 · RUV III NB, CINEMA OT, CoCoA, ARI, ASW, confounder
  8. [8] § Materials and methods › Counterfactual imputation and inference ↔ causarray/DR_estimation.py, lines 901–936 · score 0.60 · augmented inverse probability, propensity score, weighted, treatment
  9. [9] § Materials and methods › Counterfactual imputation and inference ↔ paper/methods/causarray/DR_estimation.py, lines 162–198 · score 0.60 · augmented inverse probability, propensity score, weighted, treatment
  10. [10] § Results › An in vivo Perturb-seq study › Functional analysis ↔ paper/AD/GO.R, the whole file · a weak match · score 0.60 · log fold change, ROSMAP AD, SEA AD, MTG, PFC, discovered
  11. [11] § Materials and methods › The probabilistic modeling of confounders ↔ paper/methods/causarray/gcate.py, lines 158–246 · score 0.59 · negative binomial distributions, dispersion parameter, residual, nuisance, rank, log
  12. [12] § Results › Alzheimer’s disease case-control study › An integrative analysis of excitatory neurons ↔ paper/AD/Plot.ipynb, lines 403–425 · score 0.57 · ROSMAP AD, SEA AD, GO terms, MTG, PFC, DE
  13. [13] § Materials and methods › The probabilistic modeling of confounders ↔ causarray/gcate.py, lines 203–270 · score 0.55 · negative binomial distributions, dispersion parameter, variation, GLMs, latent, confounders
  14. [14] § Results › An in vivo Perturb-seq study › Functional analysis ↔ paper/perturbseq/3-GO.R, lines 154–225 · score 0.54 · top GO terms, Satb2 perturbation, enrich, DE, RUV, causarray
  15. [15] § Materials and methods › The probabilistic modeling of confounders ↔ paper/methods/R_functions.R, lines 19–87 · score 0.53 · bulk gene expression, confounder adjustment, unmeasured confounders, LFC
  16. [16] § Results › An in vivo Perturb-seq study › Functional analysis ↔ paper/AD/Plot.ipynb, lines 194–257 · score 0.52 · fold change, SEA AD, scatter, slope, MTG, PFC
  17. [17] § Results › Alzheimer’s disease case-control study › An integrative analysis of excitatory neurons ↔ paper/AD/GO.R, the whole file · a weak match · score 0.51 · ROSMAP AD, SEA AD, MTG, PFC, discoveries, GO

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 1,068 lines · 44 KB · MIT · 2 matches

  1. import numpy as np
  2. import pandas as pd
  3. from sklearn.linear_model import LogisticRegression
  4. from sklearn.tree import DecisionTreeClassifier, DecisionTreeRegressor
  5. from sklearn_ensemble_cv import reset_random_seeds, Ensemble, ECV
  6. from causarray.gcate_glm import fit_glm
  7. import causarray.gcate_glm as _gcate_glm # module-qualified so _USE_FAST_BACKEND changes take effect at call time
  8. from causarray.utils import *
  9. from causarray.utils import _filter_params
  10. from joblib import Parallel, delayed
  11. from tqdm import tqdm
  12. import pprint
  13. import warnings
  14. from collections.abc import Mapping
  15. from sklearn.model_selection import KFold, ShuffleSplit
  16. def _get_func_ps(ps_model, **kwargs):
  17. if ps_model=='random_forest_cv':
  18. params_ps = _filter_params(fit_rf, kwargs)
  19. func_ps = lambda X, Y, X_test:fit_rf_ind_ps(X, Y[:,None], X_test=X_test, **params_ps)[:,0]
  20. elif ps_model=='logistic':
  21. clf_ps = LogisticRegression
  22. kwargs = {**{'fit_intercept':False, 'C':1e0, 'class_weight':None, 'random_state':0}, **kwargs}
  23. params_ps = _filter_params(clf_ps().get_params(), kwargs)
  24. func_ps = lambda X, Y, X_test: clf_ps(**params_ps).fit(X, Y).predict_proba(X_test)[:,1]
  25. elif ps_model=='ensemble':
  26. kwargs = dict(kwargs)
  27. logistic_class_weight = kwargs.pop('class_weight', None)
  28. params_ps_rf = _filter_params(fit_rf, kwargs)
  29. clf_ps = LogisticRegression
  30. kwargs = {**{
  31. 'fit_intercept': False, 'C': 1e0,
  32. 'class_weight': logistic_class_weight, 'random_state': 0,
  33. }, **kwargs}
  34. params_ps_lr = _filter_params(clf_ps().get_params(), kwargs)
  35. params_ps = {'params_ps_rf':params_ps_rf, 'params_ps_lr':params_ps_lr}
  36. func_ps = lambda X, Y, X_test:(fit_rf_ind(X, Y[:,None], X_test=X_test, **params_ps_rf)[:,0] + clf_ps(**params_ps_lr).fit(X, Y).predict_proba(X_test)[:,1])/2
  37. else:
  38. raise ValueError('Invalid propensity score model.')
  39. return func_ps, params_ps
  40. def _validate_clip(clip):
  41. """Validate a propensity clipping bound and return it as ``(lower, upper)``.
  42. Raises
  43. ------
  44. ValueError
  45. If ``clip`` is not a pair of numbers satisfying
  46. ``0 <= lower < upper <= 1``.
  47. """
  48. message = 'clip must be None or a pair 0 <= lower < upper <= 1'
  49. if isinstance(clip, (str, bytes)) or np.ndim(clip) != 1:
  50. raise ValueError(message)
  51. try:
  52. lower, upper = (float(bound) for bound in clip)
  53. except (TypeError, ValueError):
  54. raise ValueError(message) from None
  55. if not 0 <= lower < upper <= 1:
  56. raise ValueError(message)
  57. return lower, upper
  58. def estimate_propensity_scores(
  59. A, X_A, K=1, ps_model='logistic', mask=None, clip=None,
  60. random_state=0, verbose=False, class_weight=None, **kwargs,
  61. ):
  62. """Estimate per-treatment propensity scores.
  63. Each treatment is compared with the shared all-zero control group. With
  64. ``K > 1``, every returned score is predicted by a model that did not train
  65. on that cell. Logistic models return calibrated treatment probabilities
  66. (``class_weight=None``) by default, matching :func:`LFC`. Pass
  67. ``class_weight='balanced'`` to reproduce pre-0.1.0 fits, whose scores are
  68. centred near 0.5 regardless of prevalence.
  69. .. versionchanged:: 0.1.0
  70. Default ``class_weight`` changed from ``'balanced'`` to ``None``.
  71. Parameters
  72. ----------
  73. A : array-like, shape (n,) or (n, a)
  74. Binary treatment indicators. Rows containing only zeros are controls.
  75. X_A : array-like, shape (n, d_A)
  76. Covariates used by the propensity model, including an intercept column
  77. when ``fit_intercept=False``.
  78. K : int, optional
  79. Number of folds. ``1`` fits and predicts on all eligible cells;
  80. values greater than one produce out-of-fold predictions.
  81. ps_model : {'logistic', 'random_forest_cv', 'ensemble'}, optional
  82. Propensity model.
  83. mask : array-like or None, shape (n,) or (n, a)
  84. Optional per-treatment eligibility mask for model fitting.
  85. clip : tuple(float, float) or None, optional
  86. Bounds applied after prediction. ``None`` returns raw probabilities.
  87. random_state : int, optional
  88. Random seed used for fold construction and supported estimators.
  89. class_weight : str, dict or None, optional
  90. Class weighting for logistic propensity estimation. ``None`` (default)
  91. gives calibrated probabilities and matches :func:`LFC`;
  92. ``'balanced'`` reproduces the pre-0.1.0 behaviour.
  93. Returns
  94. -------
  95. pi_hat : ndarray, shape (n, a)
  96. Estimated probabilities ``P(A_j=1 | X_A)``.
  97. """
  98. A = np.asarray(A)
  99. if A.ndim == 1:
  100. A = A[:, None]
  101. X_A = np.asarray(X_A)
  102. if A.ndim != 2 or X_A.ndim != 2 or A.shape[0] != X_A.shape[0]:
  103. raise ValueError('A and X_A must be two-dimensional with matching rows')
  104. if not np.all(np.isin(A, (0, 1))):
  105. raise ValueError('A must contain only binary treatment indicators')
  106. try:
  107. K_int = int(K)
  108. except (TypeError, ValueError) as exc:
  109. raise ValueError('K must be a positive integer') from exc
  110. if K_int != K:
  111. raise ValueError('K must be a positive integer')
  112. K = K_int
  113. if K < 1:
  114. raise ValueError('K must be a positive integer')
  115. if K > A.shape[0]:
  116. raise ValueError('K cannot exceed the number of samples')
  117. if mask is not None:
  118. mask = np.asarray(mask, dtype=bool)
  119. if mask.ndim == 1:
  120. mask = mask[:, None]
  121. if mask.shape != A.shape:
  122. raise ValueError('Mask must have the same shape as the treatment matrix')
  123. func_ps, params_ps = _get_func_ps(
  124. ps_model, verbose=False, random_state=random_state,
  125. class_weight=class_weight, **kwargs)
  126. if verbose:
  127. pprint.pprint(params_ps)
  128. if ps_model == 'random_forest_cv':
  129. info_ecv = run_ecv(X_A, A, **params_ps)
  130. func_ps, params_ps = _get_func_ps(
  131. ps_model, verbose=False, ecv=False,
  132. kwargs_ensemble=info_ecv['best_params_ensemble'],
  133. kwargs_regr=info_ecv['best_params_regr'],
  134. )
  135. if verbose:
  136. pprint.pprint('Best parameters for the regression model:')
  137. pprint.pprint(info_ecv['best_params_regr'])
  138. pprint.pprint('Best parameters for the ensemble model:')
  139. pprint.pprint(info_ecv['best_params_ensemble'])
  140. n = A.shape[0]
  141. if K == 1:
  142. folds = [(np.arange(n), np.arange(n))]
  143. elif K == n:
  144. folds = [
  145. (np.delete(np.arange(n), j), np.array([j])) for j in range(n)
  146. ]
  147. else:
  148. folds = KFold(
  149. n_splits=K, random_state=random_state, shuffle=True,
  150. ).split(X_A)
  151. pi_hat = np.zeros(A.shape, dtype=float)
  152. for train_index, test_index in folds:
  153. A_train = A[train_index]
  154. XA_train, XA_test = X_A[train_index], X_A[test_index]
  155. i_ctrl = np.sum(A_train, axis=1) == 0
  156. for j in range(A.shape[1]):
  157. i_case = A_train[:, j] == 1
  158. eligible = mask[train_index, j] if mask is not None else (i_ctrl | i_case)
  159. y_train = A_train[eligible, j]
  160. if y_train.size == 0 or np.unique(y_train).size != 2:
  161. raise ValueError(
  162. f'Treatment {j} needs at least one eligible control and case '
  163. 'in every training fold'
  164. )
  165. x_train = XA_train[eligible]
  166. constant_design = np.all(np.ptp(x_train, axis=0) == 0)
  167. if ps_model == 'logistic' and constant_design:
  168. if class_weight == 'balanced':
  169. pi_hat[test_index, j] = 0.5
  170. elif isinstance(class_weight, dict):
  171. sample_weight = np.asarray([
  172. class_weight.get(int(value), 1.0) for value in y_train
  173. ])
  174. pi_hat[test_index, j] = np.average(
  175. y_train, weights=sample_weight)
  176. else:
  177. pi_hat[test_index, j] = np.mean(y_train)
  178. else:
  179. pi_hat[test_index, j] = func_ps(x_train, y_train, XA_test)
  180. if clip is not None:
  181. lower, upper = _validate_clip(clip)
  182. pi_hat = np.clip(pi_hat, lower, upper)
  183. return pi_hat
  184. def _arm_support_metrics(A, pi_hat, j, bins=40, mask=None):
  185. """Support metrics for one treatment column, matching
  186. :func:`~causarray.diagnostics.summarize_propensity_scores`.
  187. Scoring a single arm avoids summarizing every treatment on each candidate
  188. penalty, which dominates the cost of a search. With ``mask``, only the
  189. cells the propensity model was fitted on are scored.
  190. """
  191. from sklearn.metrics import roc_auc_score
  192. from causarray.diagnostics import _effective_sample_size
  193. A = np.asarray(A, dtype=float)
  194. ctrl = A.sum(axis=1) == 0
  195. case = A[:, j] == 1
  196. eligible = ctrl | case
  197. if mask is not None:
  198. eligible &= mask[:, j]
  199. y = case[eligible].astype(int)
  200. p = np.asarray(pi_hat)[eligible, j]
  201. p_ctrl, p_case = p[y == 0], p[y == 1]
  202. h_ctrl, edges = np.histogram(p_ctrl, bins=bins, range=(0, 1))
  203. h_case, _ = np.histogram(p_case, bins=edges)
  204. h_ctrl = h_ctrl / h_ctrl.sum() if h_ctrl.sum() else h_ctrl
  205. h_case = h_case / h_case.sum() if h_case.sum() else h_case
  206. eps = np.finfo(float).eps
  207. return {
  208. 'auc': float(roc_auc_score(y, p)) if 0 < y.sum() < len(y) else np.nan,
  209. 'overlap_ratio': float(np.minimum(h_ctrl, h_case).sum()),
  210. 'ess_treated_fraction': (
  211. _effective_sample_size(1 / np.clip(p_case, eps, None)) / max(len(p_case), 1)),
  212. 'ess_control_fraction': (
  213. _effective_sample_size(1 / np.clip(1 - p_ctrl, eps, None)) / max(len(p_ctrl), 1)),
  214. }
  215. def _meets(metrics, target):
  216. """True when every ``<metric>_lt`` / ``<metric>_gt`` condition holds.
  217. The prefix must name a metric exactly, for example
  218. ``ess_treated_fraction_lt``, not ``ess_treated_lt``.
  219. """
  220. for key, bound in target.items():
  221. if key.endswith('_lt'):
  222. name, op = key[:-3], 'lt'
  223. elif key.endswith('_gt'):
  224. name, op = key[:-3], 'gt'
  225. else:
  226. raise ValueError(f"condition {key!r} must end with '_lt' or '_gt'")
  227. if name not in metrics:
  228. raise ValueError(
  229. f"condition {key!r} names no metric; available metrics are "
  230. f"{sorted(metrics)}")
  231. value = metrics[name]
  232. if not np.isfinite(value):
  233. return False
  234. if op == 'lt' and not value < bound:
  235. return False
  236. if op == 'gt' and not value > bound:
  237. return False
  238. return True
  239. def tune_penalty_factor(
  240. A, X_A, covariate, treatment_names=None, covariate_names=None,
  241. trigger=None, target=None, bracket=(1.0, 1e4), tol=0.15,
  242. on_infeasible='best',
  243. K=1, ps_model='logistic', mask=None, random_state=0, verbose=False,
  244. class_weight=None, **kwargs,
  245. ):
  246. """Choose a per-treatment L2 penalty for one propensity covariate.
  247. Some perturbations shift a covariate so strongly that the propensity model
  248. separates them from the controls, and their inverse-probability weights
  249. collapse onto a few cells. :func:`refit_propensity_scores` can penalize
  250. that coefficient for selected treatments; this picks the factor.
  251. For each triggered treatment the search is bracketed by two endpoints: the
  252. unpenalized fit and the fit with the covariate dropped, which is the
  253. infinite-penalty limit. **If the dropped fit misses ``target``, no finite
  254. penalty can reach it**, so the treatment is reported as infeasible after a
  255. single fit instead of an exhausted search. Otherwise the factor is found by
  256. bisection on a log scale and the *smallest* qualifying value is returned, so
  257. the covariate keeps as much of its adjustment role as the data support.
  258. Target a monotone metric. ``auc`` and ``overlap_ratio`` move monotonically
  259. with the penalty; ``ess_treated_fraction`` does not, because an arm that is
  260. completely separated has near-uniform weights and a deceptively high ESS
  261. that *falls* as the penalty restores genuine overlap.
  262. Penalizing a covariate is a soft version of dropping it, so the two are
  263. endpoints of one continuum. When the covariate is affected by treatment,
  264. no factor makes that contrast identified; the choice trades a known bias
  265. against precision and belongs in the analysis plan, not in this search.
  266. Parameters
  267. ----------
  268. A : array-like, shape (n,) or (n, a)
  269. Binary treatment indicators; all-zero rows are the shared controls.
  270. X_A : array-like, shape (n, d_A)
  271. Propensity covariates, including the intercept column.
  272. covariate : str or int
  273. The single covariate whose penalty is tuned.
  274. treatment_names, covariate_names : sequence, optional
  275. Labels for ``A`` columns and ``X_A`` columns.
  276. trigger : mapping or None
  277. Conditions selecting which treatments to tune, as ``{'auc_gt': 0.9,
  278. 'ess_treated_fraction_lt': 0.5}``. A treatment is triggered when **any**
  279. condition holds. Defaults to ``{'auc_gt': 0.9}``.
  280. target : mapping or None
  281. Conditions a factor must satisfy, in the same form, combined with
  282. **and**. Defaults to ``{'auc_lt': 0.9}``.
  283. bracket : tuple(float, float)
  284. Smallest and largest factors considered. The lower end is evaluated as
  285. the unpenalized fit and the upper end as the dropped-covariate fit.
  286. on_infeasible : {'best', 'none'}
  287. What to do when the dropped-covariate endpoint already misses
  288. ``target``, so no finite factor can reach it. ``'best'`` (default)
  289. applies the largest factor in ``bracket``, giving the arm the closest
  290. support the covariate allows. ``'none'`` leaves it unpenalized, which
  291. keeps the arm at its worst-case weights -- the trigger has already said
  292. its support is inadequate, so doing nothing is not a neutral choice.
  293. tol : float
  294. Bisection stops when the bracket spans less than ``tol`` in natural log
  295. units.
  296. K, ps_model, mask, random_state, verbose, class_weight, **kwargs
  297. Passed to :func:`estimate_propensity_scores` for every candidate fit.
  298. Returns
  299. -------
  300. penalty_factors_by_treatment : dict
  301. ``{treatment: {covariate: factor}}`` for feasible triggered treatments,
  302. ready to hand to :func:`refit_propensity_scores`. Treatments that were
  303. not triggered, already met ``target`` unpenalized, or are infeasible
  304. with ``on_infeasible='none'`` are absent.
  305. report : DataFrame
  306. One row per triggered treatment with the chosen factor, whether the
  307. target was feasible (``feasible``) and met at that factor
  308. (``target_met``), the number of fits used, and the metrics
  309. unpenalized, at the chosen factor, and with the covariate dropped.
  310. Examples
  311. --------
  312. >>> factors, report = tune_penalty_factor(
  313. ... A, X_A, 'log_library_size', treatment_names=names,
  314. ... covariate_names=cov, target={'auc_lt': 0.9}) # doctest: +SKIP
  315. >>> pi, _ = refit_propensity_scores(
  316. ... A, X_A, pi_hat=pi, treatment_names=names, covariate_names=cov,
  317. ... penalty_factors_by_treatment=factors) # doctest: +SKIP
  318. """
  319. trigger = {'auc_gt': 0.9} if trigger is None else dict(trigger)
  320. target = {'auc_lt': 0.9} if target is None else dict(target)
  321. low, high = (float(b) for b in bracket)
  322. if not 1.0 <= low < high:
  323. raise ValueError('bracket must satisfy 1 <= low < high')
  324. if tol <= 0:
  325. raise ValueError('tol must be positive')
  326. A_arr = np.asarray(A, dtype=float)
  327. if A_arr.ndim == 1:
  328. A_arr = A_arr[:, None]
  329. if treatment_names is None:
  330. treatment_names = (list(A.columns) if hasattr(A, 'columns')
  331. else list(range(A_arr.shape[1])))
  332. treatment_names = list(treatment_names)
  333. if covariate_names is None:
  334. covariate_names = (list(X_A.columns) if hasattr(X_A, 'columns')
  335. else [f'covariate_{j + 1}' for j in range(np.shape(X_A)[1])])
  336. covariate_names = list(covariate_names)
  337. if covariate not in covariate_names:
  338. if isinstance(covariate, (int, np.integer)) and 0 <= covariate < len(covariate_names):
  339. covariate = covariate_names[covariate]
  340. else:
  341. raise ValueError(f'covariate {covariate!r} is not in covariate_names')
  342. mask_arr = None
  343. if mask is not None:
  344. mask_arr = np.asarray(mask, dtype=bool)
  345. if mask_arr.ndim == 1:
  346. mask_arr = mask_arr[:, None]
  347. fit_kwargs = dict(K=K, ps_model=ps_model, mask=mask, clip=None,
  348. random_state=random_state, verbose=False,
  349. class_weight=class_weight, **kwargs)
  350. pi_base = estimate_propensity_scores(A_arr, X_A, **fit_kwargs)
  351. def scored(pi, j):
  352. return _arm_support_metrics(A_arr, pi, j, mask=mask_arr)
  353. rows, factors = [], {}
  354. for j, name in enumerate(treatment_names):
  355. base = scored(pi_base, j)
  356. if not any(_meets(base, {key: bound}) for key, bound in trigger.items()):
  357. continue
  358. n_fits = 1
  359. def evaluate(factor, name=name, j=j):
  360. pi_try, _ = refit_propensity_scores(
  361. A_arr, X_A, pi_hat=pi_base.copy(), treatment_names=treatment_names,
  362. covariate_names=covariate_names,
  363. penalty_factors_by_treatment={name: {covariate: float(factor)}},
  364. **fit_kwargs)
  365. return scored(pi_try, j)
  366. # factor -> infinity is the covariate dropped; it bounds what any
  367. # finite penalty can achieve.
  368. pi_drop, _ = refit_propensity_scores(
  369. A_arr, X_A, pi_hat=pi_base.copy(), treatment_names=treatment_names,
  370. covariate_names=covariate_names, drop_by_treatment={name: [covariate]},
  371. **fit_kwargs)
  372. dropped = scored(pi_drop, j)
  373. n_fits += 1
  374. feasible = _meets(dropped, target)
  375. if not feasible:
  376. if on_infeasible == 'best':
  377. chosen_factor, chosen = high, evaluate(high)
  378. n_fits += 1
  379. factors[name] = {covariate: high}
  380. elif on_infeasible == 'none':
  381. chosen_factor, chosen = 1.0, base
  382. else:
  383. raise ValueError("on_infeasible must be 'best' or 'none'")
  384. elif _meets(base, target):
  385. chosen_factor, chosen = low, base # nothing to do beyond the trigger
  386. else:
  387. lo_log, hi_log = np.log(low), np.log(high)
  388. chosen_factor, chosen = None, None
  389. while hi_log - lo_log > tol:
  390. mid_log = 0.5 * (lo_log + hi_log)
  391. metrics = evaluate(np.exp(mid_log))
  392. n_fits += 1
  393. if _meets(metrics, target):
  394. hi_log, chosen_factor, chosen = mid_log, float(np.exp(mid_log)), metrics
  395. else:
  396. lo_log = mid_log
  397. if chosen is None:
  398. # No bisection point qualified; score the largest factor
  399. # itself rather than borrowing the dropped fit's metrics.
  400. chosen_factor, chosen = high, evaluate(high)
  401. n_fits += 1
  402. factors[name] = {covariate: chosen_factor}
  403. rows.append({'treatment': name, 'penalty_factor': chosen_factor,
  404. 'feasible': feasible, 'target_met': _meets(chosen, target),
  405. 'n_fits': n_fits,
  406. **{f'{k}_unpenalized': v for k, v in base.items()},
  407. **{f'{k}_chosen': v for k, v in chosen.items()},
  408. **{f'{k}_dropped': v for k, v in dropped.items()}})
  409. report = pd.DataFrame(rows)
  410. if verbose and len(report):
  411. print(f'[tune_penalty_factor] {len(report)} treatments triggered, '
  412. f'{int(report.feasible.sum())} feasible, '
  413. f'{report.n_fits.sum()} fits', flush=True)
  414. return factors, report
  415. def refit_propensity_scores(
  416. A, X_A, drop_by_treatment=None, pi_hat=None, treatment_names=None,
  417. covariate_names=None, penalty_factors_by_treatment=None, K=1,
  418. ps_model='logistic', mask=None, clip=None, random_state=0, verbose=False,
  419. class_weight=None, **kwargs,
  420. ):
  421. """Refit propensity scores with treatment-specific covariate filtering.
  422. This helper supports sensitivity analyses in which different treatments
  423. omit different observed covariates or latent factors, or apply stronger L2
  424. regularization to selected covariates. When existing scores are supplied,
  425. only treatments named in either treatment-specific mapping are refitted;
  426. the remaining columns are carried over from ``pi_hat`` unchanged, up to the
  427. shared ``clip`` described below. Outcome models are not fit by this function
  428. and cached ``Y_hat`` values can be reused in :func:`LFC`.
  429. Parameters
  430. ----------
  431. A : array-like, shape (n, a)
  432. Binary treatment indicator matrix.
  433. X_A : array-like, shape (n, d_A)
  434. Full propensity-model design before treatment-specific filtering.
  435. drop_by_treatment : mapping or None
  436. Treatment names or indices mapped to covariate names or indices to
  437. remove for that treatment. Defaults to no removals.
  438. pi_hat : array-like or None, shape (n, a)
  439. Existing raw propensity scores. If supplied, treatments absent from
  440. ``drop_by_treatment`` are not refitted.
  441. treatment_names, covariate_names : sequence, optional
  442. Column labels, inferred from DataFrames when possible.
  443. penalty_factors_by_treatment : mapping or None
  444. Treatment names or indices mapped to ``{covariate: factor}`` mappings.
  445. A factor greater than one applies that multiple of the ordinary L2
  446. penalty to the named coefficient. This is implemented by dividing the
  447. covariate by ``sqrt(factor)`` during both fitting and prediction and is
  448. available only for ``ps_model='logistic'`` with an L2 penalty. A factor
  449. of one leaves the covariate unchanged. Interpret relative penalties on
  450. a common scale; ``prep_causarray_data`` standardizes log-library size.
  451. K, ps_model, mask, random_state, verbose, class_weight, **kwargs
  452. Passed to :func:`estimate_propensity_scores` for each refitted model.
  453. clip : tuple(float, float) or None, optional
  454. Bounds applied to the **whole** returned matrix, refitted columns and
  455. carried-over columns alike, so that a single consistent bound reaches
  456. :func:`LFC`. Pass ``None`` to leave carried-over scores exactly as
  457. supplied.
  458. Returns
  459. -------
  460. pi_updated : ndarray, shape (n, a)
  461. Updated propensity scores.
  462. report : DataFrame
  463. Audit table of refitted treatments, retained/dropped covariates, and
  464. the resulting score spread. ``degenerate_design`` flags a treatment
  465. whose retained design is constant, and ``score_std`` reports the
  466. standard deviation of its refitted scores on eligible rows; both make a
  467. collapsed propensity model visible without reading the warning stream.
  468. Warns
  469. -----
  470. RuntimeWarning
  471. If filtering leaves a constant design for some treatment, in which case
  472. its scores carry no covariate information.
  473. Notes
  474. -----
  475. A very large penalty factor shrinks a coefficient towards zero without
  476. making the design constant. That case leaves ``degenerate_design`` False
  477. but drives ``score_std`` towards zero.
  478. .. versionadded:: 0.0.9
  479. .. versionchanged:: 0.1.0
  480. Default ``class_weight`` changed from ``'balanced'`` to ``None`` to
  481. match :func:`estimate_propensity_scores` and :func:`LFC`.
  482. """
  483. if drop_by_treatment is None:
  484. drop_by_treatment = {}
  485. if penalty_factors_by_treatment is None:
  486. penalty_factors_by_treatment = {}
  487. if not isinstance(drop_by_treatment, Mapping):
  488. raise ValueError('drop_by_treatment must be a mapping')
  489. if not isinstance(penalty_factors_by_treatment, Mapping):
  490. raise ValueError('penalty_factors_by_treatment must be a mapping')
  491. if penalty_factors_by_treatment:
  492. if ps_model != 'logistic':
  493. raise ValueError(
  494. 'penalty_factors_by_treatment is supported only for logistic models'
  495. )
  496. if kwargs.get('penalty', 'l2') != 'l2':
  497. raise ValueError(
  498. 'penalty_factors_by_treatment requires logistic penalty="l2"'
  499. )
  500. if isinstance(A, pd.DataFrame):
  501. if treatment_names is None:
  502. treatment_names = list(A.columns)
  503. A_array = A.to_numpy()
  504. else:
  505. A_array = np.asarray(A)
  506. if A_array.ndim == 1:
  507. A_array = A_array[:, None]
  508. if A_array.ndim != 2 or not np.all(np.isin(A_array, (0, 1))):
  509. raise ValueError('A must be a one- or two-dimensional binary matrix')
  510. if treatment_names is None:
  511. treatment_names = list(range(A_array.shape[1]))
  512. treatment_names = list(treatment_names)
  513. if len(treatment_names) != A_array.shape[1]:
  514. raise ValueError('treatment_names must match the number of treatments')
  515. if len(set(treatment_names)) != len(treatment_names):
  516. raise ValueError('treatment_names must be unique')
  517. if isinstance(X_A, pd.DataFrame):
  518. if covariate_names is None:
  519. covariate_names = list(X_A.columns)
  520. X_array = X_A.to_numpy()
  521. else:
  522. X_array = np.asarray(X_A)
  523. if X_array.ndim != 2 or X_array.shape[0] != A_array.shape[0]:
  524. raise ValueError('X_A must be two-dimensional with the same rows as A')
  525. try:
  526. X_array = np.asarray(X_array, dtype=float)
  527. except (TypeError, ValueError) as exc:
  528. raise ValueError('X_A must contain numeric covariates') from exc
  529. if not np.all(np.isfinite(X_array)):
  530. raise ValueError('X_A must contain only finite values')
  531. if covariate_names is None:
  532. covariate_names = [f'covariate_{j + 1}' for j in range(X_array.shape[1])]
  533. covariate_names = list(covariate_names)
  534. if len(covariate_names) != X_array.shape[1]:
  535. raise ValueError('covariate_names must match the number of covariates')
  536. if len(set(covariate_names)) != len(covariate_names):
  537. raise ValueError('covariate_names must be unique')
  538. treatment_lookup = {name: j for j, name in enumerate(treatment_names)}
  539. covariate_lookup = {name: j for j, name in enumerate(covariate_names)}
  540. def resolve_treatment(value):
  541. if value in treatment_lookup:
  542. return treatment_lookup[value]
  543. if isinstance(value, (int, np.integer)) and 0 <= int(value) < A_array.shape[1]:
  544. return int(value)
  545. raise ValueError(f'Unknown treatment: {value}')
  546. def resolve_covariate(value):
  547. if value in covariate_lookup:
  548. return covariate_lookup[value]
  549. if isinstance(value, (int, np.integer)) and 0 <= int(value) < X_array.shape[1]:
  550. return int(value)
  551. raise ValueError(f'Unknown covariate: {value}')
  552. drops = {}
  553. for treatment, covariates in drop_by_treatment.items():
  554. j = resolve_treatment(treatment)
  555. if j in drops:
  556. raise ValueError(f'Treatment {treatment_names[j]} is specified more than once')
  557. if isinstance(covariates, (str, bytes)):
  558. covariates = [covariates]
  559. indices = [resolve_covariate(value) for value in covariates]
  560. if len(set(indices)) != len(indices):
  561. raise ValueError(
  562. f'Duplicate dropped covariates for treatment {treatment_names[j]}'
  563. )
  564. drops[j] = set(indices)
  565. penalty_factors = {}
  566. for treatment, factors in penalty_factors_by_treatment.items():
  567. j = resolve_treatment(treatment)
  568. if j in penalty_factors:
  569. raise ValueError(f'Treatment {treatment_names[j]} is specified more than once')
  570. if not isinstance(factors, Mapping):
  571. raise ValueError(
  572. f'Penalty factors for treatment {treatment_names[j]} must be a mapping'
  573. )
  574. resolved = {}
  575. for covariate, factor in factors.items():
  576. k = resolve_covariate(covariate)
  577. try:
  578. factor = float(factor)
  579. except (TypeError, ValueError) as exc:
  580. raise ValueError('Penalty factors must be finite numbers at least 1') from exc
  581. if not np.isfinite(factor) or factor < 1:
  582. raise ValueError('Penalty factors must be finite numbers at least 1')
  583. if k in resolved:
  584. raise ValueError(
  585. f'Duplicate penalty factors for covariate {covariate_names[k]}'
  586. )
  587. resolved[k] = factor
  588. penalty_factors[j] = resolved
  589. for j in set(drops).intersection(penalty_factors):
  590. conflict = drops[j].intersection(penalty_factors[j])
  591. if conflict:
  592. names = [covariate_names[k] for k in sorted(conflict)]
  593. raise ValueError(
  594. f'Covariates cannot be both dropped and penalized for '
  595. f'{treatment_names[j]}: {names}'
  596. )
  597. if pi_hat is None:
  598. pi_updated = np.empty(A_array.shape, dtype=float)
  599. refit_indices = range(A_array.shape[1])
  600. else:
  601. pi_updated = np.asarray(pi_hat, dtype=float).copy()
  602. if pi_updated.shape != A_array.shape:
  603. raise ValueError('pi_hat must have the same shape as A')
  604. if (not np.all(np.isfinite(pi_updated))
  605. or np.any((pi_updated < 0) | (pi_updated > 1))):
  606. raise ValueError('pi_hat must contain finite probabilities in [0, 1]')
  607. refit_indices = sorted(set(drops).union(penalty_factors))
  608. mask_array = None
  609. if mask is not None:
  610. mask_array = np.asarray(mask, dtype=bool)
  611. if mask_array.ndim == 1:
  612. mask_array = mask_array[:, None]
  613. if mask_array.shape != A_array.shape:
  614. raise ValueError('Mask must have the same shape as the treatment matrix')
  615. ctrl = np.sum(A_array, axis=1) == 0
  616. rows = []
  617. for j in refit_indices:
  618. dropped = drops.get(j, set())
  619. retained = [k for k in range(X_array.shape[1]) if k not in dropped]
  620. if not retained:
  621. raise ValueError(
  622. f'Filtering removes every covariate for treatment {treatment_names[j]}'
  623. )
  624. eligible = ctrl | (A_array[:, j] == 1)
  625. if mask_array is not None:
  626. eligible &= mask_array[:, j]
  627. X_treatment = X_array[:, retained].copy()
  628. treatment_penalties = penalty_factors.get(j, {})
  629. for position, k in enumerate(retained):
  630. factor = treatment_penalties.get(k, 1.0)
  631. X_treatment[:, position] /= np.sqrt(factor)
  632. degenerate = bool(np.all(np.ptp(X_treatment[eligible], axis=0) == 0))
  633. if degenerate:
  634. warnings.warn(
  635. f'The propensity design for treatment {treatment_names[j]} is '
  636. 'constant after filtering; its scores fall back to a '
  637. 'class-weighted prevalence and carry no covariate information.',
  638. RuntimeWarning, stacklevel=2,
  639. )
  640. scores = estimate_propensity_scores(
  641. A_array[:, [j]], X_treatment, K=K,
  642. ps_model=ps_model, mask=eligible[:, None], clip=None,
  643. random_state=random_state, verbose=verbose,
  644. class_weight=class_weight, **kwargs,
  645. )
  646. pi_updated[:, j] = scores[:, 0]
  647. rows.append({
  648. 'treatment': treatment_names[j],
  649. 'dropped_covariates': [covariate_names[k] for k in sorted(dropped)],
  650. 'retained_covariates': [covariate_names[k] for k in retained],
  651. 'penalty_factors': {
  652. covariate_names[k]: treatment_penalties[k]
  653. for k in retained if treatment_penalties.get(k, 1.0) != 1.0
  654. },
  655. 'n_retained': len(retained),
  656. 'degenerate_design': degenerate,
  657. 'score_std': float(np.std(scores[eligible, 0])),
  658. })
  659. if clip is not None:
  660. lower, upper = _validate_clip(clip)
  661. pi_updated = np.clip(pi_updated, lower, upper)
  662. report = pd.DataFrame(rows, columns=[
  663. 'treatment', 'dropped_covariates', 'retained_covariates',
  664. 'penalty_factors', 'n_retained', 'degenerate_design', 'score_std',
  665. ])
  666. return pi_updated, report
  667. def cross_fitting(
  668. Y, A, X, X_A, family='poisson', K=1, glm_alpha=1e-4,
  669. ps_model='logistic', ps_class_weight=None,
  670. Y_hat=None, pi_hat=None, mask=None, ps_clip='auto',
  671. return_raw_pi=False, verbose=False, **kwargs):
  672. '''
  673. Cross-fitting for causal estimands.
  674. Parameters
  675. ----------
  676. Y : array
  677. Outcomes.
  678. A : array
  679. Binary treatment indicator.
  680. X : array
  681. Covariates.
  682. X_A : array
  683. Covariates for the propensity score model.
  684. family : str, optional
  685. The family of the generalized linear model. The default is 'poisson'.
  686. K : int, optional
  687. The number of folds for cross-validation. The default is 1.
  688. glm_alpha : float, optional
  689. The regularization parameter for the generalized linear model. The default is 1e-4.
  690. ps_model : str, optional
  691. The propensity score model. The default is 'logistic'.
  692. ps_class_weight : str, dict or None, optional
  693. Class weighting used by the propensity model. ``None`` (default since
  694. 0.1.0) gives calibrated treatment probabilities; ``'balanced'``
  695. reproduces the pre-0.1.0 nuisance fit.
  696. Y_hat : array, optional
  697. Estimated potential outcome of shape (n, p, a, 2). The default is None.
  698. pi_hat : array, optional
  699. Propensity score of shape (n, a). The default is None.
  700. mask : array, optional
  701. Boolean mask of shape (n, a) for the treatment, indicating which samples are used for
  702. propensity-model fitting and the downstream estimand.
  703. ps_clip : {'auto'}, tuple(float, float), (lower_array, upper_array) or None, optional
  704. Bounds applied to scores used by AIPW. ``'auto'`` (default) resolves
  705. to a prevalence-aware bound per treatment (see
  706. :func:`causarray.DR_learner._resolve_ps_clip`); a pair of scalars
  707. applies one bound to all treatments; a pair of length-``a`` arrays
  708. gives per-treatment bounds; ``None`` disables clipping.
  709. return_raw_pi : bool, optional
  710. Return raw scores as a third result when true.
  711. **kwargs : dict
  712. Additional arguments to pass to the model.
  713. Returns
  714. -------
  715. Y_hat : array
  716. Estimated potential outcome under control.
  717. pi_hat : array
  718. Estimated propensity score.
  719. pi_hat_raw : array
  720. Unclipped propensity score, returned only when ``return_raw_pi=True``.
  721. '''
  722. kwargs = dict(kwargs)
  723. if 'class_weight' in kwargs:
  724. legacy_class_weight = kwargs.pop('class_weight')
  725. warnings.warn(
  726. 'Passing class_weight through LFC/cross_fitting is deprecated; '
  727. 'use ps_class_weight instead.',
  728. FutureWarning, stacklevel=2,
  729. )
  730. ps_class_weight = legacy_class_weight
  731. params_glm = _filter_params(fit_glm, {**kwargs, 'verbose': verbose})
  732. if verbose:
  733. pprint.pprint(params_glm)
  734. if K > 1:
  735. n_samples = X.shape[0]
  736. if K >= n_samples:
  737. # Use Leave-One-Out Cross-Validation
  738. folds = [([i for i in range(n_samples) if i != j], [j]) for j in range(n_samples)]
  739. else:
  740. # Initialize KFold cross-validator
  741. kf = KFold(n_splits=int(K), random_state=0, shuffle=True)
  742. folds = kf.split(X)
  743. else:
  744. folds = [(np.arange(X.shape[0]), np.arange(X.shape[0]))]
  745. # Initialize lists to store results
  746. if pi_hat is None:
  747. if verbose:
  748. pprint.pprint('Fit propensity score models...')
  749. ps_kwargs = {k: v for k, v in kwargs.items() if k != 'random_state'}
  750. if ps_model in ('logistic', 'ensemble'):
  751. ps_kwargs['class_weight'] = ps_class_weight
  752. pi_hat_raw = estimate_propensity_scores(
  753. A, X_A, K=K, ps_model=ps_model, mask=mask,
  754. random_state=kwargs.get('random_state', 0), verbose=verbose,
  755. **ps_kwargs,
  756. )
  757. else:
  758. pi_hat_raw = np.asarray(pi_hat, dtype=float).reshape(A.shape)
  759. if ps_clip is None:
  760. pi_hat = pi_hat_raw.copy()
  761. else:
  762. if isinstance(ps_clip, str):
  763. from causarray.DR_learner import _resolve_ps_clip
  764. ps_clip = _resolve_ps_clip(ps_clip, A, mask)
  765. if len(ps_clip) != 2:
  766. raise ValueError(
  767. "ps_clip must be 'auto', None, or a pair 0 <= lower < upper <= 1")
  768. lower = np.broadcast_to(np.asarray(ps_clip[0], dtype=float), (A.shape[1],))
  769. upper = np.broadcast_to(np.asarray(ps_clip[1], dtype=float), (A.shape[1],))
  770. if not (np.all(0 <= lower) and np.all(lower < upper) and np.all(upper <= 1)):
  771. raise ValueError(
  772. "ps_clip must be 'auto', None, or a pair 0 <= lower < upper <= 1")
  773. pi_hat = np.clip(pi_hat_raw, lower[None, :], upper[None, :])
  774. fit_Y = True if Y_hat is None else False
  775. if fit_Y:
  776. _yhat_gb = Y.shape[0] * Y.shape[1] * A.shape[1] * 2 * 8 / 1e9
  777. _mem_limit_gb = kwargs.get('mem_limit_gb', None)
  778. if _mem_limit_gb is not None and _yhat_gb > _mem_limit_gb:
  779. warnings.warn(
  780. f"Y_hat allocation ({_yhat_gb:.1f} GB as float64) exceeds "
  781. f"mem_limit_gb={_mem_limit_gb} GB; using float32 to halve peak memory.",
  782. ResourceWarning, stacklevel=3,
  783. )
  784. Y_hat = np.zeros((Y.shape[0], Y.shape[1], A.shape[1], 2), dtype=np.float32)
  785. else:
  786. Y_hat = np.zeros((Y.shape[0], Y.shape[1], A.shape[1], 2), dtype=float)
  787. # Perform cross-fitting
  788. for train_index, test_index in folds:
  789. # Split data
  790. X_train, X_test = X[train_index], X[test_index]
  791. XA_train, XA_test = X_A[train_index], X_A[test_index]
  792. A_train, A_test = A[train_index], A[test_index]
  793. Y_train, Y_test = Y[train_index], Y[test_index]
  794. if fit_Y:
  795. if verbose: pprint.pprint('Fit outcome models...')
  796. # Subset offset to training fold (for fitting) and test fold (for
  797. # imputation) when it is a pre-computed array, so that
  798. # ``fit_glm_auto`` receives arrays with matching leading
  799. # dimensions in both stages.
  800. params_glm_fold = params_glm
  801. offset_test_arr = None
  802. if 'offset' in params_glm and isinstance(params_glm['offset'], np.ndarray):
  803. params_glm_fold = dict(params_glm)
  804. params_glm_fold['offset'] = params_glm['offset'][train_index]
  805. offset_test_arr = params_glm['offset'][test_index]
  806. # Fit GLM on training data and predict on test data
  807. res = _gcate_glm.fit_glm_auto(Y_train, X_train, A_train, family=family, alpha=glm_alpha,
  808. impute=X_test, offset_test=offset_test_arr, **params_glm_fold)
  809. Y_hat[test_index,:,:,0] = res[1][0]
  810. Y_hat[test_index,:,:,1] = res[1][1]
  811. Y_hat = np.clip(Y_hat, None, 1e5)
  812. if return_raw_pi:
  813. return Y_hat, pi_hat, pi_hat_raw
  814. return Y_hat, pi_hat
  815. def AIPW_mean(Y, A, mu, pi):
  816. '''
  817. Augmented inverse probability weighted estimator (AIPW)
  818. Parameters
  819. ----------
  820. Y : array
  821. Outcomes of shape (n, p).
  822. A : array
  823. Binary treatment indicator of shape (n, a, 2).
  824. mu : array
  825. Conditional outcome distribution estimate of shape (n, p, a, 2).
  826. pi : array
  827. Propensity score of shape (n, a, 2).
  828. Returns
  829. -------
  830. tau : array
  831. A point estimate of the expected potential outcome of shape (p, a, 2).
  832. pseudo_y : array
  833. Pseudo-outcome of shape (n, p, a, 2).
  834. '''
  835. with np.errstate(divide='ignore', invalid='ignore', over='ignore'):
  836. weight = A / pi
  837. weight = weight[:, None, ...]
  838. Y = Y[:, :, None, None]
  839. # Influence-function values are intentionally left unconstrained. Even
  840. # for a nonnegative outcome, individual AIPW pseudo-outcomes may be
  841. # negative; projecting them cell by cell changes their mean and biases the
  842. # estimator. Parameter-space constraints belong after aggregation.
  843. pseudo_y = weight * (Y - mu) + mu
  844. tau = np.mean(pseudo_y, axis=0, dtype=np.float64)
  845. return tau, pseudo_y
  846. def run_ecv(
  847. X, y, M=200, M_max=1000,
  848. # fixed parameters for bagging regressor
  849. kwargs_ensemble={},
  850. # fixed parameters for decision tree
  851. kwargs_regr={},
  852. # grid search parameters
  853. grid_regr={},
  854. grid_ensemble={}
  855. ):
  856. """
  857. Runs Ensemble Cross-Validation (ECV) to find the best hyperparameters.
  858. """
  859. kwargs_ensemble = {**{'verbose': 1, 'bootstrap': True}, **kwargs_ensemble}
  860. kwargs_regr = {**{'min_samples_split': 20, 'min_samples_leaf': 10, 'max_features': 'sqrt', 'ccp_alpha': 0.02, 'class_weight': 'balanced'}, **kwargs_regr}
  861. grid_regr = {**{'max_depth': [3, 5, 7]}, **grid_regr}
  862. grid_ensemble = {**{'random_state': 0, 'max_samples': [0.4, 0.6, 0.8, 1.]}, **grid_ensemble}
  863. # Validate integer parameters
  864. M = int(M)
  865. M_max = int(M_max)
  866. # Make sure y is 2D
  867. y = y.reshape(-1, 1) if y.ndim == 1 else y
  868. # Run ECV
  869. _, info_ecv = ECV(
  870. X, y, DecisionTreeClassifier, grid_regr, grid_ensemble,
  871. kwargs_regr, kwargs_ensemble,
  872. M=M, M0=M, M_max=M_max, return_df=True
  873. )
  874. # Replace the in-sample best parameter for 'n_estimators' with extrapolated best parameter
  875. info_ecv['best_params_ensemble']['n_estimators'] = info_ecv['best_n_estimators_extrapolate']
  876. return info_ecv
  877. def fit_rf(
  878. X, y, X_test=None, M=100, M_max=1000, ecv=True,
  879. # fixed parameters for bagging regressor
  880. kwargs_ensemble={},
  881. # fixed parameters for decision tree
  882. kwargs_regr={},
  883. # grid search parameters
  884. grid_regr={},
  885. grid_ensemble={}
  886. ):
  887. """
  888. Fits a Random Forest model using parameters found by ECV.
  889. """
  890. kwargs_ensemble = {**{'verbose': 1, 'bootstrap': True}, **kwargs_ensemble}
  891. kwargs_regr = {**{'min_samples_split': 20, 'min_samples_leaf': 10, 'max_features': 'sqrt', 'ccp_alpha': 0.02, 'class_weight': 'balanced'}, **kwargs_regr}
  892. grid_regr = {**{'max_depth': [3, 5, 7]}, **grid_regr}
  893. grid_ensemble = {**{'random_state': 0, 'max_samples': [0.4, 0.6, 0.8, 1.]}, **grid_ensemble}
  894. # Make sure y is 2D
  895. y_2d = y.reshape(-1, 1) if y.ndim == 1 else y
  896. if ecv:
  897. # Get best parameters from ECV
  898. info_ecv = run_ecv(
  899. X, y_2d, M=M, M_max=M_max,
  900. kwargs_ensemble=kwargs_ensemble,
  901. kwargs_regr=kwargs_regr,
  902. grid_regr=grid_regr,
  903. grid_ensemble=grid_ensemble
  904. )
  905. params_regr = info_ecv['best_params_regr']
  906. params_ensemble = info_ecv['best_params_ensemble']
  907. else:
  908. params_regr = kwargs_regr
  909. params_ensemble = kwargs_ensemble
  910. # Fit the ensemble with the best CV parameters
  911. regr = Ensemble(
  912. estimator=DecisionTreeClassifier(**params_regr), **params_ensemble).fit(X, y_2d)
  913. # Predict
  914. if X_test is None:
  915. X_test = X
  916. return regr.predict(X_test).reshape(-1, y_2d.shape[1])
  917. def fit_rf_ind(X, Y, *args, **kwargs):
  918. Y_hat = Parallel(n_jobs=-1)(delayed(fit_rf)(X, Y[:,j], *args, **kwargs)
  919. for j in tqdm(range(Y.shape[1])))
  920. Y_pred = np.concatenate(Y_hat, axis=-1)
  921. return Y_pred
  922. def fit_rf_ind_ps(X, Y, *args, **kwargs):
  923. i_ctrl = (np.sum(Y, axis=1) == 0.)
  924. if 'X_test' not in kwargs:
  925. kwargs['X_test'] = X
  926. def _fit(X, y, i_ctrl, *args, **kwargs):
  927. i_case = (y == 1.)
  928. i_cells = i_ctrl | i_case
  929. return fit_rf(X[i_cells], y[i_cells], *args, **kwargs)
  930. Y_hat = Parallel(n_jobs=-1)(delayed(_fit)(X, Y[:,j], i_ctrl, *args, **kwargs)
  931. for j in tqdm(range(Y.shape[1])))
  932. Y_pred = np.concatenate(Y_hat, axis=-1)
  933. return Y_pred
  934. def fit_rf_ind_outcome(W, Y, A, *args, **kwargs):
  935. d = W.shape[1]
  936. a = A.shape[1]
  937. X = np.c_[W, A]
  938. X_test = np.tile(np.c_[W, np.zeros_like(A)][:,None,:], (1,1+a,1))
  939. for j in range(a):
  940. X_test[:,1+j,d+j] = 1
  941. X_test = X_test.reshape(-1, X_test.shape[-1])
  942. Y_pred = fit_rf_ind(X, Y, X_test=X_test)
  943. Y_pred = Y_pred.reshape(X.shape[0],1+a,Y.shape[1])
  944. Yhat_1 = Y_pred[:,1:,:].transpose(0,2,1)
  945. Yhat_0 = np.tile(Y_pred[:,0,:][:,:,None], (1,1,a))
  946. return Yhat_0, Yhat_1

DR_estimation.py at commit 14d4828, under MIT · at the source

Overview

Authors: Jin-Hong Du1,2, Maya Shen3, Hansruedi Mathys4, Kathryn Roeder3,5
ORCID iDs: Jin-Hong Du
  1. Department of Statistics and Actuarial Science, The University of Hong Kong, Pok Fu Lam, Hong Kong SAR 00000, China
  2. Musketeers Foundation Institute of Data Science, The University of Hong Kong, Pok Fu Lam, Hong Kong SAR 00000, China
  3. Department of Statistics and Data Science, Carnegie Mellon University, 5000 Forbes Ave, Pittsburgh, PA 15213, United States
  4. Department of Neurobiology, University of Pittsburgh, 4200 Fifth Ave, Pittsburgh, PA 15261, United States
  5. Computational Biology Department, Carnegie Mellon University, 5000 Forbes Ave, Pittsburgh, PA 15213, United States
Institutions: University of Hong Kong (Hong Kong SAR China); Carnegie Mellon University (United States); University of Pittsburgh (United States)
Journal: Briefings in bioinformatics, volume 27, issue 2, article bbag175
Dates: received 12 November 2025; accepted 15 March 2026; published online 15 April 2026; in print March 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1093/bib/bbag175 · PMID 41985059 · PMCID PMC13082396 · OpenAlex W7154448802
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: human (organism), mouse (organism), Alzheimer's / dementia (population), autism (population)
Methods: Machine learning
Keywords: causal inference, confounder adjustment, counterfactual, differential expression analysis, semiparametric inference
MeSH: Alzheimer Disease*, Genomics*, Single-Cell Analysis*, Animals, Autistic Disorder, Brain, Humans, Mice (* major topic)
Topic: Single-cell and spatial transcriptomics (Molecular Biology, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Citations: cited by 2 papers (Europe PMC); 53 references in the paper

Abstract

Advances in single-cell sequencing and Clustered Regularly Interspaced Short Palindromic Repeats (CRISPR) technologies have enabled detailed case-control comparisons and experimental perturbations at single-cell resolution. However, uncovering causal relationships in observational genomic data remains challenging due to selection bias and inadequate adjustment for unmeasured confounders, particularly in heterogeneous datasets. To address these challenges, we introduce causarray, a robust causal inference framework for analyzing array-based genomic data at both pseudo-bulk and single-cell levels under unmeasured confounding. causarray integrates a generalized confounder adjustment method to account for unmeasured confounders and employs semiparametric inference with flexible machine learning techniques to ensure robust statistical estimation of treatment effects. Benchmarking results show that causarray robustly separates treatment effects from confounders while preserving biological signals across diverse settings. We also apply causarray to two single-cell genomic studies: (i) an in vivo Perturb-seq study of autism risk genes in developing mouse brains and (ii) a case-control study of Alzheimer’s disease (AD) using three human brain transcriptomic datasets. In these applications, causarray identifies clustered causal effects of multiple autism risk genes and consistent causally affected genes across AD datasets, uncovering biologically relevant pathways directly linked to neuronal development and synaptic functions that are critical for understanding disease pathology.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repository

Its files are read in the Code ↔ Paper reader above, with 17 matches between paragraphs and lines of code.

jaydu1/causarray

License: MIT
State: the link answers, verified on 30 September 2026
Evidence: files inventoried
Commit: 14d482803af83879330625e27ed64c49a9b0b9e0, 25 September 2026
Languages: Python (57), R (15), Jupyter (9), Shell (5)
Size: 152 files, 86 scripts
Software Heritage: not archived
Found in: “Data and code availability”
Holds: README, license file, environment (environment-r.yaml, environment.yaml, pyproject.toml, setup.cfg, docs/requirements.txt, paper/environment.yml), tests, continuous integration, documentation, 10 notebooks
Not found: CITATION.cff
Tools: NumPy (55 files), pandas (39 files), SciPy (36 files), Matplotlib (16 files), statsmodels (15 files), tidyverse (15 files), Scanpy (12 files), scikit-learn (12 files), seaborn (11 files), anndata (10 files), ggplot2 (10 files), h5py (9 files), Seurat (9 files), reticulate (7 files), DESeq2 (6 files), clusterProfiler (4 files), patchwork (4 files), Numba (3 files), caret (1 file), data.table (1 file), Harmony (1 file), igraph (1 file), Plotly (1 file), SingleCellExperiment (1 file)
Availability: 1 check, the latest on 30 September 2026: the link answers
  • 30 September 2026: the link answers
88 files

The paper's code and data availability statement is in the Data section.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 86 scripts, each with its path and the digest of its content;
  • 17 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

No dataset and no data link were found in the paper.

Data and code availability

All datasets used in this paper are previously published and freely available, except the metadata for donors from the ROSMAP cohort. The Perturb-seq dataset is available through the Broad single cell portal (https://singlecell.broadinstitute.org/single_cell/study/SCP1184) as txt files. The gene expression count matrices of ROSMAP-AD datasets [42] can be obtained from supplementary website (https://compbio.mit.edu/ad_aging_brain/##processed-count-matrices), which have been deidentified to protect confidentiality—the mapping to ROSMAP IDs and complete metadata can be found on Synapse (https://www.synapse.org/##!Synapse:syn52293417) as Seurat objects (rds files). The SEA-AD datasets of nuclei-by-gene matrices with counts and normalized expression values from the snRNA-seq assay [43] are available through the Open Data Registry (https://registry.opendata.aws/allen-sea-ad-atlas/) in an AWS bucket (sea-ad-single-cell-profiling) as AnnData objects (h5ad files). The code for reproducing the results in the paper and the causarray package can be accessed at https://github.com/jaydu1/causarray.

Reproduced under the paper's license (CC BY), from the paper cited above.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 30 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 4 authors, 5 keywords, 8 MeSH terms, 3 funders, 49 references.

Cite

This paper

Du, J.-H., Shen, M., Mathys, H., & Roeder, K. (2026). Uncovering causal relationships in single-cell omic studies with causarray. Briefings in bioinformatics, 27(2), bbag175. https://doi.org/10.1093/bib/bbag175

BibTeX

@article{du2026uncovering,
author = {Du, Jin-Hong and Shen, Maya and Mathys, Hansruedi and Roeder, Kathryn},
title = {{Uncovering causal relationships in single-cell omic studies with causarray}},
journal = {Briefings in bioinformatics},
year = {2026},
month = mar,
volume = {27},
number = {2},
pages = {bbag175},
publisher = {Oxford University Press},
issn = {1467-5463},
doi = {10.1093/bib/bbag175},
url = {https://doi.org/10.1093/bib/bbag175},
pmid = {41985059},
pmcid = {PMC13082396}
}

RIS

TY - JOUR
AU - Du, Jin-Hong
AU - Shen, Maya
AU - Mathys, Hansruedi
AU - Roeder, Kathryn
TI - Uncovering causal relationships in single-cell omic studies with causarray
T2 - Briefings in bioinformatics
J2 - Brief Bioinform
PY - 2026
DA - 2026/03/01
VL - 27
IS - 2
SP - bbag175
SN - 1467-5463
PB - Oxford University Press
DO - 10.1093/bib/bbag175
UR - https://doi.org/10.1093/bib/bbag175
LA - en
ER -

CSL-JSON

{
"id": "10.1093/bib/bbag175",
"type": "article-journal",
"title": "Uncovering causal relationships in single-cell omic studies with causarray",
"container-title": "Briefings in bioinformatics",
"author": [
{
"family": "Du",
"given": "Jin-Hong"
},
{
"family": "Shen",
"given": "Maya"
},
{
"family": "Mathys",
"given": "Hansruedi"
},
{
"family": "Roeder",
"given": "Kathryn"
}
],
"container-title-short": "Brief Bioinform",
"volume": "27",
"issue": "2",
"page": "bbag175",
"DOI": "10.1093/bib/bbag175",
"PMID": "41985059",
"PMCID": "PMC13082396",
"ISSN": "1467-5463",
"publisher": "Oxford University Press",
"URL": "https://doi.org/10.1093/bib/bbag175",
"language": "en",
"issued": {
"date-parts": [
[
2026,
3,
1
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.1038/s44318-026-00818-9 [code]
FAM134B-mediated ER-phagy degrades APP and suppresses Alzheimer's disease pathology.
Journal: The EMBO journal
In common: Harmony, SingleCellExperiment, reticulate, 18 other tools, Alzheimer's / dementia, mouse
[2] doi:10.1016/j.xcrm.2026.102766 [code]
A longitudinal single-cell and spatial multiomic atlas of pediatric high-grade glioma.
Journal: Cell reports. Medicine
In common: Harmony, SingleCellExperiment, anndata, 18 other tools, 1 reference
[3] doi:10.1038/s41586-026-10629-x [code]
Whole-genome duplication shaped cell-type evolution in the vertebrate brain.
Journal: Nature
In common: Harmony, reticulate, anndata, 16 other tools, mouse, 1 reference
[4] doi:10.1016/j.xcrm.2026.102651 [code]
Integrative CSF profiling identifies disease-specific immune responses in leptomeningeal disease.
Journal: Cell reports. Medicine
In common: Harmony, SingleCellExperiment, reticulate, 16 other tools
[5] doi:10.1038/s41586-026-10214-2 [code]
Multidimensional profiling of heterogeneity in supratentorial ependymomas.
Journal: Nature
In common: Harmony, SingleCellExperiment, reticulate, 16 other tools, mouse
[6] doi:10.1002/imt2.70163 [code]
Spatial multi-omics unveils sphingolipid metabolic reprogramming within the retinal pathological niche.
Journal: iMeta
In common: Harmony, SingleCellExperiment, anndata, 16 other tools, mouse
[7] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: caret, reticulate, igraph, 16 other tools, mouse
[8] doi:10.7554/elife.93640 [code]
Sibling chimerism among microglia in marmosets.
Journal: eLife
In common: Harmony, SingleCellExperiment, reticulate, 15 other tools
[9] doi:10.1016/j.isci.2026.116055 [code]
Mapping the transcriptional diversity of calcium signaling in the mouse and human brain.
Journal: iScience
In common: Harmony, SingleCellExperiment, anndata, 15 other tools, mouse
[10] doi:10.1038/s44320-026-00208-7 [code]
Interpretable deep generative ensemble learning for single-cell omics with Hydra.
Journal: Molecular systems biology
In common: SingleCellExperiment, reticulate, anndata, 15 other tools

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.