Tae Hyun Kim (Lowell)

X-Learner

5분 읽기 #causal-inference#hte#meta-learner

정의

X-Learner는 대치된 처치효과(imputed treatment effect)를 활용하는 3단계 알고리즘으로, 그룹 간 불균형과 CATE의 구조적 특성을 효과적으로 이용하는 Meta-learners의 한 방법이다.

알고리즘:

Stage 1: 반응함수(response function) 추정 μ^0(x)=E^[YX=x,W=0]\hat{\mu}_0(x) = \hat{E}[Y | X = x, W = 0] μ^1(x)=E^[YX=x,W=1]\hat{\mu}_1(x) = \hat{E}[Y | X = x, W = 1]

Stage 2: 대치 처치효과 계산과 CATE 추정 D~1i:=Y1iμ^0(X1i)(treatment group)\tilde{D}_{1i} := Y_{1i} - \hat{\mu}_0(X_{1i}) \quad \text{(treatment group)} D~0i:=μ^1(X0i)Y0i(control group)\tilde{D}_{0i} := \hat{\mu}_1(X_{0i}) - Y_{0i} \quad \text{(control group)}

각 그룹에서 τ(x)\tau(x)를 추정한다.

  • τ^1(x)\hat{\tau}_1(x): 처치군의 대치 효과로 학습
  • τ^0(x)\hat{\tau}_0(x): 대조군의 대치 효과로 학습

Stage 3: 가중 결합 τ^X(x)=g(x)τ^0(x)+(1g(x))τ^1(x)\hat{\tau}_X(x) = g(x)\hat{\tau}_0(x) + (1 - g(x))\hat{\tau}_1(x)

여기서 g(x)[0,1]g(x) \in [0, 1]은 가중 함수이며, 보통 성향점수(propensity score) e^(x)\hat{e}(x)를 쓴다.

직관적 이해

핵심 아이디어:

관측된 결과와 상대 그룹의 예측값으로 “대치된” 처치효과를 만들고, 각 그룹의 관점에서 CATE를 추정한 뒤 결합한다.

Stage 1: Estimate response functions (like T-learner)
           μ̂₀(x), μ̂₁(x)

Stage 2: Impute treatment effects
  Treatment group: D̃₁ᵢ = Y₁ᵢ - μ̂₀(X₁ᵢ)  (observed - predicted control)
  Control group:   D̃₀ᵢ = μ̂₁(X₀ᵢ) - Y₀ᵢ  (predicted treatment - observed)

         Train τ̂₁(x) on D̃₁, τ̂₀(x) on D̃₀

Stage 3: Weighted combination
         τ̂(x) = g(x)·τ̂₀(x) + (1-g(x))·τ̂₁(x)

이름이 “X”인 이유

  • 처치군의 정보가 대조군의 CATE 추정에 쓰이고, 그 반대도 마찬가지다.
  • 정보가 “교차(cross)“하며 전달된다.

핵심 성질

CATE 구조의 활용

  • CATE가 단순할 때(예: 선형) 그 구조를 활용할 수 있다.
  • 반응함수가 복잡해도 CATE가 단순하면 빠른 수렴 속도(convergence rate)를 달성한다.

불균형 그룹 처리

  • 가중 함수 g(x)g(x)로 그룹 크기 불균형을 조절한다.
  • 큰 그룹의 정보를 더 많이 활용한다.

수렴 속도 (Conjecture 1)

조건이 충족되면 다음이 성립한다.

  • τ^0\hat{\tau}_0: O(maτ+naμ)O(m^{-a_\tau} + n^{-a_\mu})
  • τ^1\hat{\tau}_1: O(maμ+naτ)O(m^{-a_\mu} + n^{-a_\tau})

여기서 각 기호는 다음을 뜻한다.

  • aμa_\mu: 반응함수의 평활도(smoothness)
  • aτa_\tau: CATE 함수의 평활도
  • mm: 대조군 크기, nn: 처치군 크기

모수적 수렴 속도 (Theorem 2)

조건:

  • 선형 CATE: τ(x)=xTβ\tau(x) = x^T \beta
  • Lipschitz 반응함수
  • mc3n1/am \geq c_3 n^{1/a}

결과: g0g \equiv 0인 X-learner는 처치군 크기에 대해 모수적 수렴 속도 O(n1)O(n^{-1})을 달성한다.

알고리즘 상세

