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
────────────────────────────────────────────────────────────────────