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
12 changes: 12 additions & 0 deletions docs/docs/tutorials/oregon.md
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
40 changes: 40 additions & 0 deletions dte_adj/local.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
"""
Expand All @@ -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:
Expand Down Expand Up @@ -139,6 +147,8 @@ def predict_ldte(
locations,
alpha,
display_progress,
variance_type=variance_type,
n_bootstrap=n_bootstrap,
)

def predict_lpte(
Expand All @@ -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]:
"""
Expand All @@ -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:
Expand Down Expand Up @@ -217,6 +235,8 @@ def predict_lpte(
locations,
alpha,
display_progress,
variance_type=variance_type,
n_bootstrap=n_bootstrap,
)


Expand Down Expand Up @@ -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]:
"""
Expand All @@ -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:
Expand Down Expand Up @@ -335,6 +363,8 @@ def predict_ldte(
locations,
alpha,
display_progress,
variance_type=variance_type,
n_bootstrap=n_bootstrap,
)

def predict_lpte(
Expand All @@ -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]:
"""
Expand All @@ -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:
Expand Down Expand Up @@ -414,4 +452,6 @@ def predict_lpte(
locations,
alpha,
display_progress,
variance_type=variance_type,
n_bootstrap=n_bootstrap,
)
91 changes: 91 additions & 0 deletions dte_adj/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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.
Expand All @@ -267,13 +325,22 @@ 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:
- Expected effects (beta)
- 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
Expand Down Expand Up @@ -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
Expand All @@ -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.
Expand All @@ -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:
Expand All @@ -429,6 +512,8 @@ def compute_ldte(
alpha,
use_intervals=False,
display_progress=display_progress,
variance_type=variance_type,
n_bootstrap=n_bootstrap,
)


Expand All @@ -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.
Expand All @@ -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:
Expand All @@ -465,4 +554,6 @@ def compute_lpte(
alpha,
use_intervals=True,
display_progress=display_progress,
variance_type=variance_type,
n_bootstrap=n_bootstrap,
)
73 changes: 73 additions & 0 deletions tests/test_local_estimators.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Loading
Loading