Interactive Regression Model (IRM) Tutorial
This tutorial demonstrates how to use the DoubleMLIRM model for estimating treatment effects with binary treatments.
Overview
The Interactive Regression Model assumes:
$$Y = g_0(D, X) + \zeta, \quad \text{where } D \in \{0, 1\}$$
Where:
\(Y\) is the outcome variable
\(D\) is a binary treatment variable (0 or 1)
\(X\) are control variables (covariates)
\(g_0(D, X)\) is the conditional mean function
IRM allows for heterogeneous treatment effects and uses doubly robust estimation.
Load packages and import ML models
begin
using DoubleML
using StableRNGs
using MLJ
using TreeParzen
using EvoTrees
end
begin
RandomForestRegressor = @load RandomForestRegressor pkg = DecisionTree verbosity = 0
EvoTreeRegressor = @load EvoTreeRegressor pkg = EvoTrees verbosity = 0
EvoTreeClassifier = @load EvoTreeClassifier pkg = EvoTrees verbosity = 0
RandomForestClassifier = @load RandomForestClassifier pkg = DecisionTree verbosity = 0
end
MLJDecisionTreeInterface.RandomForestClassifier
Generate IRM data
# IRM Data
data_irm = DoubleML.make_irm_data(1000, theta = 0.5, dim_x = 100, rng = StableRNG(42))
DoubleMLData{Float32, CategoricalArrays.CategoricalVector{Float32, UInt32, Float32, CategoricalArrays.CategoricalValue{Float32, UInt32}, Union{}}}(Float32[-0.6748344, 0.3039633, 0.6760439, 1.4083221, -0.8368299, -0.61595523, 3.0014334, -0.036194555, 1.864324, -0.28150964 … 0.8323356, -1.2143124, -1.2397234, -1.3907428, -0.14038181, 2.120363, -1.8365414, -0.623508, 1.898227, 1.4702643], CategoricalArrays.CategoricalValue{Float32, UInt32}[CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 1), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 1), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 1), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 1), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 2), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 2), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 2), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 1), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 2), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 1) … CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 2), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 1), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 1), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 1), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 2), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 2), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 2), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 1), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 1), CategoricalValue(CategoricalArrays.CategoricalPool{Float32, UInt32}([0.0f0, 1.0f0]), 1)], Float32[-0.8857411 -0.54276526 … -0.4588549 0.062866904; -0.9109771 -0.48209068 … -0.4358175 1.0236975; … ; -0.8479842 -1.4113976 … 0.59593326 0.031701036; -0.9461808 -1.1490784 … -1.3875467 -1.6245216], 1000, 100, :y, :d, [:X1, :X2, :X3, :X4, :X5, :X6, :X7, :X8, :X9, :X10 … :X91, :X92, :X93, :X94, :X95, :X96, :X97, :X98, :X99, :X100])
View what models are available for our data
begin
# Find matching models for y
models() do model
matching(model, data_irm.x, data_irm.y)
end
end
12-element Vector{NamedTuple{(:name, :package_name, :is_supervised, :abstract_type, :constructor, :deep_properties, :docstring, :fit_data_scitype, :human_name, :hyperparameter_ranges, :hyperparameter_types, :hyperparameters, :implemented_methods, :inverse_transform_scitype, :is_pure_julia, :is_wrapper, :iteration_parameter, :load_path, :package_license, :package_url, :package_uuid, :predict_scitype, :prediction_type, :reporting_operations, :reports_feature_importances, :supports_class_weights, :supports_online, :supports_training_losses, :supports_weights, :tags, :target_in_fit, :transform_scitype, :input_scitype, :target_scitype, :output_scitype)}}:
(name = CatBoostRegressor, package_name = CatBoost, ... )
(name = DecisionTreeRegressor, package_name = BetaML, ... )
(name = EvoTreeGaussian, package_name = EvoTrees, ... )
(name = EvoTreeMLE, package_name = EvoTrees, ... )
(name = EvoTreeRegressor, package_name = EvoTrees, ... )
(name = GaussianMixtureRegressor, package_name = BetaML, ... )
(name = NeuralNetworkRegressor, package_name = BetaML, ... )
(name = NeuralNetworkRegressor, package_name = MLJFlux, ... )
(name = PartLS, package_name = PartitionedLS, ... )
(name = RandomForestRegressor, package_name = BetaML, ... )
(name = SRRegressor, package_name = SymbolicRegression, ... )
(name = SRTestRegressor, package_name = SymbolicRegression, ... )
begin
# Find matching models for d
models() do model
matching(model, data_irm.x, data_irm.d)
end
end
11-element Vector{NamedTuple{(:name, :package_name, :is_supervised, :abstract_type, :constructor, :deep_properties, :docstring, :fit_data_scitype, :human_name, :hyperparameter_ranges, :hyperparameter_types, :hyperparameters, :implemented_methods, :inverse_transform_scitype, :is_pure_julia, :is_wrapper, :iteration_parameter, :load_path, :package_license, :package_url, :package_uuid, :predict_scitype, :prediction_type, :reporting_operations, :reports_feature_importances, :supports_class_weights, :supports_online, :supports_training_losses, :supports_weights, :tags, :target_in_fit, :transform_scitype, :input_scitype, :target_scitype, :output_scitype)}}:
(name = CatBoostClassifier, package_name = CatBoost, ... )
(name = DecisionTreeClassifier, package_name = BetaML, ... )
(name = EvoTreeClassifier, package_name = EvoTrees, ... )
(name = GaussianNBClassifier, package_name = NaiveBayes, ... )
(name = KernelPerceptronClassifier, package_name = BetaML, ... )
(name = NeuralNetworkBinaryClassifier, package_name = MLJFlux, ... )
(name = NeuralNetworkClassifier, package_name = BetaML, ... )
(name = NeuralNetworkClassifier, package_name = MLJFlux, ... )
(name = PegasosClassifier, package_name = BetaML, ... )
(name = PerceptronClassifier, package_name = BetaML, ... )
(name = RandomForestClassifier, package_name = BetaML, ... )
Run a simple model
begin
# Simple IRM with RandomForest
ml_g = RandomForestRegressor(rng = StableRNG(42))
ml_m = RandomForestClassifier(rng = StableRNG(42))
dml_irm_simple = DoubleML.DoubleMLIRM(data_irm, ml_g, ml_m, score = :ATE)
fit!(dml_irm_simple)
end
DoubleMLIRM{Float32, MLJDecisionTreeInterface.RandomForestRegressor, MLJDecisionTreeInterface.RandomForestClassifier}
==========================
StatsBase.CoefTable(Any[[0.8683320879936218], [0.06620793044567108], [13.115227699279785], [2.6937196905670033e-39], [0.7385669288291734], [0.9980972471580702]], ["Estimate", "Std. Error", "z value", "Pr(>|z|)", "Lower 95.0%", "Upper 95.0%"], ["d"], 4, 3)
coeftable(dml_irm_simple)
──────────────────────────────────────────────────────────────────── Estimate Std. Error z value Pr(>|z|) Lower 95.0% Upper 95.0% ──────────────────────────────────────────────────────────────────── d 0.868332 0.0662079 13.12 <1e-38 0.738567 0.998097 ────────────────────────────────────────────────────────────────────
Advanced example: self-tuning models
begin
# IRM with TreeParzen hyperparameter tuning
space = Dict(
:max_depth => HP.QuantUniform(:max_depth, 2.0, 8.0, 1.0)
)
tuned_ml_g = TunedModel(
model = EvoTreeRegressor(seed = 42),
tuning = MLJTreeParzenTuning(),
resampling = Holdout(),
range = space,
measure = MLJ.rmse,
acceleration = CPUProcesses(),
)
tuned_ml_m = TunedModel(
model = EvoTreeClassifier(seed = 42),
tuning = MLJTreeParzenTuning(),
resampling = Holdout(),
range = space,
measure = MLJ.cross_entropy,
acceleration = CPUProcesses(),
)
dml_irm = DoubleML.DoubleMLIRM(data_irm, tuned_ml_g, tuned_ml_m)
fit!(dml_irm, verbose = 0)
end
DoubleMLIRM{Float32, MLJTuning.DeterministicTunedModel{MLJTreeParzenTuning, EvoTrees.EvoTreeRegressor, Nothing}, MLJTuning.ProbabilisticTunedModel{MLJTreeParzenTuning, EvoTrees.EvoTreeClassifier, Nothing}}
==========================
StatsBase.CoefTable(Any[[1.0225872993469238], [0.2876145541667938], [3.5554087162017822], [0.0003773919595837042], [0.4588731317504624], [1.5863014669433853]], ["Estimate", "Std. Error", "z value", "Pr(>|z|)", "Lower 95.0%", "Upper 95.0%"], ["d"], 4, 3)
coeftable(dml_irm)
──────────────────────────────────────────────────────────────────── Estimate Std. Error z value Pr(>|z|) Lower 95.0% Upper 95.0% ──────────────────────────────────────────────────────────────────── d 1.02259 0.287615 3.56 0.0004 0.458873 1.5863 ────────────────────────────────────────────────────────────────────