diff-diff 1.3.1__py3-none-any.whl → 1.4.0__py3-none-any.whl

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.
diff_diff/__init__.py CHANGED
@@ -103,7 +103,7 @@ from diff_diff.visualization import (
103
103
  plot_sensitivity,
104
104
  )
105
105
 
106
- __version__ = "1.3.1"
106
+ __version__ = "1.4.0"
107
107
  __all__ = [
108
108
  # Estimators
109
109
  "DifferenceInDifferences",
diff_diff/estimators.py CHANGED
@@ -17,12 +17,12 @@ from typing import Any, Dict, List, Optional, Tuple
17
17
  import numpy as np
18
18
  import pandas as pd
19
19
 
20
+ from diff_diff.linalg import compute_r_squared, compute_robust_vcov, solve_ols
20
21
  from diff_diff.results import DiDResults, MultiPeriodDiDResults, PeriodEffect
21
22
  from diff_diff.utils import (
22
23
  WildBootstrapResults,
23
24
  compute_confidence_interval,
24
25
  compute_p_value,
25
- compute_robust_se,
26
26
  validate_binary,
27
27
  wild_bootstrap_se,
28
28
  )
@@ -261,8 +261,11 @@ class DifferenceInDifferences:
261
261
  X = np.column_stack([X, dummies[col].values.astype(float)])
262
262
  var_names.append(col)
263
263
 
264
- # Fit OLS
265
- coefficients, residuals, fitted, r_squared = self._fit_ols(X, y)
264
+ # Fit OLS using unified backend
265
+ coefficients, residuals, fitted, vcov = solve_ols(
266
+ X, y, return_fitted=True, return_vcov=False
267
+ )
268
+ r_squared = compute_r_squared(y, residuals)
266
269
 
267
270
  # Extract ATT (coefficient on interaction term)
268
271
  att_idx = 3 # Index of interaction term
@@ -285,13 +288,13 @@ class DifferenceInDifferences:
285
288
  )
286
289
  elif self.cluster is not None:
287
290
  cluster_ids = data[self.cluster].values
288
- vcov = compute_robust_se(X, residuals, cluster_ids)
291
+ vcov = compute_robust_vcov(X, residuals, cluster_ids)
289
292
  se = np.sqrt(vcov[att_idx, att_idx])
290
293
  t_stat = att / se
291
294
  p_value = compute_p_value(t_stat, df=df)
292
295
  conf_int = compute_confidence_interval(att, se, self.alpha, df=df)
293
296
  elif self.robust:
294
- vcov = compute_robust_se(X, residuals)
297
+ vcov = compute_robust_vcov(X, residuals)
295
298
  se = np.sqrt(vcov[att_idx, att_idx])
296
299
  t_stat = att / se
297
300
  p_value = compute_p_value(t_stat, df=df)
@@ -300,7 +303,7 @@ class DifferenceInDifferences:
300
303
  # Classical OLS standard errors
301
304
  n = len(y)
302
305
  k = X.shape[1]
303
- mse = np.sum(residuals ** 2) / (n - k)
306
+ mse = np.sum(residuals**2) / (n - k)
304
307
  # Use solve() instead of inv() for numerical stability
305
308
  # solve(A, B) computes X where AX=B, so this yields (X'X)^{-1} * mse
306
309
  vcov = np.linalg.solve(X.T @ X, mse * np.eye(k))
@@ -352,10 +355,15 @@ class DifferenceInDifferences:
352
355
 
353
356
  return self.results_
354
357
 
355
- def _fit_ols(self, X: np.ndarray, y: np.ndarray) -> Tuple[np.ndarray, np.ndarray, np.ndarray, float]:
358
+ def _fit_ols(
359
+ self, X: np.ndarray, y: np.ndarray
360
+ ) -> Tuple[np.ndarray, np.ndarray, np.ndarray, float]:
356
361
  """
357
362
  Fit OLS regression.
358
363
 
364
+ This method is kept for backwards compatibility. Internally uses the
365
+ unified solve_ols from diff_diff.linalg for optimized computation.
366
+
359
367
  Parameters
360
368
  ----------
361
369
  X : np.ndarray
@@ -367,32 +375,12 @@ class DifferenceInDifferences:
367
375
  -------
368
376
  tuple
369
377
  (coefficients, residuals, fitted_values, r_squared)
370
-
371
- Raises
372
- ------
373
- ValueError
374
- If design matrix is rank-deficient (perfect multicollinearity).
375
378
  """
376
- # Check for rank deficiency (perfect multicollinearity)
377
- rank = np.linalg.matrix_rank(X)
378
- if rank < X.shape[1]:
379
- raise ValueError(
380
- f"Design matrix is rank-deficient (rank {rank} < {X.shape[1]} columns). "
381
- "This indicates perfect multicollinearity. Check your fixed effects "
382
- "and covariates for linear dependencies."
383
- )
384
-
385
- # Solve normal equations: β = (X'X)^(-1) X'y
386
- coefficients = np.linalg.lstsq(X, y, rcond=None)[0]
387
-
388
- # Compute fitted values and residuals
389
- fitted = X @ coefficients
390
- residuals = y - fitted
391
-
392
- # Compute R-squared
393
- ss_res = np.sum(residuals ** 2)
394
- ss_tot = np.sum((y - np.mean(y)) ** 2)
395
- r_squared = 1 - (ss_res / ss_tot) if ss_tot > 0 else 0.0
379
+ # Use unified OLS backend
380
+ coefficients, residuals, fitted, _ = solve_ols(
381
+ X, y, return_fitted=True, return_vcov=False
382
+ )
383
+ r_squared = compute_r_squared(y, residuals)
396
384
 
397
385
  return coefficients, residuals, fitted, r_squared
398
386
 
@@ -442,7 +430,7 @@ class DifferenceInDifferences:
442
430
  t_stat = bootstrap_results.t_stat_original
443
431
 
444
432
  # Also compute vcov for storage (using cluster-robust for consistency)
445
- vcov = compute_robust_se(X, residuals, cluster_ids)
433
+ vcov = compute_robust_vcov(X, residuals, cluster_ids)
446
434
 
447
435
  return se, p_value, conf_int, t_stat, vcov, bootstrap_results
448
436
 
@@ -889,8 +877,11 @@ class MultiPeriodDiD(DifferenceInDifferences):
889
877
  X = np.column_stack([X, dummies[col].values.astype(float)])
890
878
  var_names.append(col)
891
879
 
892
- # Fit OLS
893
- coefficients, residuals, fitted, r_squared = self._fit_ols(X, y)
880
+ # Fit OLS using unified backend
881
+ coefficients, residuals, fitted, _ = solve_ols(
882
+ X, y, return_fitted=True, return_vcov=False
883
+ )
884
+ r_squared = compute_r_squared(y, residuals)
894
885
 
895
886
  # Degrees of freedom
896
887
  df = len(y) - X.shape[1] - n_absorbed_effects
@@ -900,13 +891,13 @@ class MultiPeriodDiD(DifferenceInDifferences):
900
891
  # For now, we use analytical inference even if inference="wild_bootstrap"
901
892
  if self.cluster is not None:
902
893
  cluster_ids = data[self.cluster].values
903
- vcov = compute_robust_se(X, residuals, cluster_ids)
894
+ vcov = compute_robust_vcov(X, residuals, cluster_ids)
904
895
  elif self.robust:
905
- vcov = compute_robust_se(X, residuals)
896
+ vcov = compute_robust_vcov(X, residuals)
906
897
  else:
907
898
  n = len(y)
908
899
  k = X.shape[1]
909
- mse = np.sum(residuals ** 2) / (n - k)
900
+ mse = np.sum(residuals**2) / (n - k)
910
901
  # Use solve() instead of inv() for numerical stability
911
902
  # solve(A, B) computes X where AX=B, so this yields (X'X)^{-1} * mse
912
903
  vcov = np.linalg.solve(X.T @ X, mse * np.eye(k))
