causarray 0.0.6__tar.gz → 0.0.7__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -1,6 +1,9 @@
1
1
  # Planning documents
2
2
  plan/
3
3
 
4
+ # Local Codex project guidance
5
+ /AGENTS.md
6
+
4
7
  # Byte-compiled / optimized / DLL files
5
8
  __pycache__/
6
9
  *.py[cod]
@@ -170,11 +173,15 @@ cython_debug/
170
173
  tests/benchmark_adamson.py
171
174
 
172
175
  causarray/___*.py
173
- /docs/source/tutorial/adamson
176
+
177
+ # Adamson tutorial is currently a local analysis and is not part of the
178
+ # committed documentation.
179
+ /docs/source/tutorial/adamson/
174
180
 
175
181
  # Replogle tutorial: large data files (regenerate with prep_tutorial_data.py or
176
182
  # download from the project data repository; replogle_subset.h5ad is ~2 GB)
177
183
  docs/source/tutorial/replogle/replogle_subset.h5ad
184
+ docs/source/tutorial/replogle/replogle_subset_norm*.h5ad
178
185
  docs/source/tutorial/replogle/replogle_normed.h5ad
179
186
  docs/source/tutorial/replogle/replogle_results.h5
180
187
 
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: causarray
3
- Version: 0.0.6
3
+ Version: 0.0.7
4
4
  Summary: causarray is a Python module for simultaneous causal inference with an array of outcomes.
5
5
  Author-email: Jin-Hong Du <jinhongd@hku.hk>, Maya Shen <myshen@andrew.cmu.edu>, Hansruedi Mathys <mathysh@pitt.edu>, Kathryn Roeder <jinhongd@hku.hk>
6
6
  Maintainer-email: Jin-Hong Du <jinhongd@hku.hk>
@@ -30,7 +30,7 @@ Description-Content-Type: text/markdown
30
30
 
31
31
  # causarray
32
32
 
33
- Advances in single-cell sequencing and 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` [Du25], a doubly robust causal inference framework for analyzing array-based genomic data at both bulk-cell and single-cell levels. `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.
33
+ Advances in single-cell sequencing and 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` [Du26], a doubly robust causal inference framework for analyzing array-based genomic data at both bulk-cell and single-cell levels. `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.
34
34
 
35
35
 
36
36
  ## Usage
