Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 17 additions & 6 deletions dte_adj/stratified.py
Original file line number Diff line number Diff line change
Expand Up @@ -246,8 +246,8 @@ def _compute_cumulative_distribution(
prediction = np.zeros((n_records, n_loc))
treatment_mask = treatment_arms == target_treatment_arm
folds = np.random.randint(self.folds, size=n_records)
_check_folds_have_training_data(folds, self.folds, treatment_mask)
strata = self.strata
_check_folds_have_training_data(folds, self.folds, treatment_mask, strata)
s_list = np.unique(strata)
if self.is_multi_task:
binomial = (outcomes.reshape(-1, 1) <= locations) * 1 # (n_records, n_loc)
Expand All @@ -260,16 +260,23 @@ def _compute_cumulative_distribution(
s_mask = strata == s
weight = (s_mask & treatment_mask).sum() / s_mask.sum()
superset_mask = (folds == fold) & s_mask
if not superset_mask.any():
continue
subset_train_mask = (folds != fold) & s_mask & treatment_mask
covariates_train = covariates[subset_train_mask]
binomial_train = binomial[subset_train_mask]
if len(np.unique(binomial_train)) > 1:
self.model = deepcopy(self.base_model)
self.model.fit(covariates_train, binomial_train)

pred = self._compute_model_prediction(
self.model, covariates[superset_mask]
)
pred = self._compute_model_prediction(
self.model, covariates[superset_mask]
)
else:
# All training labels are identical, so predict that constant
# rather than reusing a model fit on another fold or stratum.
pred = np.broadcast_to(
binomial_train[0], (superset_mask.sum(), n_loc)
)
prediction[superset_mask] = (
pred
+ treatment_mask[superset_mask].reshape(-1, 1)
Expand All @@ -295,6 +302,8 @@ def _compute_cumulative_distribution(
s_mask = strata == s
weight = (s_mask & treatment_mask).sum() / s_mask.sum()
superset_mask = (folds == fold) & s_mask
if not superset_mask.any():
continue
subset_train_mask = (folds != fold) & s_mask & treatment_mask
covariates_train = covariates[subset_train_mask]
binomial_train = binomial[subset_train_mask]
Expand Down Expand Up @@ -358,8 +367,8 @@ def _compute_interval_probability(
prediction = np.zeros((n_records, n_loc - 1))
treatment_mask = treatment_arms == target_treatment_arm
folds = np.random.randint(self.folds, size=n_records)
_check_folds_have_training_data(folds, self.folds, treatment_mask)
strata = self.strata
_check_folds_have_training_data(folds, self.folds, treatment_mask, strata)
s_list = np.unique(strata)
binominals = (outcomes[:, np.newaxis] <= locations) * 1 # (n_records, n_loc)
interval_iter = range(len(locations) - 1)
Expand All @@ -378,6 +387,8 @@ def _compute_interval_probability(
s_mask = strata == s
weight = (s_mask & treatment_mask).sum() / s_mask.sum()
superset_mask = (folds == fold) & s_mask
if not superset_mask.any():
continue
subset_train_mask = (folds != fold) & s_mask & treatment_mask
covariates_train = covariates[subset_train_mask]
binomial_train = binomial[subset_train_mask]
Expand Down
40 changes: 35 additions & 5 deletions dte_adj/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -91,18 +91,48 @@ def _prepare_fit_inputs(


def _check_folds_have_training_data(
folds: np.ndarray, n_folds: int, treatment_mask: np.ndarray
folds: np.ndarray,
n_folds: int,
treatment_mask: np.ndarray,
strata: np.ndarray,
) -> None:
"""Raise an informative error if cross-fitting would train on no data."""
"""Raise an informative error if cross-fitting would train on no data.

Every fold that contains observations of a stratum needs at least one observation of
the target treatment arm from the same stratum in the remaining folds, since the
held-out fold is predicted from them.
"""
advice = (
"This can happen by chance when the sample (or a treatment arm or stratum) is "
"small relative to the number of folds. Reduce `folds` (e.g. folds=2), merge "
"small strata, or use more data."
)
for fold in range(n_folds):
if not ((folds != fold) & treatment_mask).any():
raise ValueError(
f"Cross-fitting produced a fold ({fold} of {n_folds}) whose "
"complementary training set contains no observations of the target "
"treatment arm. This can happen by chance when the sample (or the "
"treatment arm) is small relative to the number of folds. "
"Reduce `folds` (e.g. folds=2) or use more data."
f"treatment arm. {advice}"
)
for s in np.unique(strata):
s_mask = strata == s
n_target = (s_mask & treatment_mask).sum()
if n_target == 0:
raise ValueError(
f"Stratum {s} contains no observations of the target treatment arm, so "
"its distribution function cannot be estimated. Merge it with another "
"stratum or drop it."
)
for fold in range(n_folds):
if (folds == fold)[s_mask].any() and not (
(folds != fold) & s_mask & treatment_mask
).any():
raise ValueError(
f"Cross-fitting produced a fold ({fold} of {n_folds}) for which "
f"stratum {s} has no training observations of the target treatment "
f"arm (the stratum has {n_target} such observation(s) in total). "
f"{advice}"
)


def _infer_default_locations(
Expand Down
62 changes: 61 additions & 1 deletion tests/test_simple_estimator.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,12 @@
import unittest
import numpy as np
from unittest.mock import patch, MagicMock
from sklearn.linear_model import LogisticRegression
from sklearn.linear_model import LogisticRegression, LinearRegression
from dte_adj import (
SimpleDistributionEstimator,
SimpleStratifiedDistributionEstimator,
AdjustedDistributionEstimator,
AdjustedStratifiedDistributionEstimator,
)

np.random.seed(123)
Expand Down Expand Up @@ -350,3 +351,62 @@ def test_empty_training_fold_raises_informative_error(self):
with self.assertRaises(ValueError) as cm:
est.predict(1, np.array([2.0]), display_progress=False)
self.assertIn("Reduce `folds`", str(cm.exception))


class TestSmallSampleFolds(unittest.TestCase):
def setUp(self):
self.X = np.random.RandomState(0).randn(8, 2)
self.D = np.array([0, 1, 0, 1, 0, 1, 0, 1])
self.Y = np.arange(8.0)

def test_empty_prediction_fold_is_skipped(self):
# folds=3 but no record is assigned to fold 2
est = AdjustedDistributionEstimator(LogisticRegression(), folds=3).fit(
self.X, self.D, self.Y
)
with patch(
"numpy.random.randint", return_value=np.array([0, 0, 1, 1, 0, 0, 1, 1])
):
dte, lower, upper = est.predict_dte(
1, 0, np.array([2.0, 4.0]), display_progress=False
)
self.assertTrue(np.all(np.isfinite(dte)))

def test_stratum_without_training_data_raises_informative_error(self):
strata = np.array([0, 0, 0, 0, 1, 1, 1, 1])
est = AdjustedStratifiedDistributionEstimator(
LogisticRegression(), folds=2
).fit(self.X, self.D, self.Y, strata)
# Both treated units of stratum 1 (indices 5 and 7) are in fold 0
with patch(
"numpy.random.randint", return_value=np.array([0, 0, 1, 1, 1, 0, 1, 0])
):
with self.assertRaises(ValueError) as cm:
est.predict(1, np.array([2.0]), display_progress=False)
self.assertIn("stratum 1", str(cm.exception))
self.assertIn("Reduce `folds`", str(cm.exception))

def test_stratum_without_target_arm_raises(self):
strata = np.array([0, 0, 0, 0, 1, 0, 1, 0]) # stratum 1 has only control units
est = AdjustedStratifiedDistributionEstimator(
LogisticRegression(), folds=2
).fit(self.X, self.D, self.Y, strata)
with self.assertRaises(ValueError) as cm:
est.predict(1, np.array([2.0]), display_progress=False)
self.assertIn("Stratum 1 contains no observations", str(cm.exception))

def test_multi_task_constant_labels_do_not_reuse_stale_model(self):
# In fold 0 / stratum 1 the training labels are all 1 (outcomes below every
# location), so no model is fit there and the constant must be predicted.
strata = np.array([0, 0, 0, 0, 1, 1, 1, 1])
Y = np.array([5.0, 6.0, 7.0, 8.0, 0.0, 0.0, 0.0, 0.0])
est = AdjustedStratifiedDistributionEstimator(
LinearRegression(), folds=2, is_multi_task=True
).fit(self.X, self.D, Y, strata)
with patch(
"numpy.random.randint", return_value=np.array([0, 0, 1, 1, 0, 0, 1, 1])
):
_, _, superset = est._compute_cumulative_distribution(
1, np.array([1.0, 2.0]), est.covariates, est.treatment_arms, Y
)
np.testing.assert_allclose(superset[4:], 1.0)
Loading