diff --git a/docs/docs/tutorials/oregon.md b/docs/docs/tutorials/oregon.md index b417cb1..da0d6c9 100644 --- a/docs/docs/tutorials/oregon.md +++ b/docs/docs/tutorials/oregon.md @@ -182,6 +182,18 @@ plt.tight_layout() plt.show() ``` +By default the confidence intervals are pointwise and analytic (`variance_type="moment"`). As with `predict_dte`, `predict_ldte` and `predict_lpte` also accept `variance_type="multiplier"` (pointwise multiplier bootstrap) and `variance_type="uniform"` (a band that holds simultaneously over all locations), with `n_bootstrap` controlling the number of draws: + +```python +ldte_ml, lower_uniform, upper_uniform = ml_local_estimator.predict_ldte( + target_treatment_arm=1, + control_treatment_arm=0, + locations=outcome_ed_costs_locations, + variance_type="uniform", + n_bootstrap=500, +) +``` + The analysis produces the following local distribution treatment effects visualization: ![Oregon Health Insurance Experiment LDTE Analysis](../assets/oregon_ldte_costs_comparison.png) diff --git a/dte_adj/local.py b/dte_adj/local.py index f1a8b59..7d4a48b 100644 --- a/dte_adj/local.py +++ b/dte_adj/local.py @@ -73,6 +73,8 @@ def predict_ldte( control_treatment_arm: int, locations: Optional[np.ndarray] = None, alpha: float = 0.05, + variance_type: str = "moment", + n_bootstrap: int = 500, display_progress: bool = True, ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: """ @@ -91,6 +93,12 @@ def predict_ldte( distribution via ``np.histogram_bin_edges(outcomes, bins='auto')``. The actual array used is stored on ``self.last_locations``. alpha (float, optional): Significance level of the confidence bound. Defaults to 0.05. + variance_type (str, optional): Variance type to be used to compute confidence intervals. + Available values are "moment" (analytic, pointwise), "multiplier" (pointwise + multiplier bootstrap), and "uniform" (uniform band over all locations via + multiplier bootstrap). Defaults to "moment". + n_bootstrap (int, optional): Number of bootstrap draws for "multiplier" and + "uniform". Defaults to 500. display_progress (bool, optional): Whether to display a progress bar. Defaults to True. Returns: @@ -139,6 +147,8 @@ def predict_ldte( locations, alpha, display_progress, + variance_type=variance_type, + n_bootstrap=n_bootstrap, ) def predict_lpte( @@ -147,6 +157,8 @@ def predict_lpte( control_treatment_arm: int, locations: Optional[np.ndarray] = None, alpha: float = 0.05, + variance_type: str = "moment", + n_bootstrap: int = 500, display_progress: bool = True, ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: """ @@ -167,6 +179,12 @@ def predict_lpte( ``np.histogram_bin_edges(outcomes, bins='auto')``. The actual array used is stored on ``self.last_locations``. alpha (float, optional): Significance level of the confidence bound. Defaults to 0.05. + variance_type (str, optional): Variance type to be used to compute confidence intervals. + Available values are "moment" (analytic, pointwise), "multiplier" (pointwise + multiplier bootstrap), and "uniform" (uniform band over all locations via + multiplier bootstrap). Defaults to "moment". + n_bootstrap (int, optional): Number of bootstrap draws for "multiplier" and + "uniform". Defaults to 500. display_progress (bool, optional): Whether to display a progress bar. Defaults to True. Returns: @@ -217,6 +235,8 @@ def predict_lpte( locations, alpha, display_progress, + variance_type=variance_type, + n_bootstrap=n_bootstrap, ) @@ -269,6 +289,8 @@ def predict_ldte( control_treatment_arm: int, locations: Optional[np.ndarray] = None, alpha: float = 0.05, + variance_type: str = "moment", + n_bootstrap: int = 500, display_progress: bool = True, ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: """ @@ -286,6 +308,12 @@ def predict_ldte( distribution via ``np.histogram_bin_edges(outcomes, bins='auto')``. The actual array used is stored on ``self.last_locations``. alpha (float, optional): Significance level of the confidence bound. Defaults to 0.05. + variance_type (str, optional): Variance type to be used to compute confidence intervals. + Available values are "moment" (analytic, pointwise), "multiplier" (pointwise + multiplier bootstrap), and "uniform" (uniform band over all locations via + multiplier bootstrap). Defaults to "moment". + n_bootstrap (int, optional): Number of bootstrap draws for "multiplier" and + "uniform". Defaults to 500. display_progress (bool, optional): Whether to display a progress bar. Defaults to True. Returns: @@ -335,6 +363,8 @@ def predict_ldte( locations, alpha, display_progress, + variance_type=variance_type, + n_bootstrap=n_bootstrap, ) def predict_lpte( @@ -343,6 +373,8 @@ def predict_lpte( control_treatment_arm: int, locations: Optional[np.ndarray] = None, alpha: float = 0.05, + variance_type: str = "moment", + n_bootstrap: int = 500, display_progress: bool = True, ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: """ @@ -362,6 +394,12 @@ def predict_lpte( ``np.histogram_bin_edges(outcomes, bins='auto')``. The actual array used is stored on ``self.last_locations``. alpha (float, optional): Significance level of the confidence bound. Defaults to 0.05. + variance_type (str, optional): Variance type to be used to compute confidence intervals. + Available values are "moment" (analytic, pointwise), "multiplier" (pointwise + multiplier bootstrap), and "uniform" (uniform band over all locations via + multiplier bootstrap). Defaults to "moment". + n_bootstrap (int, optional): Number of bootstrap draws for "multiplier" and + "uniform". Defaults to 500. display_progress (bool, optional): Whether to display a progress bar. Defaults to True. Returns: @@ -414,4 +452,6 @@ def predict_lpte( locations, alpha, display_progress, + variance_type=variance_type, + n_bootstrap=n_bootstrap, ) diff --git a/dte_adj/util.py b/dte_adj/util.py index 0b518c4..574e91f 100644 --- a/dte_adj/util.py +++ b/dte_adj/util.py @@ -247,6 +247,62 @@ def compute_confidence_intervals( raise ValueError(f"Invalid variance type was specified: {variance_type}") +def _multiplier_bootstrap_bands( + estimate: np.ndarray, + influence_function: np.ndarray, + alpha: float, + variance_type: str, + n_bootstrap: int, +) -> Tuple[np.ndarray, np.ndarray]: + """Confidence bands from a multiplier bootstrap of per-observation influence functions. + + Each draw reweights the influence functions with i.i.d. multipliers of mean zero and + unit variance, ``xi = eta1 / sqrt(2) + (eta2 ** 2 - 1) / 2`` (the same multipliers as + :func:`compute_confidence_intervals`), so stratum structure that is already encoded in + the influence functions is preserved without refitting any model. + + Args: + estimate (np.ndarray): Point estimates, shape (n_loc,). + influence_function (np.ndarray): Influence function of each observation, shape + (n_obs, n_loc), such that ``mean(influence_function**2, axis=0) / n_obs`` is the + asymptotic variance of ``estimate``. + alpha (float): Significance level. + variance_type (str): "multiplier" for pointwise bands or "uniform" for uniform bands + (simultaneous over all locations, via the max-t statistic). + n_bootstrap (int): Number of bootstrap draws. + + Returns: + Tuple[np.ndarray, np.ndarray]: Lower and upper bounds. + """ + num_obs = influence_function.shape[0] + omega = (influence_function**2).mean(axis=0) + + boot_draw = np.zeros((n_bootstrap, influence_function.shape[1])) + for b in range(n_bootstrap): + eta1 = np.random.normal(0, 1, num_obs) + eta2 = np.random.normal(0, 1, num_obs) + xi = eta1 / np.sqrt(2) + (eta2**2 - 1) / 2 + boot_draw[b] = (xi[:, np.newaxis] * influence_function).mean(axis=0) + + if variance_type == "multiplier": + se = boot_draw.std(axis=0) + return estimate + norm.ppf(alpha / 2) * se, estimate + norm.ppf( + 1 - alpha / 2 + ) * se + + # Uniform band: critical value from the max of studentized draws, ignoring locations + # with (numerically) zero variance, e.g. a CDF evaluated at or above the maximum outcome. + valid = omega > 1e-12 * max(omega.max(), 1e-300) + if not valid.any(): + return estimate.copy(), estimate.copy() + tstats = np.abs(boot_draw[:, valid]) / np.sqrt(omega[valid] / num_obs) + critical_value = np.quantile(tstats.max(axis=1), 1 - alpha) + se = (np.quantile(boot_draw, 0.75, axis=0) - np.quantile(boot_draw, 0.25, axis=0)) / ( + norm.ppf(0.75) - norm.ppf(0.25) + ) + return estimate - critical_value * se, estimate + critical_value * se + + def _compute_local_treatment_effects_core( estimator: "SimpleStratifiedDistributionEstimator | AdjustedLocalDistributionEstimator", target_treatment_arm: int, @@ -255,6 +311,8 @@ def _compute_local_treatment_effects_core( alpha: float, use_intervals: bool = False, display_progress: bool = False, + variance_type: str = "moment", + n_bootstrap: int = 500, ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: """ Core computation logic shared between LDTE and LPTE. @@ -267,6 +325,9 @@ def _compute_local_treatment_effects_core( alpha (float): Significance level of the confidence bound. use_intervals (bool): If True, compute interval probabilities (LPTE), else cumulative (LDTE). display_progress (bool): Whether to display a progress bar. + variance_type (str): "moment" (analytic), "multiplier" (pointwise multiplier + bootstrap) or "uniform" (uniform band via multiplier bootstrap). + n_bootstrap (int): Number of bootstrap draws for "multiplier" and "uniform". Returns: Tuple[np.ndarray, np.ndarray, np.ndarray]: A tuple containing: @@ -274,6 +335,12 @@ def _compute_local_treatment_effects_core( - Lower bounds - Upper bounds """ + if variance_type not in ("moment", "multiplier", "uniform"): + raise ValueError( + f"Invalid variance type was specified: {variance_type}. " + "Available values are moment, multiplier, and uniform." + ) + X = estimator.covariates Z = estimator.treatment_arms D = estimator.treatment_indicator @@ -383,6 +450,18 @@ def xi(s): xi_2_dict = {s: xi(s) for s in s_list} xi_2 = np.array([xi_2_dict[s] for s in S]) + if variance_type != "moment": + # Per-observation influence function. xi_t / xi_c are centered within each + # stratum x arm cell and xi_2 is constant within a stratum, so the cross terms + # vanish and mean(influence**2) equals the analytic sigma below. + influence = ( + Z.reshape(-1, 1) * xi_t + (1 - Z).reshape(-1, 1) * xi_c + xi_2 + ) / psi_b.mean() + lower_bound, upper_bound = _multiplier_bootstrap_bands( + beta, influence, alpha, variance_type, n_bootstrap + ) + return beta, lower_bound, upper_bound + sigma = ( Z.reshape(-1, 1) * xi_t**2 + (1 - Z).reshape(-1, 1) * xi_c**2 + xi_2**2 ).mean(axis=0) / (psi_b.mean()) ** 2 @@ -403,6 +482,8 @@ def compute_ldte( locations: np.ndarray, alpha: float = 0.05, display_progress: bool = False, + variance_type: str = "moment", + n_bootstrap: int = 500, ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: """ Compute Local Distribution Treatment Effects (LDTE) using the provided formula. @@ -414,6 +495,8 @@ def compute_ldte( locations (np.ndarray): Scalar values to be used for computing the cumulative distribution. alpha (float, optional): Significance level of the confidence bound. Defaults to 0.05. display_progress (bool, optional): Whether to display a progress bar. Defaults to False. + variance_type (str, optional): "moment", "multiplier", or "uniform". Defaults to "moment". + n_bootstrap (int, optional): Number of bootstrap draws. Defaults to 500. Returns: Tuple[np.ndarray, np.ndarray, np.ndarray]: A tuple containing: @@ -429,6 +512,8 @@ def compute_ldte( alpha, use_intervals=False, display_progress=display_progress, + variance_type=variance_type, + n_bootstrap=n_bootstrap, ) @@ -439,6 +524,8 @@ def compute_lpte( locations: np.ndarray, alpha: float = 0.05, display_progress: bool = False, + variance_type: str = "moment", + n_bootstrap: int = 500, ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: """ Compute Local Probability Treatment Effects (LPTE) using the provided formula. @@ -450,6 +537,8 @@ def compute_lpte( locations (np.ndarray): Scalar values to be used for computing the interval probabilities. alpha (float, optional): Significance level of the confidence bound. Defaults to 0.05. display_progress (bool, optional): Whether to display a progress bar. Defaults to False. + variance_type (str, optional): "moment", "multiplier", or "uniform". Defaults to "moment". + n_bootstrap (int, optional): Number of bootstrap draws. Defaults to 500. Returns: Tuple[np.ndarray, np.ndarray, np.ndarray]: A tuple containing: @@ -465,4 +554,6 @@ def compute_lpte( alpha, use_intervals=True, display_progress=display_progress, + variance_type=variance_type, + n_bootstrap=n_bootstrap, ) diff --git a/tests/test_local_estimators.py b/tests/test_local_estimators.py index 5d5d96e..8937222 100644 --- a/tests/test_local_estimators.py +++ b/tests/test_local_estimators.py @@ -391,3 +391,76 @@ def test_e2e(self): ), "Adjusted estimator does not have narrower intervals", ) + + +class TestLocalBootstrapInference(unittest.TestCase): + @classmethod + def setUpClass(cls): + np.random.seed(0) + data = generate_data(n=2000) + cls.locations = np.arange(0, 12, 2.0) + cls.estimator = SimpleLocalDistributionEstimator().fit( + data["X"], data["Z"], data["D"], data["Y"], data["strata"] + ) + cls.data = data + cls.moment = cls.estimator.predict_ldte( + 1, 0, cls.locations, display_progress=False + ) + + def test_default_is_moment(self): + explicit = self.estimator.predict_ldte( + 1, 0, self.locations, variance_type="moment", display_progress=False + ) + for a, b in zip(self.moment, explicit): + np.testing.assert_allclose(a, b) + + def test_multiplier_matches_moment(self): + np.random.seed(1) + beta, lower, upper = self.estimator.predict_ldte( + 1, + 0, + self.locations, + variance_type="multiplier", + n_bootstrap=2000, + display_progress=False, + ) + np.testing.assert_allclose(beta, self.moment[0]) + moment_half = (self.moment[2] - self.moment[1]) / 2 + np.testing.assert_allclose((upper - lower) / 2, moment_half, rtol=0.1) + + def test_uniform_band_is_wider_than_pointwise(self): + np.random.seed(2) + beta, lower, upper = self.estimator.predict_ldte( + 1, + 0, + self.locations, + variance_type="uniform", + n_bootstrap=1000, + display_progress=False, + ) + self.assertTrue(np.all(lower <= beta) and np.all(beta <= upper)) + self.assertTrue(np.all(upper - lower > self.moment[2] - self.moment[1])) + + def test_lpte_bootstrap_for_adjusted_estimator(self): + np.random.seed(3) + d = self.data + estimator = AdjustedLocalDistributionEstimator( + LogisticRegression(), folds=2 + ).fit(d["X"], d["Z"], d["D"], d["Y"], d["strata"]) + for variance_type in ("multiplier", "uniform"): + beta, lower, upper = estimator.predict_lpte( + 1, + 0, + self.locations, + variance_type=variance_type, + n_bootstrap=100, + display_progress=False, + ) + self.assertEqual(beta.shape, (len(self.locations) - 1,)) + self.assertTrue(np.all(lower <= beta) and np.all(beta <= upper)) + + def test_invalid_variance_type(self): + with self.assertRaises(ValueError): + self.estimator.predict_ldte( + 1, 0, self.locations, variance_type="simple", display_progress=False + ) diff --git a/tests/test_utils.py b/tests/test_utils.py index aace50f..2b19bd8 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -92,3 +92,32 @@ def test_for_intervals_constant_outcomes(self): outcomes = np.full(50, 3.0) result = _infer_default_locations(outcomes, for_intervals=True) self.assertLess(result[0], 3.0) + + +class TestMultiplierBootstrapBands(unittest.TestCase): + def test_pointwise_se_matches_influence_variance(self): + from dte_adj.util import _multiplier_bootstrap_bands + + np.random.seed(0) + n = 2000 + influence = np.random.randn(n, 3) * np.array([1.0, 2.0, 0.5]) + estimate = np.zeros(3) + lower, upper = _multiplier_bootstrap_bands( + estimate, influence, 0.05, "multiplier", 2000 + ) + expected_se = np.sqrt((influence**2).mean(axis=0) / n) + np.testing.assert_allclose((upper - lower) / (2 * 1.959964), expected_se, rtol=0.1) + + def test_uniform_ignores_zero_variance_locations(self): + from dte_adj.util import _multiplier_bootstrap_bands + + np.random.seed(0) + influence = np.random.randn(500, 3) + influence[:, 2] = 0.0 + estimate = np.array([0.1, 0.2, 1.0]) + lower, upper = _multiplier_bootstrap_bands( + estimate, influence, 0.05, "uniform", 200 + ) + self.assertTrue(np.all(np.isfinite(lower)) and np.all(np.isfinite(upper))) + self.assertAlmostEqual(lower[2], 1.0) + self.assertAlmostEqual(upper[2], 1.0)