User Guide
Installation
This package is still under active development. To install it, from the Julia REPL, run:
using Pkg
Pkg.add(url = "https://github.com/masonrhayes/DoubleML.jl")Required MLJ Models
DoubleML.jl uses MLJ.jl for machine learning. Install model packages:
using Pkg
Pkg.add(["MLJ", "MLJLinearModels", "DecisionTree", "DataFrames"])What is Double Machine Learning?
Double Machine Learning (DML) estimates causal effects under the presence of (potentially) high-dimensional control variables. It solves the "regularization bias" problem by using two separate ML models to orthogonalize the problem:
- One model predicts outcome
Yfrom controls:E[Y|X] - Another model predicts treatment
Dfrom controls:E[D|X]
By taking residuals from both, we can isolate the causal effect of the treatment on the outcome.
Cross-Fitting
Cross-fitting prevents overfitting by ensuring predictions are made on data not used for training:
- Split data into K folds
- For each fold k: train on other folds, predict on fold k
- Combine predictions across all folds
This is essential for valid inference - without it, regularization bias returns.
Model Types
Partially Linear Regression (PLR)
When to use: Continuous or binary treatment, constant treatment effect
\[Y = \theta D + g_0(X) + \varepsilon \\ D = m_0(X) + v\]
Learners: ml_l for E[Y|X], ml_m for E[D|X]
model = DoubleMLPLR(data, ml_l, ml_m, n_folds=5, score=:partialling_out)
fit!(model)Score functions:
:partialling_out(default) - standard orthogonalization:IV_type- requires additionalml_glearner for endogenous treatment
Interactive Regression Model (IRM)
When to use: Binary treatment ($D \in {0,1}$), allows heterogeneous effects
\[Y = g_0(D, X) + \zeta, \space \\ \text{where} \space D \in \{0, 1\}\]
Learners: ml_g for E[Y|X,D], ml_m (classifier) for P(D=1|X)
model = DoubleMLIRM(data, ml_g, ml_m, n_folds=5, score=:ATE)
fit!(model)Estimands:
:ATE- Average Treatment Effect:ATTE- Average Treatment Effect on the Treated
Logistic Partially Linear Regression (LPLR) ⚠️ Experimental
When to use: Binary outcome ($Y \in {0,1}$), treatment effect on log-odds scale
\[E[Y|D,X] = \text{expit}(\beta_0 D + r_0(X)), \\ \text{where} \space Y \in \{0, 1\}\]
Learners:
ml_M(classifier) for P(Y=1|D,X)ml_tfor $E[\text{logit}(M)|X]$ml_mfor nuisance estimation for $E[D|X]$ml_a(optional) alterantive for $E[D|X]$
model = DoubleMLLPLR(data, ml_M, ml_t, ml_m, n_folds=5, score=:nuisance_space)
fit!(model)Score functions:
:nuisance_space(default) - fitsml_mon $Y=0$ observations only:instrument- uses weighted estimation with $M*(1-M)$ weights
Note: This model is experimental and may change in future versions.
Learner Naming Convention
| Model | Learner 1 | Learner 2 | Learner 3 |
|---|---|---|---|
| DoubleMLPLR | ml_l ($E[Y|X]$) | ml_m ($E[D|X]$) | ml_g (for :IV_type score only) |
| DoubleMLIRM | ml_g ($E[Y|X,D]$) | ml_m ($E[D|X]$) | –––––––––– |
| DoubleMLLPLR | ml_M ($P(Y=1|D,X)$) | ml_t ($E[\text{logit}(M)|X]$) | ml_m ($E[D|X]$) |
Workflow
1. Prepare Data
using DoubleML, DataFrames
# From DataFrame
df = DataFrame(y=..., d=..., x1=..., x2=...)
data = DoubleMLData(
df,
y_col=:y,
d_col=:d,
x_cols=[:x1, :x2]
)
# Or use built-in generators (set `return_type` = DataFrame to get a DataFrame)
data = make_plr_CCDDHNR2018(500, alpha=0.5)2. Set Up Learners
using MLJ
# For PLR: both can be regressors
RandomForestRegressor = @load RandomForestRegressor pkg=DecisionTree verbosity=0
ml_l = RandomForestRegressor(max_depth=10)
ml_m = RandomForestRegressor(max_depth=5)
# For IRM: ml_m must be a classifier
LogisticClassifier = @load LogisticClassifier pkg=MLJLinearModels verbosity=0
ml_g = RandomForestRegressor()
ml_m = LogisticClassifier()3. Fit Model
model = DoubleMLPLR(data, ml_l, ml_m, n_folds=5, n_rep=1)
fit!(model)4. Extract Results
summary(model) # Print summary of results
θ = coef(model)[1] # Point estimate
se = stderror(model)[1] # Standard error
ci = confint(model) # 95% CI
ct = coeftable(model) # Summary table5. Bootstrap (Optional)
bootstrap!(model, n_rep_boot=1000, method=:normal)
joint_ci = confint(model, joint=true) # Get joint confidence intervalsStatsAPI Interface
All models implement StatsAPI:
coef(model) # Treatment effect(s)
stderror(model) # Standard error(s)
confint(model) # Confidence intervals
vcov(model) # Variance-covariance matrix
nobs(model) # Number of observations
coeftable(model) # Formatted table