diff_diff/linalg.py ADDED
@@ -0,0 +1,271 @@
1
+ """
2
+ Unified linear algebra backend for diff-diff.
3
+
4
+ This module provides optimized OLS and variance estimation that can be
5
+ swapped to a compiled backend (Rust/C++) for maximum performance.
6
+
7
+ The key optimizations are:
8
+ 1. scipy.linalg.lstsq with 'gelsy' driver (QR-based, faster than SVD)
9
+ 2. Vectorized cluster-robust SE via groupby (eliminates O(n*clusters) loop)
10
+ 3. Single interface for all estimators (reduces code duplication)
11
+
12
+ Future: This module can be extended with a Rust backend for additional speedup.
13
+ """
14
+
15
+ from typing import Optional, Tuple, Union
16
+
17
+ import numpy as np
18
+ import pandas as pd
19
+ from scipy.linalg import lstsq as scipy_lstsq
20
+
21
+
22
+ def solve_ols(
23
+ X: np.ndarray,
24
+ y: np.ndarray,
25
+ *,
26
+ cluster_ids: Optional[np.ndarray] = None,
27
+ return_vcov: bool = True,
28
+ return_fitted: bool = False,
29
+ check_finite: bool = True,
30
+ ) -> Union[
31
+ Tuple[np.ndarray, np.ndarray, Optional[np.ndarray]],
32
+ Tuple[np.ndarray, np.ndarray, np.ndarray, Optional[np.ndarray]],
33
+ ]:
34
+ """
35
+ Solve OLS regression with optional clustered standard errors.
36
+
37
+ This is the unified OLS solver for all diff-diff estimators. It uses
38
+ scipy's optimized LAPACK routines and vectorized variance estimation.
39
+
40
+ Parameters
41
+ ----------
42
+ X : ndarray of shape (n, k)
43
+ Design matrix (should include intercept if desired).
44
+ y : ndarray of shape (n,)
45
+ Response vector.
46
+ cluster_ids : ndarray of shape (n,), optional
47
+ Cluster identifiers for cluster-robust standard errors.
48
+ If None, HC1 (heteroskedasticity-robust) SEs are computed.
49
+ return_vcov : bool, default True
50
+ Whether to compute and return the variance-covariance matrix.
51
+ Set to False for faster computation when SEs are not needed.
52
+ return_fitted : bool, default False
53
+ Whether to return fitted values in addition to residuals.
54
+ check_finite : bool, default True
55
+ Whether to check that X and y contain only finite values (no NaN/Inf).
56
+ Set to False for faster computation if you are certain your data is clean.
57
+
58
+ Returns
59
+ -------
60
+ coefficients : ndarray of shape (k,)
61
+ OLS coefficient estimates.
62
+ residuals : ndarray of shape (n,)
63
+ Residuals (y - X @ coefficients).
64
+ fitted : ndarray of shape (n,), optional
65
+ Fitted values (X @ coefficients). Only returned if return_fitted=True.
66
+ vcov : ndarray of shape (k, k) or None
67
+ Variance-covariance matrix (HC1 or cluster-robust).
68
+ None if return_vcov=False.
69
+
70
+ Notes
71
+ -----
72
+ This function uses scipy.linalg.lstsq with the 'gelsy' driver, which is
73
+ QR-based and typically faster than NumPy's default SVD-based solver for
74
+ well-conditioned matrices.
75
+
76
+ The cluster-robust standard errors use the sandwich estimator with the
77
+ standard small-sample adjustment: (G/(G-1)) * ((n-1)/(n-k)).
78
+
79
+ Examples
80
+ --------
81
+ >>> import numpy as np
82
+ >>> from diff_diff.linalg import solve_ols
83
+ >>> X = np.column_stack([np.ones(100), np.random.randn(100)])
84
+ >>> y = 2 + 3 * X[:, 1] + np.random.randn(100)
85
+ >>> coef, resid, vcov = solve_ols(X, y)
86
+ >>> print(f"Intercept: {coef[0]:.2f}, Slope: {coef[1]:.2f}")
87
+ """
88
+ # Validate inputs
89
+ X = np.asarray(X, dtype=np.float64)
90
+ y = np.asarray(y, dtype=np.float64)
91
+
92
+ if X.ndim != 2:
93
+ raise ValueError(f"X must be 2-dimensional, got shape {X.shape}")
94
+ if y.ndim != 1:
95
+ raise ValueError(f"y must be 1-dimensional, got shape {y.shape}")
96
+ if X.shape[0] != y.shape[0]:
97
+ raise ValueError(
98
+ f"X and y must have same number of observations: "
99
+ f"{X.shape[0]} vs {y.shape[0]}"
100
+ )
101
+
102
+ n, k = X.shape
103
+ if n < k:
104
+ raise ValueError(
105
+ f"Fewer observations ({n}) than parameters ({k}). "
106
+ "Cannot solve underdetermined system."
107
+ )
108
+
109
+ # Check for NaN/Inf values if requested
110
+ if check_finite:
111
+ if not np.isfinite(X).all():
112
+ raise ValueError(
113
+ "X contains NaN or Inf values. "
114
+ "Clean your data or set check_finite=False to skip this check."
115
+ )
116
+ if not np.isfinite(y).all():
117
+ raise ValueError(
118
+ "y contains NaN or Inf values. "
119
+ "Clean your data or set check_finite=False to skip this check."
120
+ )
121
+
122
+ # Solve OLS using scipy's optimized solver
123
+ # 'gelsy' uses QR with column pivoting, faster than default 'gelsd' (SVD)
124
+ # Note: gelsy doesn't reliably report rank, so we don't check for deficiency
125
+ coefficients = scipy_lstsq(X, y, lapack_driver="gelsy", check_finite=False)[0]
126
+
127
+ # Compute residuals and fitted values
128
+ fitted = X @ coefficients
129
+ residuals = y - fitted
130
+
131
+ # Compute variance-covariance matrix if requested
132
+ vcov = None
133
+ if return_vcov:
134
+ vcov = compute_robust_vcov(X, residuals, cluster_ids)
135
+
136
+ if return_fitted:
137
+ return coefficients, residuals, fitted, vcov
138
+ else:
139
+ return coefficients, residuals, vcov
140
+
141
+
142
+ def compute_robust_vcov(
143
+ X: np.ndarray,
144
+ residuals: np.ndarray,
145
+ cluster_ids: Optional[np.ndarray] = None,
146
+ ) -> np.ndarray:
147
+ """
148
+ Compute heteroskedasticity-robust or cluster-robust variance-covariance matrix.
149
+
150
+ Uses the sandwich estimator: (X'X)^{-1} * meat * (X'X)^{-1}
151
+
152
+ Parameters
153
+ ----------
154
+ X : ndarray of shape (n, k)
155
+ Design matrix.
156
+ residuals : ndarray of shape (n,)
157
+ OLS residuals.
158
+ cluster_ids : ndarray of shape (n,), optional
159
+ Cluster identifiers. If None, computes HC1 robust SEs.
160
+
161
+ Returns
162
+ -------
163
+ vcov : ndarray of shape (k, k)
164
+ Variance-covariance matrix.
165
+
166
+ Notes
167
+ -----
168
+ For HC1 (no clustering):
169
+ meat = X' * diag(u^2) * X
170
+ adjustment = n / (n - k)
171
+
172
+ For cluster-robust:
173
+ meat = sum_g (X_g' u_g)(X_g' u_g)'
174
+ adjustment = (G / (G-1)) * ((n-1) / (n-k))
175
+
176
+ The cluster-robust computation is vectorized using pandas groupby,
177
+ which is much faster than a Python loop over clusters.
178
+ """
179
+ n, k = X.shape
180
+ XtX = X.T @ X
181
+
182
+ if cluster_ids is None:
183
+ # HC1 (heteroskedasticity-robust) standard errors
184
+ adjustment = n / (n - k)
185
+ u_squared = residuals**2
186
+ # Vectorized meat computation: X' diag(u^2) X = (X * u^2)' X
187
+ meat = X.T @ (X * u_squared[:, np.newaxis])
188
+ else:
189
+ # Cluster-robust standard errors (vectorized via groupby)
190
+ cluster_ids = np.asarray(cluster_ids)
191
+ unique_clusters = np.unique(cluster_ids)
192
+ n_clusters = len(unique_clusters)
193
+
194
+ if n_clusters < 2:
195
+ raise ValueError(
196
+ f"Need at least 2 clusters for cluster-robust SEs, got {n_clusters}"
197
+ )
198
+
199
+ # Small-sample adjustment
200
+ adjustment = (n_clusters / (n_clusters - 1)) * ((n - 1) / (n - k))
201
+
202
+ # Compute cluster-level scores: sum of X_i * u_i within each cluster
203
+ # scores[i] = X[i] * residuals[i] for each observation
204
+ scores = X * residuals[:, np.newaxis] # (n, k)
205
+
206
+ # Sum scores within each cluster using pandas groupby (vectorized)
207
+ # This is much faster than looping over clusters
208
+ cluster_scores = pd.DataFrame(scores).groupby(cluster_ids).sum().values # (G, k)
209
+
210
+ # Meat is the outer product sum: sum_g (score_g)(score_g)'
211
+ # Equivalent to cluster_scores.T @ cluster_scores
212
+ meat = cluster_scores.T @ cluster_scores # (k, k)
213
+
214
+ # Sandwich estimator: (X'X)^{-1} meat (X'X)^{-1}
215
+ # Solve (X'X) temp = meat, then solve (X'X) vcov' = temp'
216
+ # More stable than explicit inverse
217
+ try:
218
+ temp = np.linalg.solve(XtX, meat)
219
+ vcov = adjustment * np.linalg.solve(XtX, temp.T).T
220
+ except np.linalg.LinAlgError as e:
221
+ if "Singular" in str(e):
222
+ raise ValueError(
223
+ "Design matrix is rank-deficient (singular X'X matrix). "
224
+ "This indicates perfect multicollinearity. Check your fixed effects "
225
+ "and covariates for linear dependencies."
226
+ ) from e
227
+ raise
228
+
229
+ return vcov
230
+
231
+
232
+ def compute_r_squared(
233
+ y: np.ndarray,
234
+ residuals: np.ndarray,
235
+ adjusted: bool = False,
236
+ n_params: int = 0,
237
+ ) -> float:
238
+ """
239
+ Compute R-squared or adjusted R-squared.
240
+
241
+ Parameters
242
+ ----------
243
+ y : ndarray of shape (n,)
244
+ Response vector.
245
+ residuals : ndarray of shape (n,)
246
+ OLS residuals.
247
+ adjusted : bool, default False
248
+ If True, compute adjusted R-squared.
249
+ n_params : int, default 0
250
+ Number of parameters (including intercept). Required if adjusted=True.
251
+
252
+ Returns
253
+ -------
254
+ r_squared : float
255
+ R-squared or adjusted R-squared.
256
+ """
257
+ ss_res = np.sum(residuals**2)
258
+ ss_tot = np.sum((y - np.mean(y)) ** 2)
259
+
260
+ if ss_tot == 0:
261
+ return 0.0
262
+
263
+ r_squared = 1 - (ss_res / ss_tot)
264
+
265
+ if adjusted:
266
+ n = len(y)
267
+ if n <= n_params:
268
+ return r_squared
269
+ r_squared = 1 - (1 - r_squared) * (n - 1) / (n - n_params)
270
+
271
+ return r_squared
diff_diff/staggered.py CHANGED
@@ -13,12 +13,16 @@ import numpy as np
13
13
  import pandas as pd