def x_learner(X, W, Y, base_learner, propensity_model=None):
    # Split data by treatment
    X_ctrl, Y_ctrl = X[W == 0], Y[W == 0]
    X_treat, Y_treat = X[W == 1], Y[W == 1]

    # Stage 1: Estimate response functions
    model_0 = base_learner.fit(X_ctrl, Y_ctrl)
    model_1 = base_learner.fit(X_treat, Y_treat)

    # Stage 2: Compute imputed treatment effects
    # For treatment group: observed - predicted control
    D_tilde_1 = Y_treat - model_0.predict(X_treat)
    # For control group: predicted treatment - observed
    D_tilde_0 = model_1.predict(X_ctrl) - Y_ctrl

    # Fit CATE models on imputed effects
    tau_model_1 = base_learner.fit(X_treat, D_tilde_1)
    tau_model_0 = base_learner.fit(X_ctrl, D_tilde_0)

    # Stage 3: Estimate propensity score for weighting
    if propensity_model is None:
        # Simple estimate or use provided
        g = lambda x: len(X_ctrl) / len(X)
    else:
        propensity_model.fit(X, W)
        g = lambda x: 1 - propensity_model.predict_proba(x)[:, 1]

    # Weighted combination
    def predict_cate(X_new):
        tau_0 = tau_model_0.predict(X_new)
        tau_1 = tau_model_1.predict(X_new)
        weights = g(X_new)
        return weights * tau_0 + (1 - weights) * tau_1

    return predict_cate

S/T-Learner와의 비교

AspectS-LearnerT-LearnerX-Learner
Models124 (2 + 2)
Data sharingFullNoneCross-group
Best whenCATE ≈ 0Different μ0,μ1\mu_0, \mu_1Imbalanced groups, smooth CATE
Rate depends onaμa_\muaμa_\muCan depend on aτa_\tau
ComplexityLowMediumHigh

활용

적합한 경우

  • 그룹 크기가 불균형할 때: 가중으로 조절할 수 있다.
  • CATE가 반응함수보다 단순할 때: aτ>aμa_\tau > a_\mu
  • 구조적 가정이 있을 때: 선형 CATE, 평활도 등

부적합한 경우

  • 두 그룹 크기가 균형적이고 단순한 경우: T-learner로 충분하다.
  • CATE가 0에 가까울 때: S-learner가 더 적합하다.
  • 계산 비용이 문제일 때: 모델 4개를 학습해야 한다.

가중 함수 선택

g(x)g(x)를 고르는 방법은 다음과 같다.

  1. 성향점수: g(x)=1e^(x)g(x) = 1 - \hat{e}(x)

    • 대조군이 많으면(e(x)e(x)가 작으면) τ^0\hat{\tau}_0에 더 의존한다.
  2. 상수: g(x)=m/(m+n)g(x) = m / (m + n)

    • 전역적 그룹 비율을 쓴다.
  3. 최적 선택: 조건에 따라 g0g \equiv 0 또는 g1g \equiv 1이 최적이다.

    • 한쪽 그룹이 압도적으로 클 때 해당한다.

반사실 결과 추정

중요: 반사실(counterfactual) 결과 추정에는 Stage 1의 결과만 쓴다.

  • Y^i(0)=μ^0(Xi)\hat{Y}_i(0) = \hat{\mu}_0(X_i) (if treated)
  • Y^i(1)=μ^1(Xi)\hat{Y}_i(1) = \hat{\mu}_1(X_i) (if control)

τ^0,τ^1\hat{\tau}_0, \hat{\tau}_1g(x)g(x)는 CATE 추정에만 쓴다.

관련 개념

  • Meta-learners - 전체 framework
  • S-Learner - 대안: 단일 모델
  • T-Learner - 대안: 별도의 두 모델
  • R-Learner - 대안: 잔차화(residualization) 회귀
  • DR-Learner - 대안: 이중 강건(doubly robust) pseudo-outcome
  • CATE - 추정 대상
  • Propensity Score - 가중 함수에 사용

구현

Python (econml):

from econml.metalearners import XLearner
from sklearn.ensemble import RandomForestRegressor

x_learner = XLearner(models=RandomForestRegressor())
x_learner.fit(Y, T, X=X)
cate = x_learner.effect(X_test)

R (causalToolbox):

library(causalToolbox)
x_rf <- X_RF(feat = X, tr = W, yobs = Y)
cate <- EstimateCate(x_rf, X_test)

참고 문헌

  • kunzelMetalearnersEstimatingHeterogeneous2019 - X-learner 제안과 이론적 분석

연결 그래프