@@ -59,10 +59,19 @@ The documentation and tutorials using both `Python` and `R` are available at [ca
59
59
 
60
60
 
61
61
 
62
- ## Batch fitting for large-scale screens
62
+ ## Tutorials
63
63
 
64
- For screens with hundreds to thousands of perturbations, use the batch API
65
- so that peak memory is bounded by one batch at a time:
64
+ | Tutorial | Language | Description | Link |
65
+ |----------|----------|-------------|------|
66
+ | Perturb-seq [Jin20] | Python | CRISPR screen analysis on excitatory neurons | [Notebook](https://causarray.readthedocs.io/en/latest/tutorial/perturbseq/perturbseq-py.html) |
67
+ | Perturb-seq [Jin20] | R | Same analysis using `reticulate` | [Notebook](https://causarray.readthedocs.io/en/latest/tutorial/perturbseq/perturbseq-r.html) |
68
+ | Genome-wide CRISPRi screen [Replogle22] | Python | Batch fitting on 200 perturbations from a K562 genome-wide CRISPRi screen | [Notebook](https://causarray.readthedocs.io/en/latest/tutorial/replogle/replogle-py.html) |
69
+ | Case-control: SEA-AD [Gabitto24] | Python | Causal inference on observational single-cell data (Alzheimer's disease) | [Notebook](https://causarray.readthedocs.io/en/latest/tutorial/case_control/sea_ad_case_control.html) |
70
+
71
+ ### Batch fitting API
72
+
73
+ For screens with hundreds to thousands of perturbations, use `gcate_lfc_batch` so
74
+ that peak memory is bounded by one batch at a time:
66
75
 
67
76
  ```python
68
77
  from causarray import gcate_lfc_batch
@@ -77,16 +86,12 @@ df_res = gcate_lfc_batch(
77
86
  )
78
87
  ```
79
88
 
80
- See the [Replogle-E-K562 tutorial](https://causarray.readthedocs.io/en/latest/)
89
+ See the [Replogle-E-K562 tutorial](https://causarray.readthedocs.io/en/latest/tutorial/replogle/replogle-py.html)
81
90
  for a demonstration on 200 perturbations from a genome-wide CRISPRi screen.
82
91
 
83
92
  ## Changelog
84
93
 
85
- - [x] (2025-01-30) Python package released on PyPI
86
- - [x] (2025-02-01) Code for reproducing figures in paper
87
- - [x] (2025-02-02) Tutorial for Python and R
88
- - [x] (2026-05-31) Batch fitting API (`gcate_lfc_batch`) for large-scale screens
89
- - [x] (2026-05-31) Documentation at [causarray.readthedocs.io](https://causarray.readthedocs.io/en/latest/)
94
+ See [CHANGELOG](https://causarray.readthedocs.io/en/latest/changelog.html) for a full version history.
90
95
 
91
96
 
92
97
  <!--
@@ -128,4 +133,10 @@ rmarkdown::render("perturbseq.Rmd", rmarkdown::md_document(variant = "markdown_g
128
133
 
129
134
 
130
135
  ## References
131
- [Du25] Jin-Hong Du, Maya Shen, Hansruedi Mathys, and Kathryn Roeder (2025). Causal differential expression analysis under unmeasured confounders with causarray. bioRxiv, 2025-01.
136
+ [Du26] Jin-Hong Du, Maya Shen, Hansruedi Mathys, and Kathryn Roeder. "Uncovering causal relationships in single cell omic studies with causarray". In: Briefings in Bioinformatics (2026).
137
+
138
+ [Gabitto24] Mariano I. Gabitto et al. "Integrated multimodal cell atlas of Alzheimer's disease". In: Nature Neuroscience (2024).
139
+
140
+ [Jin20] Xin Jin et al. "In vivo Perturb-seq reveals neuronal and glial abnormalities associated with autism risk genes". In: Nature Neuroscience (2020).
141
+
142
+ [Replogle22] Joseph M. Replogle et al. "Mapping information-rich genotype-phenotype landscapes with genome-scale Perturb-seq". In: Cell (2022).
@@ -5,7 +5,7 @@
5
5
 
6
6
  # causarray
7
7
 
8
- Advances in single-cell sequencing and 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` [Du25], a doubly robust causal inference framework for analyzing array-based genomic data at both bulk-cell and single-cell levels. `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.
8
+ Advances in single-cell sequencing and 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` [Du26], a doubly robust causal inference framework for analyzing array-based genomic data at both bulk-cell and single-cell levels. `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.
9
9
 
10
10
 
11
11
  ## Usage
@@ -34,10 +34,19 @@ The documentation and tutorials using both `Python` and `R` are available at [ca
34
34
 
35
35
 
36
36
 
37
- ## Batch fitting for large-scale screens
37
+ ## Tutorials
38
38
 
39
- For screens with hundreds to thousands of perturbations, use the batch API
40
- so that peak memory is bounded by one batch at a time:
39
+ | Tutorial | Language | Description | Link |
40
+ |----------|----------|-------------|------|
41
+ | Perturb-seq [Jin20] | Python | CRISPR screen analysis on excitatory neurons | [Notebook](https://causarray.readthedocs.io/en/latest/tutorial/perturbseq/perturbseq-py.html) |
42
+ | Perturb-seq [Jin20] | R | Same analysis using `reticulate` | [Notebook](https://causarray.readthedocs.io/en/latest/tutorial/perturbseq/perturbseq-r.html) |
43
+ | Genome-wide CRISPRi screen [Replogle22] | Python | Batch fitting on 200 perturbations from a K562 genome-wide CRISPRi screen | [Notebook](https://causarray.readthedocs.io/en/latest/tutorial/replogle/replogle-py.html) |
44
+ | Case-control: SEA-AD [Gabitto24] | Python | Causal inference on observational single-cell data (Alzheimer's disease) | [Notebook](https://causarray.readthedocs.io/en/latest/tutorial/case_control/sea_ad_case_control.html) |
45
+
46
+ ### Batch fitting API
47
+
48
+ For screens with hundreds to thousands of perturbations, use `gcate_lfc_batch` so
49
+ that peak memory is bounded by one batch at a time:
41
50
 
42
51
  ```python
43
52
  from causarray import gcate_lfc_batch
@@ -52,16 +61,12 @@ df_res = gcate_lfc_batch(
52
61
  )
53
62
  ```
54
63
 
55
- See the [Replogle-E-K562 tutorial](https://causarray.readthedocs.io/en/latest/)
64
+ See the [Replogle-E-K562 tutorial](https://causarray.readthedocs.io/en/latest/tutorial/replogle/replogle-py.html)
56
65
  for a demonstration on 200 perturbations from a genome-wide CRISPRi screen.
57
66
 
58
67
  ## Changelog
59
68
 
60
- - [x] (2025-01-30) Python package released on PyPI
61
- - [x] (2025-02-01) Code for reproducing figures in paper
62
- - [x] (2025-02-02) Tutorial for Python and R
63
- - [x] (2026-05-31) Batch fitting API (`gcate_lfc_batch`) for large-scale screens
64
- - [x] (2026-05-31) Documentation at [causarray.readthedocs.io](https://causarray.readthedocs.io/en/latest/)
69
+ See [CHANGELOG](https://causarray.readthedocs.io/en/latest/changelog.html) for a full version history.
65
70
 
66
71
 
67
72
  <!--
@@ -103,4 +108,10 @@ rmarkdown::render("perturbseq.Rmd", rmarkdown::md_document(variant = "markdown_g
103
108
 
104
109
 
105
110
  ## References
106
- [Du25] Jin-Hong Du, Maya Shen, Hansruedi Mathys, and Kathryn Roeder (2025). Causal differential expression analysis under unmeasured confounders with causarray. bioRxiv, 2025-01.
111
+ [Du26] Jin-Hong Du, Maya Shen, Hansruedi Mathys, and Kathryn Roeder. "Uncovering causal relationships in single cell omic studies with causarray". In: Briefings in Bioinformatics (2026).
112
+
113
+ [Gabitto24] Mariano I. Gabitto et al. "Integrated multimodal cell atlas of Alzheimer's disease". In: Nature Neuroscience (2024).
114
+
115
+ [Jin20] Xin Jin et al. "In vivo Perturb-seq reveals neuronal and glial abnormalities associated with autism risk genes". In: Nature Neuroscience (2020).
116
+
117
+ [Replogle22] Joseph M. Replogle et al. "Mapping information-rich genotype-phenotype landscapes with genome-scale Perturb-seq". In: Cell (2022).
@@ -9,6 +9,7 @@ from causarray.utils import _filter_params
9
9
  from joblib import Parallel, delayed
10
10
  from tqdm import tqdm
11
11
  import pprint
12
+ import warnings
12
13
 
13
14
  from sklearn.model_selection import KFold, ShuffleSplit
14
15
 
@@ -18,13 +19,18 @@ def _get_func_ps(ps_model, **kwargs):
18
19
  func_ps = lambda X, Y, X_test:fit_rf_ind_ps(X, Y[:,None], X_test=X_test, **params_ps)[:,0]
19
20
  elif ps_model=='logistic':
20
21
  clf_ps = LogisticRegression
21
- kwargs = {**{'fit_intercept':False, 'C':1e0, 'class_weight':'balanced', 'random_state':0}, **kwargs}
22
+ kwargs = {**{'fit_intercept':False, 'C':1e0, 'class_weight':None, 'random_state':0}, **kwargs}
22
23
  params_ps = _filter_params(clf_ps().get_params(), kwargs)
23
24
  func_ps = lambda X, Y, X_test: clf_ps(**params_ps).fit(X, Y).predict_proba(X_test)[:,1]
24
25
  elif ps_model=='ensemble':
26
+ kwargs = dict(kwargs)
27
+ logistic_class_weight = kwargs.pop('class_weight', None)
25
28
  params_ps_rf = _filter_params(fit_rf, kwargs)
26
29
  clf_ps = LogisticRegression
27
- kwargs = {**{'fit_intercept':False, 'C':1e0, 'class_weight':'balanced', 'random_state':0}, **kwargs}
30
+ kwargs = {**{
31
+ 'fit_intercept': False, 'C': 1e0,
32
+ 'class_weight': logistic_class_weight, 'random_state': 0,
33
+ }, **kwargs}
28
34
  params_ps_lr = _filter_params(clf_ps().get_params(), kwargs)
29
35
  params_ps = {'params_ps_rf':params_ps_rf, 'params_ps_lr':params_ps_lr}
30
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
@@ -34,10 +40,149 @@ def _get_func_ps(ps_model, **kwargs):
34
40
  return func_ps, params_ps
35
41
 
36
42
 
43
+ def estimate_propensity_scores(
44
+ A, X_A, K=1, ps_model='logistic', mask=None, clip=None,
45
+ random_state=0, verbose=False, class_weight='balanced', **kwargs,
46
+ ):
47
+ """Estimate per-treatment propensity scores.
48
+
49
+ Each treatment is compared with the shared all-zero control group. With
50
+ ``K > 1``, every returned score is predicted by a model that did not train
51
+ on that cell. Logistic models use ``class_weight='balanced'`` by default,
52
+ matching :func:`LFC` and historical causarray fits. Pass
53
+ ``class_weight=None`` for calibrated treatment probabilities.
54
+
55
+ Parameters
56
+ ----------
57
+ A : array-like, shape (n,) or (n, a)
58
+ Binary treatment indicators. Rows containing only zeros are controls.
59
+ X_A : array-like, shape (n, d_A)
60
+ Covariates used by the propensity model, including an intercept column
61
+ when ``fit_intercept=False``.
62
+ K : int, optional
63
+ Number of folds. ``1`` fits and predicts on all eligible cells;
64
+ values greater than one produce out-of-fold predictions.
65
+ ps_model : {'logistic', 'random_forest_cv', 'ensemble'}, optional
66
+ Propensity model.
67
+ mask : array-like or None, shape (n,) or (n, a)
68
+ Optional per-treatment eligibility mask for model fitting.
69
+ clip : tuple(float, float) or None, optional
70
+ Bounds applied after prediction. ``None`` returns raw probabilities.
71
+ random_state : int, optional
72
+ Random seed used for fold construction and supported estimators.
73
+ class_weight : str, dict or None, optional
74
+ Class weighting for logistic propensity estimation. The default
75
+ ``'balanced'`` matches :func:`LFC`; pass ``None`` for calibrated
76
+ probabilities.
77
+
78
+ Returns
79
+ -------
80
+ pi_hat : ndarray, shape (n, a)
81
+ Estimated probabilities ``P(A_j=1 | X_A)``.
82
+ """
83
+ A = np.asarray(A)
84
+ if A.ndim == 1:
85
+ A = A[:, None]
86
+ X_A = np.asarray(X_A)
87
+ if A.ndim != 2 or X_A.ndim != 2 or A.shape[0] != X_A.shape[0]:
88
+ raise ValueError('A and X_A must be two-dimensional with matching rows')
89
+ if not np.all(np.isin(A, (0, 1))):
90
+ raise ValueError('A must contain only binary treatment indicators')
91
+
92
+ try:
93
+ K_int = int(K)
94
+ except (TypeError, ValueError) as exc:
95
+ raise ValueError('K must be a positive integer') from exc
96
+ if K_int != K:
97
+ raise ValueError('K must be a positive integer')
98
+ K = K_int
99
+ if K < 1:
100
+ raise ValueError('K must be a positive integer')
101
+ if K > A.shape[0]:
102
+ raise ValueError('K cannot exceed the number of samples')
103
+
104
+ if mask is not None:
105
+ mask = np.asarray(mask, dtype=bool)
106
+ if mask.ndim == 1:
107
+ mask = mask[:, None]
108
+ if mask.shape != A.shape:
109
+ raise ValueError('Mask must have the same shape as the treatment matrix')
110
+
111
+ func_ps, params_ps = _get_func_ps(
112
+ ps_model, verbose=False, random_state=random_state,
113
+ class_weight=class_weight, **kwargs)
114
+ if verbose:
115
+ pprint.pprint(params_ps)
116
+
117
+ if ps_model == 'random_forest_cv':
118
+ info_ecv = run_ecv(X_A, A, **params_ps)
119
+ func_ps, params_ps = _get_func_ps(
120
+ ps_model, verbose=False, ecv=False,
121
+ kwargs_ensemble=info_ecv['best_params_ensemble'],
122
+ kwargs_regr=info_ecv['best_params_regr'],
123
+ )
124
+ if verbose:
125
+ pprint.pprint('Best parameters for the regression model:')
126
+ pprint.pprint(info_ecv['best_params_regr'])
127
+ pprint.pprint('Best parameters for the ensemble model:')
128
+ pprint.pprint(info_ecv['best_params_ensemble'])
129
+
130
+ n = A.shape[0]
131
+ if K == 1:
132
+ folds = [(np.arange(n), np.arange(n))]
133
+ elif K == n:
134
+ folds = [
135
+ (np.delete(np.arange(n), j), np.array([j])) for j in range(n)
136
+ ]
137
+ else:
138
+ folds = KFold(
139
+ n_splits=K, random_state=random_state, shuffle=True,
140
+ ).split(X_A)
141
+
142
+ pi_hat = np.zeros(A.shape, dtype=float)
143
+ for train_index, test_index in folds:
144
+ A_train = A[train_index]
145
+ XA_train, XA_test = X_A[train_index], X_A[test_index]
146
+ i_ctrl = np.sum(A_train, axis=1) == 0
147
+
148
+ for j in range(A.shape[1]):
149
+ i_case = A_train[:, j] == 1
150
+ eligible = mask[train_index, j] if mask is not None else (i_ctrl | i_case)
151
+ y_train = A_train[eligible, j]
152
+ if y_train.size == 0 or np.unique(y_train).size != 2:
153
+ raise ValueError(
154
+ f'Treatment {j} needs at least one eligible control and case '
155
+ 'in every training fold'
156
+ )
157
+
158
+ x_train = XA_train[eligible]
159
+ constant_design = np.all(np.ptp(x_train, axis=0) == 0)
160
+ if ps_model == 'logistic' and constant_design:
161
+ if class_weight == 'balanced':
162
+ pi_hat[test_index, j] = 0.5
163
+ elif isinstance(class_weight, dict):
164
+ sample_weight = np.asarray([
165
+ class_weight.get(int(value), 1.0) for value in y_train
166
+ ])
167
+ pi_hat[test_index, j] = np.average(
168
+ y_train, weights=sample_weight)
169
+ else:
170
+ pi_hat[test_index, j] = np.mean(y_train)
171
+ else:
172
+ pi_hat[test_index, j] = func_ps(x_train, y_train, XA_test)
173
+
174
+ if clip is not None:
175
+ if len(clip) != 2 or not 0 <= clip[0] < clip[1] <= 1:
176
+ raise ValueError('clip must be None or a pair 0 <= lower < upper <= 1')
177
+ pi_hat = np.clip(pi_hat, clip[0], clip[1])
178
+ return pi_hat
179
+
180
+
37
181
  def cross_fitting(
38
182
  Y, A, X, X_A, family='poisson', K=1, glm_alpha=1e-4,
39
- ps_model='logistic',
40
- Y_hat=None, pi_hat=None, mask=None, verbose=False, **kwargs):
183
+ ps_model='logistic', ps_class_weight='balanced',
184
+ Y_hat=None, pi_hat=None, mask=None, ps_clip=(0.01, 0.99),
185
+ return_raw_pi=False, verbose=False, **kwargs):
41
186
  '''
42
187
  Cross-fitting for causal estimands.
43
188
 
@@ -59,6 +204,10 @@ def cross_fitting(
59
204
  The regularization parameter for the generalized linear model. The default is 1e-4.
60
205
  ps_model : str, optional
61
206
  The propensity score model. The default is 'logistic'.
207
+ ps_class_weight : str, dict or None, optional
208
+ Class weighting used by the propensity model. ``'balanced'`` preserves
209
+ the established ``LFC`` nuisance fit; pass ``None`` for calibrated
210
+ treatment probabilities.
62
211
 
63
212
  Y_hat : array, optional
64
213
  Estimated potential outcome of shape (n, p, a, 2). The default is None.
@@ -66,8 +215,11 @@ def cross_fitting(
66
215
  Propensity score of shape (n, a). The default is None.
67
216
  mask : array, optional
68
217
  Boolean mask of shape (n, a) for the treatment, indicating which samples are used for
69
- the estimation of the estimand. This does not affect the estimation of pseudo-outcomes
70
- and propensity scores.
218
+ propensity-model fitting and the downstream estimand.
219
+ ps_clip : tuple(float, float) or None, optional
220
+ Bounds applied to scores used by AIPW. ``None`` disables clipping.
221
+ return_raw_pi : bool, optional
222
+ Return raw scores as a third result when true.
71
223
 
72
224
  **kwargs : dict
73
225
  Additional arguments to pass to the model.
@@ -78,12 +230,22 @@ def cross_fitting(
78
230
  Estimated potential outcome under control.
79
231
  pi_hat : array
80
232
  Estimated propensity score.
233
+ pi_hat_raw : array
234
+ Unclipped propensity score, returned only when ``return_raw_pi=True``.
81
235
  '''
82
- func_ps, params_ps = _get_func_ps(ps_model, verbose=False, **kwargs)
236
+ kwargs = dict(kwargs)
237
+ if 'class_weight' in kwargs:
238
+ legacy_class_weight = kwargs.pop('class_weight')
239
+ warnings.warn(
240
+ 'Passing class_weight through LFC/cross_fitting is deprecated; '
241
+ 'use ps_class_weight instead.',
242
+ FutureWarning, stacklevel=2,
243
+ )
244
+ ps_class_weight = legacy_class_weight
245
+
83
246
  params_glm = _filter_params(fit_glm, {**kwargs, 'verbose': verbose})
84
247
 
85
248
  if verbose:
86
- pprint.pprint(params_ps)
87
249
  pprint.pprint(params_glm)
88
250
 
89
251
  if K > 1:
@@ -99,14 +261,31 @@ def cross_fitting(
99
261
  folds = [(np.arange(X.shape[0]), np.arange(X.shape[0]))]
100
262
 
101
263
  # Initialize lists to store results
102
- fit_pi = True if pi_hat is None else False
103
- pi_hat = np.zeros_like(A, dtype=float) if fit_pi else pi_hat
264
+ if pi_hat is None:
265
+ if verbose:
266
+ pprint.pprint('Fit propensity score models...')
267
+ ps_kwargs = {k: v for k, v in kwargs.items() if k != 'random_state'}
268
+ if ps_model in ('logistic', 'ensemble'):
269
+ ps_kwargs['class_weight'] = ps_class_weight
270
+ pi_hat_raw = estimate_propensity_scores(
271
+ A, X_A, K=K, ps_model=ps_model, mask=mask,
272
+ random_state=kwargs.get('random_state', 0), verbose=verbose,
273
+ **ps_kwargs,
274
+ )
275
+ else:
276
+ pi_hat_raw = np.asarray(pi_hat, dtype=float).reshape(A.shape)
277
+ if ps_clip is None:
278
+ pi_hat = pi_hat_raw.copy()
279
+ else:
280
+ if len(ps_clip) != 2 or not 0 <= ps_clip[0] < ps_clip[1] <= 1:
281
+ raise ValueError(
282
+ 'ps_clip must be None or a pair 0 <= lower < upper <= 1')
283
+ pi_hat = np.clip(pi_hat_raw, ps_clip[0], ps_clip[1])
104
284
  fit_Y = True if Y_hat is None else False
105
285
  if fit_Y:
106
286
  _yhat_gb = Y.shape[0] * Y.shape[1] * A.shape[1] * 2 * 8 / 1e9
107
287
  _mem_limit_gb = kwargs.get('mem_limit_gb', None)
108
288
  if _mem_limit_gb is not None and _yhat_gb > _mem_limit_gb:
109
- import warnings
110
289
  warnings.warn(
111
290
  f"Y_hat allocation ({_yhat_gb:.1f} GB as float64) exceeds "
112
291
  f"mem_limit_gb={_mem_limit_gb} GB; using float32 to halve peak memory.",
@@ -116,16 +295,6 @@ def cross_fitting(
116
295
  else:
117
296
  Y_hat = np.zeros((Y.shape[0], Y.shape[1], A.shape[1], 2), dtype=float)
118
297
 
119
- # perform ECV at once
120
- if fit_pi and ps_model == 'random_forest_cv':
121
- info_ecv = run_ecv(X_A, A, **params_ps)
122
- func_ps, params_ps = _get_func_ps(ps_model, verbose=False, ecv=False,
123
- kwargs_ensemble=info_ecv['best_params_ensemble'], kwargs_regr=info_ecv['best_params_regr'])
124
- pprint.pprint('Best parameters for the regression model:')
125
- pprint.pprint(info_ecv['best_params_regr'])
126
- pprint.pprint('Best parameters for the ensemble model:')
127
- pprint.pprint(info_ecv['best_params_ensemble'])
128
-
129
298
  # Perform cross-fitting
130
299
  for train_index, test_index in folds:
131
300
  # Split data
@@ -134,28 +303,6 @@ def cross_fitting(
134
303
  A_train, A_test = A[train_index], A[test_index]
135
304
  Y_train, Y_test = Y[train_index], Y[test_index]
136
305
 
137
- if fit_pi:
138
- if verbose: pprint.pprint('Fit propensity score models...')
139
- i_ctrl = (np.sum(A_train, axis=1) == 0.)
140
-
141
- pi = np.zeros_like(A_test, dtype=float)
142
- for j in range(A.shape[1]):
143
- i_case = (A_train[:,j] == 1.)
144
-
145
- if mask is not None:
146
- i_cells = mask[:, j]
147
- else:
148
- i_ctrl = (np.sum(A_train, axis=1) == 0.)
149
- i_cells = i_ctrl | i_case
150
-
151
- if ps_model=='logistic' and XA_train.shape[1]==1 and np.all(XA_train==1):
152
- prob = np.sum(i_case)/np.sum(i_cells)
153
- pi[A_train[:,j] == 1., j] = prob
154
- pi[A_train[:,j] == 0., j] = 1 - prob
155
- else:
156
- pi[:,j] = func_ps(XA_train[i_cells], A_train[i_cells][:,j], XA_test)
157
- pi_hat[test_index] = pi
158
-
159
306
  if fit_Y:
160
307
  if verbose: pprint.pprint('Fit outcome models...')
161
308
  # Subset offset to training fold (for fitting) and test fold (for
@@ -174,15 +321,16 @@ def cross_fitting(
174
321
  Y_hat[test_index,:,:,0] = res[1][0]
175
322
  Y_hat[test_index,:,:,1] = res[1][1]
176
323
 
177
- pi_hat = np.clip(pi_hat, 0.01, 0.99)
178
324
  Y_hat = np.clip(Y_hat, None, 1e5)
325
+ if return_raw_pi:
326
+ return Y_hat, pi_hat, pi_hat_raw
179
327
  return Y_hat, pi_hat
180
328
 
181
329
 
182
330
 
183
331
 
184
332
 
185
- def AIPW_mean(Y, A, mu, pi, positive=False):
333
+ def AIPW_mean(Y, A, mu, pi):
186
334
  '''
187
335
  Augmented inverse probability weighted estimator (AIPW)
188
336
 
@@ -196,9 +344,6 @@ def AIPW_mean(Y, A, mu, pi, positive=False):
196
344
  Conditional outcome distribution estimate of shape (n, p, a, 2).
197
345
  pi : array
198
346
  Propensity score of shape (n, a, 2).
199
- positive : bool, optional
200
- Whether to restrict the pseudo-outcome to be positive.
201
-
202
347
  Returns
203
348
  -------
204
349
  tau : array
@@ -207,16 +352,18 @@ def AIPW_mean(Y, A, mu, pi, positive=False):
207
352
  Pseudo-outcome of shape (n, p, a, 2).
208
353
  '''
209
354
 
210
- weight = A / pi
355
+ with np.errstate(divide='ignore', invalid='ignore', over='ignore'):
356
+ weight = A / pi
211
357
  weight = weight[:, None, ...]
212
358
  Y = Y[:, :, None, None]
213
359
 
360
+ # Influence-function values are intentionally left unconstrained. Even
361
+ # for a nonnegative outcome, individual AIPW pseudo-outcomes may be
362
+ # negative; projecting them cell by cell changes their mean and biases the
363
+ # estimator. Parameter-space constraints belong after aggregation.
214
364
  pseudo_y = weight * (Y - mu) + mu
215
-
216
- if positive:
217
- pseudo_y = np.clip(pseudo_y, 0, None)
218
365
 
219
- tau = np.mean(pseudo_y, axis=0)
366
+ tau = np.mean(pseudo_y, axis=0, dtype=np.float64)
220
367
 
221
368
  return tau, pseudo_y
222
369
 
@@ -350,4 +497,4 @@ def fit_rf_ind_outcome(W, Y, A, *args, **kwargs):
350
497
  Y_pred = Y_pred.reshape(X.shape[0],1+a,Y.shape[1])
351
498
  Yhat_1 = Y_pred[:,1:,:].transpose(0,2,1)
352
499
  Yhat_0 = np.tile(Y_pred[:,0,:][:,:,None], (1,1,a))
353
- return Yhat_0, Yhat_1
500
+ return Yhat_0, Yhat_1
@@ -1,6 +1,8 @@
1
1
  import numpy as np
2
2
  import contextlib
3
3
  import pandas as pd
4
+ import warnings
5
+ from typing import Literal
4
6
  from causarray.DR_estimation import AIPW_mean, cross_fitting
5
7
  from causarray.gcate_glm import loess_fit, ls_fit
6
8
  import causarray.gcate_glm as _gcate_glm # for _backend_override
@@ -14,7 +16,9 @@ def compute_causal_estimand(
14
16
  Y, W, A, W_A=None, family='nb', offset=False,
15
17
  Y_hat=None, pi_hat=None, mask=None,
16
18
  fdx=False, fdx_B=1000, fdx_alpha=0.05, fdx_c=0.1,
17
- verbose=False, random_state=0, backend: str = "auto", **kwargs):
19
+ verbose=False, random_state=0, backend: str = "auto", K=1,
20
+ ps_clip=(0.01, 0.99), ps_class_weight='balanced',
21
+ **kwargs):
18
22
  """Estimate causal treatment effects using AIPW with a user-supplied estimand.
19
23
 
20
24
  Parameters
@@ -42,9 +46,17 @@ def compute_causal_estimand(
42
46
  pi_hat : array or None, shape (n, a)
43
47
  Pre-computed propensity scores. When provided, propensity fitting is
44
48
  skipped.
49
+ K : int
50
+ Number of folds used for nuisance estimation. The default ``1``
51
+ preserves in-sample fitting.
52
+ ps_clip : tuple(float, float)
53
+ Bounds applied to propensity scores used by AIPW.
54
+ ps_class_weight : str, dict or None
55
+ Class weighting for the propensity model. ``'balanced'`` preserves the
56
+ established nuisance fit; pass ``None`` for calibrated probabilities.
45
57
  mask : array or None, shape (n, a)
46
- Boolean mask indicating which cells are used for the final estimand
47
- computation. Does not affect cross-fitting or propensity estimation.
58
+ Boolean mask indicating eligible cells for each treatment. It limits
59
+ propensity-model fitting and final estimand computation.
48
60
  fdx : bool
49
61
  Whether to apply FDX control (``P(FDP > fdx_c) < fdx_alpha``).
50
62
  fdx_B : int
@@ -64,12 +76,17 @@ def compute_causal_estimand(
64
76
  Returns
65
77
  -------
66
78
  df_res : DataFrame
67
- Test results with columns ``gene_names``, ``tau``, ``std``, ``stat``,
68
- ``rej``, ``pvalue``, ``padj``, ``pvalue_emp_null_adj``,
69
- ``padj_emp_null_adj`` (and ``trt`` when ``A`` has multiple columns).
79
+ Test results produced by ``estimand``. An estimand may optionally return
80
+ a fifth dictionary whose arrays are added as diagnostic columns.
70
81
  """
71
82
  reset_random_seeds(random_state)
72
83
 
84
+ if 'clip_pseudo_outcomes' in kwargs:
85
+ raise TypeError(
86
+ 'clip_pseudo_outcomes has been removed; AIPW pseudo-outcomes are '
87
+ 'always left unclipped'
88
+ )
89
+
73
90
  ctx = _gcate_glm._backend_override(backend) if backend != "auto" else contextlib.nullcontext()
74
91
 
75
92
  # check the input data
@@ -85,7 +102,6 @@ def compute_causal_estimand(
85
102
  _mem_limit = kwargs.get('mem_limit_gb', None)
86
103
  _use_f32 = _mem_limit is not None and _yhat_gb > _mem_limit
87
104
  if _use_f32:
88
- import warnings
89
105
  warnings.warn(
90
106
  f"AIPW intermediate Y ({_yhat_gb:.1f} GB as float64) exceeds "
91
107
  f"mem_limit_gb={_mem_limit} GB; downcasting Y, A, and pi_hat to "
@@ -94,6 +110,8 @@ def compute_causal_estimand(
94
110
  )
95
111
  Y = Y.astype(np.float32 if _use_f32 else float)
96
112
  n, p = Y.shape
113
+ if not np.all(np.isfinite(Y)):
114
+ raise ValueError('Y must contain only finite values')
97
115
 
98
116
  if len(A.shape) == 1:
99
117
  A = A.reshape(-1,1)
@@ -135,11 +153,25 @@ def compute_causal_estimand(
135
153
  else:
136
154
  offset = None
137
155
  size_factors = np.ones(n)
156
+ size_factors = np.asarray(size_factors, dtype=float)
157
+ if size_factors.shape != (n,) or not np.all(np.isfinite(size_factors)) or np.any(size_factors <= 0):
158
+ raise ValueError('Size factors must be finite, positive, and have length n')
138
159
 
139
160
  with ctx:
140
- Y_hat, pi_hat = cross_fitting(Y, A, W, W_A, family=family, offset=offset,
141
- Y_hat=Y_hat, pi_hat=pi_hat, mask=mask, random_state=random_state, verbose=verbose, **kwargs)
161
+ Y_hat, pi_hat, pi_hat_raw = cross_fitting(
162
+ Y, A, W, W_A, family=family, K=K, offset=offset,
163
+ Y_hat=Y_hat, pi_hat=pi_hat, mask=mask, ps_clip=ps_clip,
164
+ ps_class_weight=ps_class_weight,
165
+ return_raw_pi=True, random_state=random_state, verbose=verbose,
166
+ **kwargs,
167
+ )
142
168
  pi_hat = pi_hat.reshape(*A.shape)
169
+ if not np.all(np.isfinite(Y_hat)):
170
+ raise ValueError('Outcome-model predictions must contain only finite values')
171
+ if not np.all(np.isfinite(pi_hat)) or np.any((pi_hat <= 0) | (pi_hat >= 1)):
172
+ raise ValueError(
173
+ 'Propensity scores used by AIPW must be finite and strictly between 0 and 1; '
174
+ 'pass ps_clip bounds or provide valid pi_hat values')
143
175
 
144
176
  if verbose: pprint.pprint('Estimating AIPW mean...')
145
177
  _aipw_dtype = Y_hat.dtype if _use_f32 else None
@@ -147,8 +179,12 @@ def compute_causal_estimand(
147
179
  pi_hat_aipw = pi_hat.astype(_aipw_dtype) if _aipw_dtype is not None else pi_hat
148
180
 
149
181
  # point estimation of the treatment effect
150
- _, etas = AIPW_mean(Y, np.stack([1-A_aipw, A_aipw], axis=-1),
151
- Y_hat, np.stack([1-pi_hat_aipw, pi_hat_aipw], axis=-1), positive=True)
182
+ _, etas = AIPW_mean(
183
+ Y,
184
+ np.stack([1-A_aipw, A_aipw], axis=-1),
185
+ Y_hat,
186
+ np.stack([1-pi_hat_aipw, pi_hat_aipw], axis=-1),
187
+ )
152
188
 
153
189
  # normalize the influence function values
154
190
  etas /= size_factors[:,None,None,None]
@@ -165,6 +201,7 @@ def compute_causal_estimand(
165
201
  _ret = estimand(etas[i_cells,:,j], A[i_cells,j], **kwargs)
166
202
  eta_est, tau_est, var_est = _ret[:3]
167
203
  df_est = _ret[3] if len(_ret) > 3 else None
204
+ estimand_info = _ret[4] if len(_ret) > 4 else None
168
205
 
169
206
  std_est = np.sqrt(var_est)
170
207
  tvalues_init = tau_est / std_est
@@ -187,20 +224,41 @@ def compute_causal_estimand(
187
224
  'pvalue_emp_null_adj': pvals_adj,
188
225
  'padj_emp_null_adj': qvals_adj,
189
226
  })
227
+ if estimand_info is not None:
228
+ for name, values in estimand_info.items():
229
+ df_res[name] = values
230
+ n_nonestimable = int(np.sum(~np.asarray(estimand_info['estimable'], dtype=bool)))
231
+ if n_nonestimable:
232
+ treatment = trt_names[j] if A.shape[1] > 1 else 0
233
+ warnings.warn(
234
+ f'{n_nonestimable} gene(s) for treatment {treatment!r} have '
235
+ 'nonfinite or nonpositive AIPW counterfactual means and are '
236
+ 'reported as non-estimable.',
237
+ RuntimeWarning, stacklevel=2,
238
+ )
190
239
  if A.shape[1]>1:
191
240
  df_res['trt'] = trt_names[j]
192
241
  res.append(df_res)
193
242
  df_res = pd.concat(res, axis=0).reset_index(drop=True)
194
- estimation = {**{'pi_hat':pi_hat, 'Y_hat':Y_hat, 'offset':offset, 'size_factors':size_factors}, **kwargs}
243
+ estimation = {**{
244
+ 'pi_hat': pi_hat,
245
+ 'pi_hat_raw': pi_hat_raw,
246
+ 'Y_hat': Y_hat,
247
+ 'offset': offset,
248
+ 'size_factors': size_factors,
249
+ 'ps_class_weight': ps_class_weight,
250
+ }, **kwargs}
195
251
  return df_res, estimation
196
252
 
197
253
 
198
254
  def LFC(
199
255
  Y, W, A, W_A=None, family='nb', offset=False,
200
- Y_hat=None, pi_hat=None, cross_est=False, mask=None, usevar='unequal',
256
+ Y_hat=None, pi_hat=None, cross_est=False, K=None, mask=None,
257
+ usevar: Literal['unequal', 'pooled'] = 'unequal',
201
258
  thres_min=1e-2, thres_diff=1e-2, eps_var=1e-4,
202
259
  fdx=False, fdx_alpha=0.05, fdx_c=0.1,
203
- verbose=False, backend: str = "auto", **kwargs):
260
+ verbose=False, backend: str = "auto", ps_clip=(0.01, 0.99),
261
+ ps_class_weight='balanced', **kwargs):
204
262
  """Estimate log-fold changes of treatment effects (LFCs) using AIPW.
205
263
 
206
264
  Fits a doubly-robust AIPW estimator for the log-ratio of counterfactual
@@ -230,17 +288,38 @@ def LFC(
230
288
  Pre-computed propensity scores. When provided, propensity fitting is
231
289
  skipped.
232
290
  cross_est : bool
233
- Whether to use cross-estimation for nuisance parameters.
291
+ Whether to use two-fold cross-estimation for nuisance parameters.
292
+ An explicit ``K`` takes precedence.
293
+ K : int or None
294
+ Number of nuisance-estimation folds. ``None`` uses 2 when
295
+ ``cross_est=True`` and 1 otherwise.
234
296
  mask : array or None, shape (n, a)
235
- Boolean mask indicating which cells are used for the final estimand
236
- computation. Does not affect cross-fitting or propensity estimation.
297
+ Boolean mask indicating eligible cells for each treatment. It limits
298
+ propensity-model fitting and final estimand computation.
237
299
  usevar : str
238
300
  Variance estimator for the AIPW pseudo-outcomes:
239
301
 
240
302
  * ``'unequal'`` (default, v0.0.6+): Welch variance
241
303
  ``s₀²/n₀ + s₁²/n₁`` with Welch-Satterthwaite degrees of
242
- freedom; p-values use the t-distribution.
304
+ freedom; p-values use the t-distribution. This is the recommended
305
+ choice for perturbation screens and for case-control, bulk, and
306
+ donor-level pseudo-bulk analyses. Independence of the rows does not
307
+ imply equal treatment-arm variances, and biological heterogeneity or
308
+ unequal effective sample sizes can make pooled inference
309
+ anti-conservative.
243
310
  * ``'pooled'``: pooled-variance estimator ``(s² + eps_var) / n``.
311
+ Treat this as an opt-in sensitivity or legacy analysis, not as the
312
+ default for small case-control studies. Use it only when equal arm
313
+ variances have a strong scientific and empirical justification.
314
+ Pooled inference can produce substantially smaller standard errors
315
+ and many more discoveries; donor-level independence alone is not a
316
+ justification for pooling.
317
+
318
+ ``'unequal'`` accommodates arm-specific variance but does not model
319
+ within-donor or within-subject correlation. Repeated cells from the
320
+ same biological unit should still be pseudo-bulked or analyzed with a
321
+ cluster-aware method; changing ``usevar`` alone does not remove
322
+ pseudoreplication.
244
323
 
245
324
  .. versionchanged:: 0.0.6
246
325
  Default changed from ``'pooled'`` to ``'unequal'``. The
@@ -267,25 +346,41 @@ def LFC(
267
346
  backend : str
268
347
  GLM backend: ``"auto"`` (default), ``"fast"`` (force crispyx),
269
348
  or ``"original"`` (force statsmodels).
349
+ ps_clip : tuple(float, float)
350
+ Bounds applied to propensity scores used by AIPW. Raw, unclipped scores
351
+ remain available as ``estimation['pi_hat_raw']``.
352
+ ps_class_weight : str, dict or None
353
+ Class weighting for the propensity model. ``'balanced'`` remains the
354
+ default to limit nuisance-model drift; pass ``None`` for calibrated
355
+ probabilities.
270
356
  **kwargs
271
357
  Additional arguments forwarded to the GLM fitting functions.
272
358
 
273
359
  Returns
274
360
  -------
275
361
  df_res : DataFrame
276
- Test results with columns ``gene_names``, ``tau``, ``std``, ``stat``,
277
- ``rej``, ``pvalue``, ``padj``, ``pvalue_emp_null_adj``,
278
- ``padj_emp_null_adj`` (and ``trt`` when ``A`` has multiple columns).
362
+ Test results with effect estimates and inference columns, raw
363
+ ``mean_control`` and ``mean_treated`` counterfactual means, and an
364
+ ``estimable`` flag (plus ``trt`` for multiple treatments).
279
365
  """
280
366
 
281
367
  def estimand(etas, A, **kwargs):
282
368
  eta_0, eta_1 = etas[..., 0], etas[..., 1]
283
- tau_0, tau_1 = np.mean(eta_0, axis=0), np.mean(eta_1, axis=0)
284
-
285
- tau_1 = np.clip(tau_1, thres_diff, None)
286
- tau_0 = np.clip(tau_0, thres_diff, None)
287
- tau_est = np.log(tau_1/tau_0)
288
- eta_est = eta_1 / tau_1[None,:] - eta_0 / tau_0[None,:]
369
+ mean_0 = np.mean(eta_0, axis=0, dtype=np.float64)
370
+ mean_1 = np.mean(eta_1, axis=0, dtype=np.float64)
371
+ finite_means = np.isfinite(mean_0) & np.isfinite(mean_1)
372
+ estimable = finite_means & (mean_0 > 0) & (mean_1 > 0)
373
+
374
+ # Apply the count-mean parameter-space constraint only after averaging
375
+ # the unmodified AIPW pseudo-outcomes. The floor makes the logarithm
376
+ # and its delta-method denominator numerically safe. Nonpositive raw
377
+ # means remain non-estimable and are filtered below, so this safeguard
378
+ # cannot turn an invalid aggregate into a discovery.
379
+ tau_0 = np.where(estimable, np.maximum(mean_0, thres_diff), 1.0)
380
+ tau_1 = np.where(estimable, np.maximum(mean_1, thres_diff), 1.0)
381
+ with np.errstate(invalid='ignore', divide='ignore', over='ignore'):
382
+ tau_est = np.log(tau_1/tau_0)
383
+ eta_est = eta_1 / tau_1[None,:] - eta_0 / tau_0[None,:]
289
384
 
290
385
  df_eff = None
291
386
  if usevar == 'pooled':
@@ -316,8 +411,13 @@ def LFC(
316
411
  else:
317
412
  raise ValueError('usevar must be either "pooled" or "unequal"')
318
413
 
319
- # filter out low-expressed genes
320
- idx = (np.maximum(np.abs(tau_0),np.abs(tau_1))<thres_min) | (np.abs(tau_1-tau_0)<thres_diff)
414
+ # Filter on the raw aggregate estimates, not on cell-level projections
415
+ # or the numerical floor used for the log transform.
416
+ idx = (
417
+ ~estimable |
418
+ (np.maximum(mean_0, mean_1) < thres_min) |
419
+ (np.abs(mean_1 - mean_0) < thres_diff)
420
+ )
321
421
  tau_est[idx] = 0.; eta_est[:,idx] = 0.; var_est[idx] = np.inf
322
422
  if df_eff is not None:
323
423
  df_eff[idx] = np.nan
@@ -338,13 +438,22 @@ def LFC(
338
438
  RuntimeWarning, stacklevel=3,
339
439
  )
340
440
 
341
- return eta_est, tau_est, var_est, df_eff
441
+ info = {
442
+ 'mean_control': mean_0,
443
+ 'mean_treated': mean_1,
444
+ 'estimable': estimable,
445
+ }
446
+ return eta_est, tau_est, var_est, df_eff, info
447
+
448
+ if K is None:
449
+ K = 2 if cross_est else 1
342
450
 
343
451
  return compute_causal_estimand(
344
452
  estimand, Y, W, A, W_A, family, offset,
345
453
  Y_hat=Y_hat, pi_hat=pi_hat, mask=mask,
346
454
  fdx=fdx, fdx_alpha=fdx_alpha, fdx_c=fdx_c,
347
- verbose=verbose, backend=backend, **kwargs)
455
+ verbose=verbose, backend=backend, K=K, ps_clip=ps_clip,
456
+ ps_class_weight=ps_class_weight, **kwargs)
348
457
 
349
458
 
350
459
 
@@ -524,7 +633,11 @@ def gcate_lfc_batch(
524
633
 
525
634
  lfc_kwargs : dict or None
526
635
  Extra keyword arguments forwarded to :func:`LFC`
527
- (e.g. ``usevar``, ``fdx``, ``thres_min``).
636
+ (e.g. ``usevar``, ``fdx``, ``thres_min``). Retain the default
637
+ ``usevar='unequal'`` for perturbation screens and for case-control,
638
+ bulk, or pseudo-bulk analyses. Pass
639
+ ``lfc_kwargs=dict(usevar='pooled')`` only for a deliberately justified
640
+ equal-variance sensitivity or legacy analysis.
528
641
  **kwargs
529
642
  Additional arguments forwarded to both :func:`fit_gcate_batch` and
530
643
  :func:`LFC`. When a key collides with ``gcate_kwargs`` /
@@ -725,4 +838,4 @@ def LFC_batch(*args, **kwargs):
725
838
  DeprecationWarning,
726
839
  stacklevel=2,
727
840
  )
728
- return gcate_lfc_batch(*args, **kwargs)
841
+ return gcate_lfc_batch(*args, **kwargs)
@@ -0,0 +1 @@
1
+ __version__ = "0.0.7"
@@ -10,10 +10,14 @@ __all__ = [
10
10
  'LFC', 'gcate_lfc_batch', 'LFC_batch',
11
11
  'fit_glm', 'fit_glm_fast', 'fit_glm_ondisk',
12
12
  'reset_random_seeds', 'fit_gcate', 'fit_gcate_batch',
13
+ 'estimate_propensity_scores', 'summarize_propensity_scores',
14
+ 'plot_propensity_scores',
13
15
  ]
14
16
 
15
17
 
16
18
  from causarray.DR_learner import LFC, gcate_lfc_batch, LFC_batch # ATE, SATE, FC
19
+ from causarray.DR_estimation import estimate_propensity_scores
20
+ from causarray.diagnostics import summarize_propensity_scores, plot_propensity_scores
17
21
  from causarray.gcate_glm import fit_glm
18
22
  from causarray.nb_glm_fast import fit_glm_fast, fit_glm_ondisk
19
23
  from causarray.utils import prep_causarray_data, reset_random_seeds, comp_size_factor
@@ -28,4 +32,4 @@ __maintainer__ = "Jin-Hong Du"
28
32
  __maintainer_email__ = "jinhongd@andrew.cmu.edu"
29
33
  __description__ = ("Causarray: A Python package for simultaneous causal inference"
30
34
  " with an array of outcomes."
31
- )
35
+ )
@@ -0,0 +1,172 @@
1
+ """Diagnostics for treatment overlap and propensity-score quality."""
2
+
3
+ import math
4
+
5
+ import matplotlib.pyplot as plt
6
+ import numpy as np
7
+ import pandas as pd
8
+ from sklearn.metrics import brier_score_loss, roc_auc_score
9
+
10
+
11
+ def _coerce_treatment_inputs(A, pi_hat, treatment_names=None):
12
+ if isinstance(A, pd.DataFrame):
13
+ if treatment_names is None:
14
+ treatment_names = list(A.columns)
15
+ A = A.values
16
+ else:
17
+ A = np.asarray(A)
18
+ if A.ndim == 1:
19
+ A = A[:, None]
20
+
21
+ pi_hat = np.asarray(pi_hat, dtype=float)
22
+ if pi_hat.ndim == 1:
23
+ pi_hat = pi_hat[:, None]
24
+ if A.shape != pi_hat.shape:
25
+ raise ValueError('A and pi_hat must have the same shape')
26
+ if not np.all(np.isin(A, (0, 1))):
27
+ raise ValueError('A must contain only binary treatment indicators')
28
+ if np.any(~np.isfinite(pi_hat)) or np.any((pi_hat < 0) | (pi_hat > 1)):
29
+ raise ValueError('pi_hat must contain finite probabilities in [0, 1]')
30
+
31
+ if treatment_names is None:
32
+ treatment_names = list(range(A.shape[1]))
33
+ if len(treatment_names) != A.shape[1]:
34
+ raise ValueError('treatment_names must match the number of treatments')
35
+ return A, pi_hat, list(treatment_names)
36
+
37
+
38
+ def _effective_sample_size(weights):
39
+ weights = np.asarray(weights, dtype=float)
40
+ denom = np.sum(weights ** 2)
41
+ return float(np.sum(weights) ** 2 / denom) if denom > 0 else np.nan
42
+
43
+
44
+ def summarize_propensity_scores(
45
+ A, pi_hat, treatment_names=None, overlap_bounds=(0.05, 0.95),
46
+ clip_bounds=(0.01, 0.99), bins=40,
47
+ ):
48
+ """Summarize overlap and inverse-weight stability for each treatment.
49
+
50
+ Other perturbations are excluded from a treatment's diagnostic comparison;
51
+ each row compares that treatment with shared all-zero controls.
52
+ """
53
+ A, pi_hat, treatment_names = _coerce_treatment_inputs(
54
+ A, pi_hat, treatment_names)
55
+ lo, hi = overlap_bounds
56
+ if not 0 <= lo < hi <= 1:
57
+ raise ValueError('overlap_bounds must satisfy 0 <= lower < upper <= 1')
58
+ if bins < 2:
59
+ raise ValueError('bins must be at least 2')
60
+
61
+ ctrl = np.sum(A, axis=1) == 0
62
+ if not np.any(ctrl):
63
+ raise ValueError('At least one all-zero control row is required')
64
+ rows = []
65
+ for j, name in enumerate(treatment_names):
66
+ case = A[:, j] == 1
67
+ if not np.any(case):
68
+ raise ValueError(f'Treatment {name} has no treated rows')
69
+ eligible = ctrl | case
70
+ y = A[eligible, j].astype(int)
71
+ p = pi_hat[eligible, j]
72
+ p_ctrl, p_case = p[y == 0], p[y == 1]
73
+
74
+ h_ctrl, edges = np.histogram(p_ctrl, bins=bins, range=(0, 1))
75
+ h_case, _ = np.histogram(p_case, bins=edges)
76
+ h_ctrl = h_ctrl / h_ctrl.sum() if h_ctrl.sum() else h_ctrl
77
+ h_case = h_case / h_case.sum() if h_case.sum() else h_case
78
+ overlap_ratio = float(np.minimum(h_ctrl, h_case).sum())
79
+
80
+ eps = np.finfo(float).eps
81
+ ess_ctrl = _effective_sample_size(1 / np.clip(1 - p_ctrl, eps, None))
82
+ ess_case = _effective_sample_size(1 / np.clip(p_case, eps, None))
83
+ clipped = np.zeros(p.shape, dtype=bool)
84
+ if clip_bounds is not None:
85
+ clipped = np.isclose(p, clip_bounds[0]) | np.isclose(p, clip_bounds[1])
86
+
87
+ rows.append({
88
+ 'treatment': name,
89
+ 'n_control': int((y == 0).sum()),
90
+ 'n_treated': int((y == 1).sum()),
91
+ 'prevalence': float(y.mean()),
92
+ 'overlap_ratio': overlap_ratio,
93
+ 'auc': float(roc_auc_score(y, p)),
94
+ 'brier_score': float(brier_score_loss(y, p)),
95
+ 'outside_overlap_fraction': float(np.mean((p < lo) | (p > hi))),
96
+ 'clipped_fraction': float(clipped.mean()),
97
+ 'ess_control': ess_ctrl,
98
+ 'ess_treated': ess_case,
99
+ 'ess_control_fraction': ess_ctrl / max(len(p_ctrl), 1),
100
+ 'ess_treated_fraction': ess_case / max(len(p_case), 1),
101
+ 'score_q01': float(np.quantile(p, 0.01)),
102
+ 'score_median': float(np.median(p)),
103
+ 'score_q99': float(np.quantile(p, 0.99)),
104
+ })
105
+ return pd.DataFrame(rows)
106
+
107
+
108
+ def plot_propensity_scores(
109
+ A, pi_hat, treatments=None, treatment_names=None, overlap_bounds=(0.05, 0.95),
110
+ bins=40, max_panels=4, axes=None,
111
+ ):
112
+ """Plot propensity distributions for treatment and control cells."""
113
+ A, pi_hat, treatment_names = _coerce_treatment_inputs(
114
+ A, pi_hat, treatment_names)
115
+ summary = summarize_propensity_scores(
116
+ A, pi_hat, treatment_names=treatment_names,
117
+ overlap_bounds=overlap_bounds, bins=bins,
118
+ )
119
+
120
+ if treatments is None:
121
+ indices = list(range(min(A.shape[1], max_panels)))
122
+ else:
123
+ indices = []
124
+ for treatment in treatments:
125
+ if treatment in treatment_names:
126
+ indices.append(treatment_names.index(treatment))
127
+ else:
128
+ index = int(treatment)
129
+ if index < 0 or index >= A.shape[1]:
130
+ raise ValueError(f'Unknown treatment: {treatment}')
131
+ indices.append(index)
132
+ if not indices:
133
+ raise ValueError('At least one treatment must be selected')
134
+
135
+ if axes is None:
136
+ ncols = min(2, len(indices))
137
+ nrows = math.ceil(len(indices) / ncols)
138
+ fig, axes = plt.subplots(
139
+ nrows, ncols, figsize=(5 * ncols, 3.5 * nrows), squeeze=False,
140
+ )
141
+ axes_flat = axes.ravel()
142
+ else:
143
+ axes_flat = np.asarray(axes, dtype=object).ravel()
144
+ if len(axes_flat) < len(indices):
145
+ raise ValueError('Not enough axes for the selected treatments')
146
+ fig = axes_flat[0].figure
147
+
148
+ ctrl = np.sum(A, axis=1) == 0
149
+ for ax, j in zip(axes_flat, indices):
150
+ case = A[:, j] == 1
151
+ ax.hist(
152
+ pi_hat[ctrl, j], bins=bins, range=(0, 1), density=True,
153
+ histtype='step', linewidth=1.6, label='control', color='#4c78a8',
154
+ )
155
+ ax.hist(
156
+ pi_hat[case, j], bins=bins, range=(0, 1), density=True,
157
+ histtype='step', linewidth=1.6, label=str(treatment_names[j]),
158
+ color='#e45756',
159
+ )
160
+ ax.axvline(overlap_bounds[0], color='#777777', linestyle='--', linewidth=0.8)
161
+ ax.axvline(overlap_bounds[1], color='#777777', linestyle='--', linewidth=0.8)
162
+ overlap = summary.loc[j, 'overlap_ratio']
163
+ ess = summary.loc[j, 'ess_treated_fraction']
164
+ ax.set_title(f'{treatment_names[j]} | overlap={overlap:.2f}, ESS={ess:.2f}')
165
+ ax.set_xlim(0, 1)
166
+ ax.set_xlabel('Estimated propensity score')
167
+ ax.set_ylabel('Density')
168
+ ax.legend(frameon=False)
169
+ for ax in axes_flat[len(indices):]:
170
+ ax.set_visible(False)
171
+ fig.tight_layout()
172
+ return fig, axes, summary
@@ -1 +0,0 @@
1
- __version__ = "0.0.6"
File without changes
File without changes
File without changes
File without changes