14
14
  from scipy import optimize
15
15
 
16
+ from diff_diff.linalg import solve_ols
16
17
  from diff_diff.results import _get_significance_stars
17
18
  from diff_diff.utils import (
18
19
  compute_confidence_interval,
19
20
  compute_p_value,
20
21
  )
21
22
 
23
+ # Type alias for pre-computed structures
24
+ PrecomputedData = Dict[str, Any]
25
+
22
26
  # =============================================================================
23
27
  # Bootstrap Weight Generators
24
28
  # =============================================================================
@@ -75,6 +79,59 @@ def _generate_bootstrap_weights(
75
79
  )
76
80
 
77
81
 
82
+ def _generate_bootstrap_weights_batch(
83
+ n_bootstrap: int,
84
+ n_units: int,
85
+ weight_type: str,
86
+ rng: np.random.Generator,
87
+ ) -> np.ndarray:
88
+ """
89
+ Generate all bootstrap weights at once (vectorized).
90
+
91
+ Parameters
92
+ ----------
93
+ n_bootstrap : int
94
+ Number of bootstrap iterations.
95
+ n_units : int
96
+ Number of units (clusters) to generate weights for.
97
+ weight_type : str
98
+ Type of weights: "rademacher", "mammen", or "webb".
99
+ rng : np.random.Generator
100
+ Random number generator.
101
+
102
+ Returns
103
+ -------
104
+ np.ndarray
105
+ Array of bootstrap weights with shape (n_bootstrap, n_units).
106
+ """
107
+ if weight_type == "rademacher":
108
+ # Rademacher: +1 or -1 with equal probability
109
+ return rng.choice([-1.0, 1.0], size=(n_bootstrap, n_units))
110
+
111
+ elif weight_type == "mammen":
112
+ # Mammen's two-point distribution
113
+ sqrt5 = np.sqrt(5)
114
+ val1 = -(sqrt5 - 1) / 2
115
+ val2 = (sqrt5 + 1) / 2
116
+ p1 = (sqrt5 + 1) / (2 * sqrt5)
117
+ return rng.choice([val1, val2], size=(n_bootstrap, n_units), p=[p1, 1 - p1])
118
+
119
+ elif weight_type == "webb":
120
+ # Webb's 6-point distribution
121
+ values = np.array([
122
+ -np.sqrt(3 / 2), -np.sqrt(2 / 2), -np.sqrt(1 / 2),
123
+ np.sqrt(1 / 2), np.sqrt(2 / 2), np.sqrt(3 / 2)
124
+ ])
125
+ probs = np.array([1, 2, 3, 3, 2, 1]) / 12
126
+ return rng.choice(values, size=(n_bootstrap, n_units), p=probs)
127
+
128
+ else:
129
+ raise ValueError(
130
+ f"weight_type must be 'rademacher', 'mammen', or 'webb', "
131
+ f"got '{weight_type}'"
132
+ )
133
+
134
+
78
135
  # =============================================================================
79
136
  # Bootstrap Results Container
80
137
  # =============================================================================
@@ -226,15 +283,8 @@ def _linear_regression(
226
283
  # Add intercept
227
284
  X_with_intercept = np.column_stack([np.ones(n), X])
228
285
 
229
- # OLS: beta = (X'X)^{-1} X'y
230
- try:
231
- beta = np.linalg.lstsq(X_with_intercept, y, rcond=None)[0]
232
- except np.linalg.LinAlgError:
233
- # Fallback: use pseudo-inverse
234
- beta = np.linalg.pinv(X_with_intercept) @ y
235
-
236
- fitted = X_with_intercept @ beta
237
- residuals = y - fitted
286
+ # Use unified OLS backend (no vcov needed)
287
+ beta, residuals, _ = solve_ols(X_with_intercept, y, return_vcov=False)
238
288
 
239
289
  return beta, residuals
240
290
 
@@ -553,6 +603,14 @@ class CallawaySantAnna:
553
603
  Number of bootstrap iterations for inference.
554
604
  If 0, uses analytical standard errors.
555
605
  Recommended: 999 or more for reliable inference.
606
+
607
+ .. note:: Memory Usage
608
+ The bootstrap stores all weights in memory as a (n_bootstrap, n_units)
609
+ float64 array. For large datasets, this can be significant:
610
+ - 1K bootstrap × 10K units = ~80 MB
611
+ - 10K bootstrap × 100K units = ~8 GB
612
+ Consider reducing n_bootstrap if memory is constrained.
613
+
556
614
  bootstrap_weights : str, default="rademacher"
557
615
  Type of weights for multiplier bootstrap:
558
616
  - "rademacher": +1/-1 with equal probability (standard choice)
@@ -694,6 +752,199 @@ class CallawaySantAnna:
694
752
  self.is_fitted_ = False
695
753
  self.results_ = None
696
754
 
755
+ def _precompute_structures(
756
+ self,
757
+ df: pd.DataFrame,
758
+ outcome: str,
759
+ unit: str,
760
+ time: str,
761
+ first_treat: str,
762
+ covariates: Optional[List[str]],
763
+ time_periods: List[Any],
764
+ treatment_groups: List[Any],
765
+ ) -> PrecomputedData:
766
+ """
767
+ Pre-compute data structures for efficient ATT(g,t) computation.
768
+
769
+ This pivots data to wide format and pre-computes:
770
+ - Outcome matrix (units x time periods)
771
+ - Covariate matrix (units x covariates) from base period
772
+ - Unit cohort membership masks
773
+ - Control unit masks
774
+
775
+ Returns
776
+ -------
777
+ PrecomputedData
778
+ Dictionary with pre-computed structures.
779
+ """
780
+ # Get unique units and their cohort assignments
781
+ unit_info = df.groupby(unit)[first_treat].first()
782
+ all_units = unit_info.index.values
783
+ unit_cohorts = unit_info.values
784
+ n_units = len(all_units)
785
+
786
+ # Create unit index mapping for fast lookups
787
+ unit_to_idx = {u: i for i, u in enumerate(all_units)}
788
+
789
+ # Pivot outcome to wide format: rows = units, columns = time periods
790
+ outcome_wide = df.pivot(index=unit, columns=time, values=outcome)
791
+ # Reindex to ensure all units are present (handles unbalanced panels)
792
+ outcome_wide = outcome_wide.reindex(all_units)
793
+ outcome_matrix = outcome_wide.values # Shape: (n_units, n_periods)
794
+ period_to_col = {t: i for i, t in enumerate(outcome_wide.columns)}
795
+
796
+ # Pre-compute cohort masks (boolean arrays)
797
+ cohort_masks = {}
798
+ for g in treatment_groups:
799
+ cohort_masks[g] = (unit_cohorts == g)
800
+
801
+ # Never-treated mask
802
+ never_treated_mask = (unit_cohorts == 0) | (unit_cohorts == np.inf)
803
+
804
+ # Pre-compute covariate matrices by time period if needed
805
+ # (covariates are retrieved from the base period of each comparison)
806
+ covariate_by_period = None
807
+ if covariates:
808
+ covariate_by_period = {}
809
+ for t in time_periods:
810
+ period_data = df[df[time] == t].set_index(unit)
811
+ period_cov = period_data.reindex(all_units)[covariates]
812
+ covariate_by_period[t] = period_cov.values # Shape: (n_units, n_covariates)
813
+
814
+ return {
815
+ 'all_units': all_units,
816
+ 'unit_to_idx': unit_to_idx,
817
+ 'unit_cohorts': unit_cohorts,
818
+ 'outcome_matrix': outcome_matrix,
819
+ 'period_to_col': period_to_col,
820
+ 'cohort_masks': cohort_masks,
821
+ 'never_treated_mask': never_treated_mask,
822
+ 'covariate_by_period': covariate_by_period,
823
+ 'time_periods': time_periods,
824
+ }
825
+
826
+ def _compute_att_gt_fast(
827
+ self,
828
+ precomputed: PrecomputedData,
829
+ g: Any,
830
+ t: Any,
831
+ covariates: Optional[List[str]],
832
+ ) -> Tuple[Optional[float], float, int, int, Optional[Dict[str, Any]]]:
833
+ """
834
+ Compute ATT(g,t) using pre-computed data structures (fast version).
835
+
836
+ Uses vectorized numpy operations on pre-pivoted outcome matrix
837
+ instead of repeated pandas filtering.
838
+ """
839
+ time_periods = precomputed['time_periods']
840
+ period_to_col = precomputed['period_to_col']
841
+ outcome_matrix = precomputed['outcome_matrix']
842
+ cohort_masks = precomputed['cohort_masks']
843
+ never_treated_mask = precomputed['never_treated_mask']
844
+ unit_cohorts = precomputed['unit_cohorts']
845
+ all_units = precomputed['all_units']
846
+ covariate_by_period = precomputed['covariate_by_period']
847
+
848
+ # Base period for comparison
849
+ base_period = g - 1 - self.anticipation
850
+ if base_period not in period_to_col:
851
+ # Find closest earlier period
852
+ earlier = [p for p in time_periods if p < g - self.anticipation]
853
+ if not earlier:
854
+ return None, 0.0, 0, 0, None
855
+ base_period = max(earlier)
856
+
857
+ # Check if periods exist in the data
858
+ if base_period not in period_to_col or t not in period_to_col:
859
+ return None, 0.0, 0, 0, None
860
+
861
+ base_col = period_to_col[base_period]
862
+ post_col = period_to_col[t]
863
+
864
+ # Get treated units mask (cohort g)
865
+ treated_mask = cohort_masks[g]
866
+
867
+ # Get control units mask
868
+ if self.control_group == "never_treated":
869
+ control_mask = never_treated_mask
870
+ else: # not_yet_treated
871
+ # Not yet treated at time t: never-treated OR first_treat > t
872
+ control_mask = never_treated_mask | (unit_cohorts > t)
873
+
874
+ # Extract outcomes for base and post periods
875
+ y_base = outcome_matrix[:, base_col]
876
+ y_post = outcome_matrix[:, post_col]
877
+
878
+ # Compute outcome changes (vectorized)
879
+ outcome_change = y_post - y_base
880
+
881
+ # Filter to units with valid data (no NaN in either period)
882
+ valid_mask = ~(np.isnan(y_base) | np.isnan(y_post))
883
+
884
+ # Get treated and control with valid data
885
+ treated_valid = treated_mask & valid_mask
886
+ control_valid = control_mask & valid_mask
887
+
888
+ n_treated = np.sum(treated_valid)
889
+ n_control = np.sum(control_valid)
890
+
891
+ if n_treated == 0 or n_control == 0:
892
+ return None, 0.0, 0, 0, None
893
+
894
+ # Extract outcome changes for treated and control
895
+ treated_change = outcome_change[treated_valid]
896
+ control_change = outcome_change[control_valid]
897
+
898
+ # Get unit IDs for influence function
899
+ treated_units = all_units[treated_valid]
900
+ control_units = all_units[control_valid]
901
+
902
+ # Get covariates if specified (from the base period)
903
+ X_treated = None
904
+ X_control = None
905
+ if covariates and covariate_by_period is not None:
906
+ cov_matrix = covariate_by_period[base_period]
907
+ X_treated = cov_matrix[treated_valid]
908
+ X_control = cov_matrix[control_valid]
909
+
910
+ # Check for missing values
911
+ if np.any(np.isnan(X_treated)) or np.any(np.isnan(X_control)):
912
+ warnings.warn(
913
+ f"Missing values in covariates for group {g}, time {t}. "
914
+ "Falling back to unconditional estimation.",
915
+ UserWarning,
916
+ stacklevel=3,
917
+ )
918
+ X_treated = None
919
+ X_control = None
920
+
921
+ # Estimation method
922
+ if self.estimation_method == "reg":
923
+ att_gt, se_gt, inf_func = self._outcome_regression(
924
+ treated_change, control_change, X_treated, X_control
925
+ )
926
+ elif self.estimation_method == "ipw":
927
+ att_gt, se_gt, inf_func = self._ipw_estimation(
928
+ treated_change, control_change,
929
+ int(n_treated), int(n_control),
930
+ X_treated, X_control
931
+ )
932
+ else: # doubly robust
933
+ att_gt, se_gt, inf_func = self._doubly_robust(
934
+ treated_change, control_change, X_treated, X_control
935
+ )
936
+
937
+ # Package influence function info with unit IDs for bootstrap
938
+ n_t = int(n_treated)
939
+ inf_func_info = {
940
+ 'treated_units': list(treated_units),
941
+ 'control_units': list(control_units),
942
+ 'treated_inf': inf_func[:n_t],
943
+ 'control_inf': inf_func[n_t:],
944
+ }
945
+
946
+ return att_gt, se_gt, int(n_treated), int(n_control), inf_func_info
947
+
697
948
  def fit(
698
949
  self,
699
950
  data: pd.DataFrame,
@@ -779,6 +1030,12 @@ class CallawaySantAnna:
779
1030
  if n_control_units == 0:
780
1031
  raise ValueError("No never-treated units found. Check 'first_treat' column.")
781
1032
 
1033
+ # Pre-compute data structures for efficient ATT(g,t) computation
1034
+ precomputed = self._precompute_structures(
1035
+ df, outcome, unit, time, first_treat,
1036
+ covariates, time_periods, treatment_groups
1037
+ )
1038
+
782
1039
  # Compute ATT(g,t) for each group-time combination
783
1040
  group_time_effects = {}
784
1041
  influence_func_info = {} # Store influence functions for bootstrap
@@ -788,9 +1045,8 @@ class CallawaySantAnna:
788
1045
  valid_periods = [t for t in time_periods if t >= g - self.anticipation]
789
1046
 
790
1047
  for t in valid_periods:
791
- att_gt, se_gt, n_treat, n_ctrl, inf_info = self._compute_att_gt(
792
- df, outcome, unit, time, first_treat, g, t,
793
- covariates, time_periods
1048
+ att_gt, se_gt, n_treat, n_ctrl, inf_info = self._compute_att_gt_fast(
1049
+ precomputed, g, t, covariates
794
1050
  )
795
1051
 
796
1052
  if att_gt is not None:
@@ -917,135 +1173,6 @@ class CallawaySantAnna:
917
1173
  self.is_fitted_ = True
918
1174
  return self.results_
919
1175
 
920
- def _compute_att_gt(
921
- self,
922
- df: pd.DataFrame,
923
- outcome: str,
924
- unit: str,
925
- time: str,
926
- first_treat: str,
927
- g: Any,
928
- t: Any,
929
- covariates: Optional[List[str]],
930
- all_periods: List[Any],
931
- ) -> Tuple[Optional[float], float, int, int, Optional[Dict[str, Any]]]:
932
- """
933
- Compute ATT(g,t) for a specific group-time combination.
934
-
935
- Uses 2x2 DiD comparing:
936
- - Treated: Units in cohort g
937
- - Control: Never-treated units (or not-yet-treated if specified)
938
- - Pre-period: g - 1 (or earlier if anticipation > 0)
939
- - Post-period: t
940
- """
941
- # Base period for comparison
942
- base_period = g - 1 - self.anticipation
943
- if base_period not in all_periods:
944
- # Find closest earlier period
945
- earlier = [p for p in all_periods if p < g - self.anticipation]
946
- if not earlier:
947
- return None, 0.0, 0, 0, None
948
- base_period = max(earlier)
949
-
950
- # Treated group: units first treated in period g
951
- treated_units = df[df[first_treat] == g][unit].unique()
952
-
953
- # Control group
954
- if self.control_group == "never_treated":
955
- control_mask = df['_never_treated']
956
- else: # not_yet_treated
957
- # Not yet treated at time t
958
- control_mask = (df['_never_treated']) | (df[first_treat] > t)
959
-
960
- control_units = df[control_mask][unit].unique()
961
-
962
- if len(treated_units) == 0 or len(control_units) == 0:
963
- return None, 0.0, 0, 0, None
964
-
965
- # Get data for the two periods
966
- df_base = df[df[time] == base_period].set_index(unit)
967
- df_post = df[df[time] == t].set_index(unit)
968
-
969
- # Compute outcome changes for treated
970
- treated_base = df_base.loc[df_base.index.isin(treated_units), outcome]
971
- treated_post = df_post.loc[df_post.index.isin(treated_units), outcome]
972
- treated_common = treated_base.index.intersection(treated_post.index)
973
-
974
- if len(treated_common) == 0:
975
- return None, 0.0, 0, 0, None
976
-
977
- treated_change = (
978
- treated_post.loc[treated_common].values -
979
- treated_base.loc[treated_common].values
980
- )
981
-
982
- # Compute outcome changes for control
983
- control_base = df_base.loc[df_base.index.isin(control_units), outcome]
984
- control_post = df_post.loc[df_post.index.isin(control_units), outcome]
985
- control_common = control_base.index.intersection(control_post.index)
986
-
987
- if len(control_common) == 0:
988
- return None, 0.0, 0, 0, None
989
-
990
- control_change = (
991
- control_post.loc[control_common].values -
992
- control_base.loc[control_common].values
993
- )
994
-
995
- # Get covariates if specified (use base period values for conditioning)
996
- X_treated = None
997
- X_control = None
998
- if covariates:
999
- try:
1000
- X_treated = df_base.loc[treated_common, covariates].values
1001
- X_control = df_base.loc[control_common, covariates].values
1002
- # Check for missing values and handle them
1003
- if np.any(np.isnan(X_treated)) or np.any(np.isnan(X_control)):
1004
- warnings.warn(
1005
- f"Missing values in covariates for group {g}, time {t}. "
1006
- "Falling back to unconditional estimation.",
1007
- UserWarning,
1008
- stacklevel=3,
1009
- )
1010
- X_treated = None
1011
- X_control = None
1012
- except KeyError:
1013
- warnings.warn(
1014
- f"Could not extract covariates for group {g}, time {t}. "
1015
- "Falling back to unconditional estimation.",
1016
- UserWarning,
1017
- stacklevel=3,
1018
- )
1019
- X_treated = None
1020
- X_control = None
1021
-
1022
- # Estimation method
1023
- if self.estimation_method == "reg":
1024
- att_gt, se_gt, inf_func = self._outcome_regression(
1025
- treated_change, control_change, X_treated, X_control
1026
- )
1027
- elif self.estimation_method == "ipw":
1028
- att_gt, se_gt, inf_func = self._ipw_estimation(
1029
- treated_change, control_change,
1030
- len(treated_common), len(control_common),
1031
- X_treated, X_control
1032
- )
1033
- else: # doubly robust
1034
- att_gt, se_gt, inf_func = self._doubly_robust(
1035
- treated_change, control_change, X_treated, X_control
1036
- )
1037
-
1038
- # Package influence function info with unit IDs for bootstrap
1039
- n_t = len(treated_common)
1040
- inf_func_info = {
1041
- 'treated_units': list(treated_common),
1042
- 'control_units': list(control_common),
1043
- 'treated_inf': inf_func[:n_t],
1044
- 'control_inf': inf_func[n_t:],
1045
- }
1046
-
1047
- return att_gt, se_gt, len(treated_common), len(control_common), inf_func_info
1048
-
1049
1176
  def _outcome_regression(
1050
1177
  self,
1051
1178
  treated_change: np.ndarray,
@@ -1539,70 +1666,80 @@ class CallawaySantAnna:
1539
1666
  gt_pairs, group_time_effects, treatment_groups
1540
1667
  )
1541
1668
 
1542
- # Bootstrap arrays to store results
1669
+ # Pre-compute unit index arrays for each (g,t) pair (done once, not per iteration)
1670
+ gt_treated_indices = []
1671
+ gt_control_indices = []
1672
+ gt_treated_inf = []
1673
+ gt_control_inf = []
1674
+
1675
+ for j, gt in enumerate(gt_pairs):
1676
+ info = influence_func_info[gt]
1677
+ treated_idx = np.array([unit_to_idx[u] for u in info['treated_units']])
1678
+ control_idx = np.array([unit_to_idx[u] for u in info['control_units']])
1679
+ gt_treated_indices.append(treated_idx)
1680
+ gt_control_indices.append(control_idx)
1681
+ gt_treated_inf.append(np.asarray(info['treated_inf']))
1682
+ gt_control_inf.append(np.asarray(info['control_inf']))
1683
+
1684
+ # Generate ALL bootstrap weights upfront: shape (n_bootstrap, n_units)
1685
+ # This is much faster than generating one at a time
1686
+ all_bootstrap_weights = _generate_bootstrap_weights_batch(
1687
+ self.n_bootstrap, n_units, self.bootstrap_weight_type, rng
1688
+ )
1689
+
1690
+ # Vectorized bootstrap ATT(g,t) computation
1691
+ # Compute all bootstrap ATTs for all (g,t) pairs using matrix operations
1543
1692
  bootstrap_atts_gt = np.zeros((self.n_bootstrap, n_gt))
1544
- bootstrap_overall = np.zeros(self.n_bootstrap)
1545
1693
 
1694
+ for j in range(n_gt):
1695
+ treated_idx = gt_treated_indices[j]
1696
+ control_idx = gt_control_indices[j]
1697
+ treated_inf = gt_treated_inf[j]
1698
+ control_inf = gt_control_inf[j]
1699
+
1700
+ # Extract weights for this (g,t)'s units across all bootstrap iterations
1701
+ # Shape: (n_bootstrap, n_treated) and (n_bootstrap, n_control)
1702
+ treated_weights = all_bootstrap_weights[:, treated_idx]
1703
+ control_weights = all_bootstrap_weights[:, control_idx]
1704
+
1705
+ # Vectorized perturbation: matrix-vector multiply
1706
+ # Shape: (n_bootstrap,)
1707
+ perturbations = (
1708
+ treated_weights @ treated_inf +
1709
+ control_weights @ control_inf
1710
+ )
1711
+
1712
+ bootstrap_atts_gt[:, j] = original_atts[j] + perturbations
1713
+
1714
+ # Vectorized overall ATT: matrix-vector multiply
1715
+ # Shape: (n_bootstrap,)
1716
+ bootstrap_overall = bootstrap_atts_gt @ overall_weights
1717
+
1718
+ # Vectorized event study aggregation
1546
1719
  if event_study_info is not None:
1547
1720
  rel_periods = sorted(event_study_info.keys())
1548
- bootstrap_event_study = {e: np.zeros(self.n_bootstrap) for e in rel_periods}
1721
+ bootstrap_event_study = {}
1722
+ for e in rel_periods:
1723
+ agg_info = event_study_info[e]
1724
+ gt_indices = agg_info['gt_indices']
1725
+ weights = agg_info['weights']
1726
+ # Vectorized: select columns and multiply by weights
1727
+ bootstrap_event_study[e] = bootstrap_atts_gt[:, gt_indices] @ weights
1549
1728
  else:
1550
1729
  bootstrap_event_study = None
1551
1730
 
1731
+ # Vectorized group aggregation
1552
1732
  if group_agg_info is not None:
1553
1733
  groups = sorted(group_agg_info.keys())
1554
- bootstrap_group = {g: np.zeros(self.n_bootstrap) for g in groups}
1734
+ bootstrap_group = {}
1735
+ for g in groups:
1736
+ agg_info = group_agg_info[g]
1737
+ gt_indices = agg_info['gt_indices']
1738
+ weights = agg_info['weights']
1739
+ bootstrap_group[g] = bootstrap_atts_gt[:, gt_indices] @ weights
1555
1740
  else:
1556
1741
  bootstrap_group = None
1557
1742
 
1558
- # Run bootstrap iterations
1559
- for b in range(self.n_bootstrap):
1560
- # Generate unit-level weights
1561
- unit_weights = _generate_bootstrap_weights(
1562
- n_units, self.bootstrap_weight_type, rng
1563
- )
1564
-
1565
- # Compute bootstrap ATT(g,t) for each group-time pair
1566
- for j, gt in enumerate(gt_pairs):
1567
- info = influence_func_info[gt]
1568
-
1569
- # Get weights for treated and control units
1570
- treated_indices = [unit_to_idx[u] for u in info['treated_units']]
1571
- control_indices = [unit_to_idx[u] for u in info['control_units']]
1572
-
1573
- treated_weights = unit_weights[treated_indices]
1574
- control_weights = unit_weights[control_indices]
1575
-
1576
- # Influence function perturbation
1577
- # Bootstrap ATT* = ATT + sum(weights * influence)
1578
- perturbation = (
1579
- np.sum(treated_weights * info['treated_inf']) +
1580
- np.sum(control_weights * info['control_inf'])
1581
- )
1582
-
1583
- bootstrap_atts_gt[b, j] = original_atts[j] + perturbation
1584
-
1585
- # Compute bootstrap overall ATT
1586
- bootstrap_overall[b] = np.sum(overall_weights * bootstrap_atts_gt[b, :])
1587
-
1588
- # Compute bootstrap event study effects
1589
- if bootstrap_event_study is not None and event_study_info is not None:
1590
- for e, agg_info in event_study_info.items():
1591
- gt_indices = agg_info['gt_indices']
1592
- weights = agg_info['weights']
1593
- bootstrap_event_study[e][b] = np.sum(
1594
- weights * bootstrap_atts_gt[b, gt_indices]
1595
- )
1596
-
1597
- # Compute bootstrap group effects
1598
- if bootstrap_group is not None and group_agg_info is not None:
1599
- for g, agg_info in group_agg_info.items():
1600
- gt_indices = agg_info['gt_indices']
1601
- weights = agg_info['weights']
1602
- bootstrap_group[g][b] = np.sum(
1603
- weights * bootstrap_atts_gt[b, gt_indices]
1604
- )
1605
-
1606
1743
  # Compute bootstrap statistics for ATT(g,t)
1607
1744
  gt_ses = {}
1608
1745
  gt_cis = {}
diff_diff/sun_abraham.py CHANGED
@@ -16,11 +16,11 @@ from typing import Any, Dict, List, Optional, Tuple
16
16
  import numpy as np
17
17
  import pandas as pd
18
18
 
19
+ from diff_diff.linalg import compute_robust_vcov
19
20
  from diff_diff.results import _get_significance_stars
20
21
  from diff_diff.utils import (
21
22
  compute_confidence_interval,
22
23
  compute_p_value,
23
- compute_robust_se,
24
24
  )
25
25
 
26
26
 
@@ -761,7 +761,7 @@ class SunAbraham:
761
761
 
762
762
  # Compute cluster-robust standard errors
763
763
  cluster_ids = df_demeaned[cluster_var].values
764
- vcov = compute_robust_se(X, residuals, cluster_ids)
764
+ vcov = compute_robust_vcov(X, residuals, cluster_ids)
765
765
 
766
766
  # Extract cohort effects and standard errors
767
767
  cohort_effects: Dict[Tuple[Any, int], float] = {}
@@ -10,6 +10,7 @@ import pandas as pd
10
10
  from numpy.linalg import LinAlgError
11
11
 
12
12
  from diff_diff.estimators import DifferenceInDifferences
13
+ from diff_diff.linalg import solve_ols
13
14
  from diff_diff.results import SyntheticDiDResults
14
15
  from diff_diff.utils import (
15
16
  compute_confidence_interval,
@@ -444,9 +445,8 @@ class SyntheticDiD(DifferenceInDifferences):
444
445
 
445
446
  y = data[outcome].values.astype(float)
446
447
 
447
- # Fit and get residuals
448
- coeffs = np.linalg.lstsq(X_full, y, rcond=None)[0]
449
- residuals = y - X_full @ coeffs
448
+ # Fit and get residuals using unified backend
449
+ coeffs, residuals, _ = solve_ols(X_full, y, return_vcov=False)
450
450
 
451
451
  # Add back the mean for interpretability
452
452
  data[outcome] = residuals + np.mean(y)
diff_diff/triple_diff.py CHANGED
@@ -37,11 +37,11 @@ import numpy as np
37
37
  import pandas as pd
38
38
  from scipy import optimize
39
39
 
40
+ from diff_diff.linalg import compute_robust_vcov, solve_ols
40
41
  from diff_diff.results import _get_significance_stars
41
42
  from diff_diff.utils import (
42
43
  compute_confidence_interval,
43
44
  compute_p_value,
44
- compute_robust_se,
45
45
  )
46
46
 
47
47
  # =============================================================================
@@ -353,15 +353,13 @@ def _linear_regression(
353
353
  n = X.shape[0]
354
354
  X_with_intercept = np.column_stack([np.ones(n), X])
355
355
 
356
- try:
357
- beta = np.linalg.lstsq(X_with_intercept, y, rcond=None)[0]
358
- except np.linalg.LinAlgError:
359
- beta = np.linalg.pinv(X_with_intercept) @ y
360
-
361
- fitted = X_with_intercept @ beta
362
- residuals = y - fitted
356
+ # Use unified OLS backend
357
+ beta, residuals, fitted, _ = solve_ols(
358
+ X_with_intercept, y, return_fitted=True, return_vcov=False
359
+ )
363
360
 
364
- ss_res = np.sum(residuals ** 2)
361
+ # Compute R-squared
362
+ ss_res = np.sum(residuals**2)
365
363
  ss_tot = np.sum((y - np.mean(y)) ** 2)
366
364
  r_squared = 1 - (ss_res / ss_tot) if ss_tot > 0 else 0.0
367
365
 
@@ -741,17 +739,13 @@ class TripleDifference:
741
739
 
742
740
  design_matrix = np.column_stack(design_cols)
743
741
 
744
- # Fit OLS
745
- try:
746
- coefficients = np.linalg.lstsq(design_matrix, y, rcond=None)[0]
747
- except np.linalg.LinAlgError:
748
- coefficients = np.linalg.pinv(design_matrix) @ y
749
-
750
- fitted = design_matrix @ coefficients
751
- residuals = y - fitted
742
+ # Fit OLS using unified backend
743
+ coefficients, residuals, fitted, _ = solve_ols(
744
+ design_matrix, y, return_fitted=True, return_vcov=False
745
+ )
752
746
 
753
747
  # R-squared
754
- ss_res = np.sum(residuals ** 2)
748
+ ss_res = np.sum(residuals**2)
755
749
  ss_tot = np.sum((y - np.mean(y)) ** 2)
756
750
  r_squared = 1 - (ss_res / ss_tot) if ss_tot > 0 else 0.0
757
751
 
@@ -1100,10 +1094,10 @@ class TripleDifference:
1100
1094
 
1101
1095
  if self.robust:
1102
1096
  # HC1 robust standard errors
1103
- vcov = compute_robust_se(X, residuals, cluster_ids=None)
1097
+ vcov = compute_robust_vcov(X, residuals, cluster_ids=None)
1104
1098
  else:
1105
1099
  # Classical OLS standard errors
1106
- mse = np.sum(residuals ** 2) / (n - k)
1100
+ mse = np.sum(residuals**2) / (n - k)
1107
1101
  try:
1108
1102
  vcov = np.linalg.solve(X.T @ X, mse * np.eye(k))
1109
1103
  except np.linalg.LinAlgError:
diff_diff/twfe.py CHANGED
@@ -12,11 +12,11 @@ if TYPE_CHECKING:
12
12
  from diff_diff.bacon import BaconDecompositionResults
13
13
 
14
14
  from diff_diff.estimators import DifferenceInDifferences
15
+ from diff_diff.linalg import compute_robust_vcov
15
16
  from diff_diff.results import DiDResults
16
17
  from diff_diff.utils import (
17
18
  compute_confidence_interval,
18
19
  compute_p_value,
19
- compute_robust_se,
20
20
  )
21
21
 
22
22
 
@@ -136,7 +136,7 @@ class TwoWayFixedEffects(DifferenceInDifferences):
136
136
  )
137
137
  else:
138
138
  # Standard cluster-robust SE
139
- vcov = compute_robust_se(X, residuals, cluster_ids)
139
+ vcov = compute_robust_vcov(X, residuals, cluster_ids)
140
140
  se = np.sqrt(vcov[att_idx, att_idx])
141
141
  t_stat = att / se
142
142
  p_value = compute_p_value(t_stat, df=df)
@@ -214,11 +214,15 @@ class TwoWayFixedEffects(DifferenceInDifferences):
214
214
  data = data.copy()
215
215
  variables = [outcome] + (covariates or [])
216
216
 
217
+ # Cache groupby objects for efficiency (avoids re-computing group indexes)
218
+ unit_grouper = data.groupby(unit, sort=False)
219
+ time_grouper = data.groupby(time, sort=False)
220
+
217
221
  for var in variables:
218
- # Unit means
219
- unit_means = data.groupby(unit)[var].transform("mean")
220
- # Time means
221
- time_means = data.groupby(time)[var].transform("mean")
222
+ # Unit means (using cached grouper)
223
+ unit_means = unit_grouper[var].transform("mean")
224
+ # Time means (using cached grouper)
225
+ time_means = time_grouper[var].transform("mean")
222
226
  # Grand mean
223
227
  grand_mean = data[var].mean()
224
228
 
diff_diff/utils.py CHANGED
@@ -10,6 +10,9 @@ import numpy as np
10
10
  import pandas as pd
11
11
  from scipy import stats
12
12
 
13
+ from diff_diff.linalg import compute_robust_vcov as _compute_robust_vcov_linalg
14
+ from diff_diff.linalg import solve_ols as _solve_ols_linalg
15
+
13
16
  # Numerical constants for optimization algorithms
14
17
  _OPTIMIZATION_MAX_ITER = 1000 # Maximum iterations for weight optimization
15
18
  _OPTIMIZATION_TOL = 1e-8 # Convergence tolerance for optimization
@@ -48,6 +51,9 @@ def compute_robust_se(
48
51
  """
49
52
  Compute heteroskedasticity-robust (HC1) or cluster-robust standard errors.
50
53
 
54
+ This function is a thin wrapper around the optimized implementation in
55
+ diff_diff.linalg for backwards compatibility.
56
+
51
57
  Parameters
52
58
  ----------
53
59
  X : np.ndarray
@@ -62,47 +68,7 @@ def compute_robust_se(
62
68
  np.ndarray
63
69
  Variance-covariance matrix of shape (k, k).
64
70
  """
65
- n, k = X.shape
66
- XtX = X.T @ X
67
-
68
- if cluster_ids is None:
69
- # HC1 robust standard errors
70
- # HC1 adjustment factor: n / (n - k)
71
- adjustment = n / (n - k)
72
-
73
- # Create diagonal matrix with squared residuals
74
- u_squared = residuals ** 2
75
-
76
- # Meat of the sandwich: X' * diag(u^2) * X
77
- meat = X.T @ (X * u_squared[:, np.newaxis])
78
-
79
- # Compute XtX^{-1} @ meat @ XtX^{-1} using solve() for numerical stability
80
- # First solve XtX @ temp = meat to get temp = XtX^{-1} @ meat
81
- temp = np.linalg.solve(XtX, meat)
82
- # Then solve XtX @ vcov = temp.T and transpose to get XtX^{-1} @ meat @ XtX^{-1}
83
- vcov = adjustment * np.linalg.solve(XtX, temp.T).T
84
- else:
85
- # Cluster-robust standard errors
86
- unique_clusters = np.unique(cluster_ids)
87
- n_clusters = len(unique_clusters)
88
-
89
- # Adjustment factor for cluster-robust SEs
90
- adjustment = (n_clusters / (n_clusters - 1)) * ((n - 1) / (n - k))
91
-
92
- # Compute the meat of the sandwich
93
- meat = np.zeros((k, k))
94
- for cluster in unique_clusters:
95
- mask = cluster_ids == cluster
96
- X_c = X[mask]
97
- u_c = residuals[mask]
98
- score_c = X_c.T @ u_c
99
- meat += np.outer(score_c, score_c)
100
-
101
- # Compute XtX^{-1} @ meat @ XtX^{-1} using solve() for numerical stability
102
- temp = np.linalg.solve(XtX, meat)
103
- vcov = adjustment * np.linalg.solve(XtX, temp.T).T
104
-
105
- return np.asarray(vcov)
71
+ return _compute_robust_vcov_linalg(X, residuals, cluster_ids)
106
72
 
107
73
 
108
74
  def compute_confidence_interval(
@@ -461,11 +427,10 @@ def wild_bootstrap_se(
461
427
  n = X.shape[0]
462
428
 
463
429
  # Step 1: Compute original coefficient and cluster-robust SE
464
- beta_hat = np.linalg.lstsq(X, y, rcond=None)[0]
430
+ beta_hat, _, vcov_original = _solve_ols_linalg(
431
+ X, y, cluster_ids=cluster_ids, return_vcov=True
432
+ )
465
433
  original_coef = beta_hat[coefficient_index]
466
-
467
- # Compute cluster-robust SE for original t-statistic
468
- vcov_original = compute_robust_se(X, residuals, cluster_ids)
469
434
  se_original = np.sqrt(vcov_original[coefficient_index, coefficient_index])
470
435
  t_stat_original = (original_coef - null_hypothesis) / se_original
471
436
 
@@ -477,8 +442,9 @@ def wild_bootstrap_se(
477
442
  # Fit restricted model (but we need to drop the column for the restricted coef)
478
443
  # Actually, for WCR bootstrap we keep all columns but impose the null via residuals
479
444
  # Re-estimate with the restricted dependent variable
480
- beta_restricted = np.linalg.lstsq(X, y_restricted, rcond=None)[0]
481
- residuals_restricted = y_restricted - X @ beta_restricted
445
+ beta_restricted, residuals_restricted, _ = _solve_ols_linalg(
446
+ X, y_restricted, return_vcov=False
447
+ )
482
448
 
483
449
  # Create cluster-to-observation mapping for efficiency
484
450
  cluster_map = {c: np.where(cluster_ids == c)[0] for c in unique_clusters}
@@ -500,13 +466,11 @@ def wild_bootstrap_se(
500
466
  # Construct bootstrap sample: y* = X @ beta_restricted + e_restricted * weights
501
467
  y_star = X @ beta_restricted + residuals_restricted * obs_weights
502
468
 
503
- # Estimate bootstrap coefficients
504
- beta_star = np.linalg.lstsq(X, y_star, rcond=None)[0]
469
+ # Estimate bootstrap coefficients with cluster-robust SE
470
+ beta_star, residuals_star, vcov_star = _solve_ols_linalg(
471
+ X, y_star, cluster_ids=cluster_ids, return_vcov=True
472
+ )
505
473
  bootstrap_coefs[b] = beta_star[coefficient_index]
506
-
507
- # Compute bootstrap residuals and cluster-robust SE
508
- residuals_star = y_star - X @ beta_star
509
- vcov_star = compute_robust_se(X, residuals_star, cluster_ids)
510
474
  se_star = np.sqrt(vcov_star[coefficient_index, coefficient_index])
511
475
 
512
476
  # Compute bootstrap t-statistic (under null hypothesis)
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: diff-diff
3
- Version: 1.3.1
3
+ Version: 1.4.0
4
4
  Summary: A library for Difference-in-Differences causal inference analysis
5
5
  Author: diff-diff contributors
6
6
  License-Expression: MIT
@@ -0,0 +1,21 @@
1
+ diff_diff/__init__.py,sha256=cnvJNAfRou91h1XguGc2n3PAUmsgcDZibzE2mncrzoA,4361
2
+ diff_diff/bacon.py,sha256=AgQOtGUN-gtMh8q6KsGmOiVF4uCCFG2IexbeKS0mBlQ,36819
3
+ diff_diff/diagnostics.py,sha256=1yOKfauW9RP7alXYlaR8oq5tmol5T24idsfqzyAxBPQ,29572
4
+ diff_diff/estimators.py,sha256=4-em3NWsXGsPiC_FZ22Erzwknc5SDuHyAQ3mNZkL9uI,35547
5
+ diff_diff/honest_did.py,sha256=J4l-JyHHhhBiUT-F9Wls_8kjhxWhEj-N3BOe02pv0AA,46925
6
+ diff_diff/linalg.py,sha256=cn5rG9okPRR759YsSr3HxoFYgwLhrO6YPSX00Cjenos,8998
7
+ diff_diff/power.py,sha256=cpdbWG-lxqAKjeAoM_8Un8gYaDwADIeYdWUIkEBfXZg,42800
8
+ diff_diff/prep.py,sha256=bTPXTWajBzdNVPcLQ_IBYjww63DhgS3Z37V9h5tsIgc,46985
9
+ diff_diff/pretrends.py,sha256=NeYjTxC9s_iIj-3beYupiDAZROYaLjFIzbm-76XhqKk,36538
10
+ diff_diff/results.py,sha256=ymN7_WVd-XTzJGmuULO8Ryda929uBOJX0FQZ341iJek,22746
11
+ diff_diff/staggered.py,sha256=6eNSSvSS1otHvHu-CPj4_hUwI_u5xjpbXLK-xJ2qRRY,74423
12
+ diff_diff/sun_abraham.py,sha256=kUlGkenlSnWWqC0X01AaR5eR1BWrF6GLBBfpChy9pq4,41713
13
+ diff_diff/synthetic_did.py,sha256=bHGvVPjcTCyEdBd9JCksl-qBHkfsBhV5VcVvBZiPf_g,26798
14
+ diff_diff/triple_diff.py,sha256=KOF4J5hPN7SryiyMaPCsS6ZeQa9ju9o7zwgoCwchQL8,44546
15
+ diff_diff/twfe.py,sha256=VhizxT6HkuVCDmKJQXJ-9nu0AwFHMIqVxqe2Dq-5ON8,12185
16
+ diff_diff/utils.py,sha256=r_wGGy3bKzppiiFeBKkE4o3RIorssctp7NgXza1SYcM,42886
17
+ diff_diff/visualization.py,sha256=1X074fUk_fZleAlHdSGYc-ihdShY9FJ6WoKmBFCV6oI,51537
18
+ diff_diff-1.4.0.dist-info/METADATA,sha256=utVY-c8vOsM16wTdsPqcAHwsxC6oFUDEe7q-QdIurCE,77719
19
+ diff_diff-1.4.0.dist-info/WHEEL,sha256=_zCd3N1l69ArxyTb8rzEoP9TpbYXkqRFSNOD5OuxnTs,91
20
+ diff_diff-1.4.0.dist-info/top_level.txt,sha256=-7mAFgjEQIA2okDLHlh5pDwBQXnsO1Z85qvtHjZILoQ,10
21
+ diff_diff-1.4.0.dist-info/RECORD,,
@@ -1,20 +0,0 @@
1
- diff_diff/__init__.py,sha256=1Izex7iZprjS_UGqIJj6iw9fO17US6ZI1HRINYW_tsA,4361
2
- diff_diff/bacon.py,sha256=AgQOtGUN-gtMh8q6KsGmOiVF4uCCFG2IexbeKS0mBlQ,36819
3
- diff_diff/diagnostics.py,sha256=1yOKfauW9RP7alXYlaR8oq5tmol5T24idsfqzyAxBPQ,29572
4
- diff_diff/estimators.py,sha256=4Jj9xAnF39EWElweFoYUa91YnyELelT8mRDzZKZWj4E,35807
5
- diff_diff/honest_did.py,sha256=J4l-JyHHhhBiUT-F9Wls_8kjhxWhEj-N3BOe02pv0AA,46925
6
- diff_diff/power.py,sha256=cpdbWG-lxqAKjeAoM_8Un8gYaDwADIeYdWUIkEBfXZg,42800
7
- diff_diff/prep.py,sha256=bTPXTWajBzdNVPcLQ_IBYjww63DhgS3Z37V9h5tsIgc,46985
8
- diff_diff/pretrends.py,sha256=NeYjTxC9s_iIj-3beYupiDAZROYaLjFIzbm-76XhqKk,36538
9
- diff_diff/results.py,sha256=ymN7_WVd-XTzJGmuULO8Ryda929uBOJX0FQZ341iJek,22746
10
- diff_diff/staggered.py,sha256=DPx6VRlF67IryT8cXYq_gfdclQZEDHIYzQakR3OSvWA,69465
11
- diff_diff/sun_abraham.py,sha256=ux6igLILW7Y1bh3TK33P781ATUWJS4fkwrjnpTARPlU,41685
12
- diff_diff/synthetic_did.py,sha256=PjVFvCWQeE0X_4YZhw8iKSkyoyS2LEy7-yiwIQQd4bE,26765
13
- diff_diff/triple_diff.py,sha256=Jb8ld5ME6E7uQNvS2pgNmi0kHY3WikrLWm4TjTe__xI,44679
14
- diff_diff/twfe.py,sha256=iB-hvmOQd97B5_3hWvoaYGCSfcwJbXfg6towhRYD_Bw,11931
15
- diff_diff/utils.py,sha256=F0vRZ32k065VsP4YcAcUdaF_dom834QOgYj3ZaGS2pI,44211
16
- diff_diff/visualization.py,sha256=1X074fUk_fZleAlHdSGYc-ihdShY9FJ6WoKmBFCV6oI,51537
17
- diff_diff-1.3.1.dist-info/METADATA,sha256=nML3_aj-pgfeQiWsz7mOI0HRAczQc1nPPOSJkqqqpNw,77719
18
- diff_diff-1.3.1.dist-info/WHEEL,sha256=_zCd3N1l69ArxyTb8rzEoP9TpbYXkqRFSNOD5OuxnTs,91
19
- diff_diff-1.3.1.dist-info/top_level.txt,sha256=-7mAFgjEQIA2okDLHlh5pDwBQXnsO1Z85qvtHjZILoQ,10
20
- diff_diff-1.3.1.dist-info/RECORD,,