Tae Hyun Kim (Lowell)

T-Learner

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

정의

T-Learner(Two Learner)는 처치군(treatment group)과 대조군(control group)에 대해 별도의 모델을 따로 학습해 CATE를 추정하는 Meta-learners다.

알고리즘:

  1. 대조군에서 μ0(x)\mu_0(x)를 추정한다: μ^0(x)=E^[YX=x,W=0]\hat{\mu}_0(x) = \hat{E}[Y | X = x, W = 0]

  2. 처치군에서 μ1(x)\mu_1(x)를 추정한다: μ^1(x)=E^[YX=x,W=1]\hat{\mu}_1(x) = \hat{E}[Y | X = x, W = 1]

  3. CATE를 추정한다: τ^T(x)=μ^1(x)μ^0(x)\hat{\tau}_T(x) = \hat{\mu}_1(x) - \hat{\mu}_0(x)

직관적 이해

핵심 아이디어:

두 그룹을 완전히 분리해 각각의 반응함수(response function)를 독립적으로 학습한다.

Control data:  (X₀, Y₀) → μ̂₀(x)
Treatment data: (X₁, Y₁) → μ̂₁(x)

CATE:          τ̂(x) = μ̂₁(x) - μ̂₀(x)

장점:

  • 각 그룹 고유의 반응 구조를 포착한다
  • μ0\mu_0μ1\mu_1이 매우 다를 때 적합하다
  • 개념적으로 명확하다

단점:

  • 데이터를 공유하지 않는다(각 모델이 절반의 데이터만 쓴다)
  • CATE가 단순해도 수렴 속도(convergence rate)가 반응함수의 복잡성에 의존한다
  • 그룹 크기가 불균형하면 비효율적이다

핵심 성질

데이터 비공유

  • 각 모델이 해당 그룹의 데이터만 쓴다
  • 대조군 mm개, 처치군 nn
  • 공통 패턴은 학습할 수 없다

수렴 속도가 반응함수에 의존

Rate=O(maμ+naμ)\text{Rate} = O(m^{-a_\mu} + n^{-a_\mu})

  • aμa_\mu는 반응함수의 매끄러움(smoothness)이다
  • CATE가 단순해도(aτ>aμa_\tau > a_\mu) 수렴 속도는 aμa_\mu에 의존한다

Minimax 최적성 (Theorem 7)

특정 조건에서 T-learner는 minimax rate optimal이다.

알고리즘 상세

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

    # Step 1: Fit control model
    model_0 = base_learner.fit(X_ctrl, Y_ctrl)

    # Step 2: Fit treatment model
    model_1 = base_learner.fit(X_treat, Y_treat)

    # Step 3: Predict CATE
    def predict_cate(X_new):
        return model_1.predict(X_new) - model_0.predict(X_new)

    return predict_cate

활용

잘 맞는 상황

  • 반응함수가 매우 다를 때: μ0(x)\mu_0(x)μ1(x)\mu_1(x)의 구조가 서로 다른 경우다
  • 그룹 크기가 균형적일 때: 각 모델이 충분한 데이터를 확보한다
  • 처치효과(treatment effect)가 복잡할 때: 각 그룹의 복잡성을 따로 모델링한다

맞지 않는 상황

  • CATE는 단순하지만 반응이 복잡할 때: 수렴 속도가 불필요하게 느려진다
  • 그룹 크기가 불균형할 때: 작은 그룹의 추정이 부정확해진다
  • 공통 패턴이 많을 때: 데이터 공유의 이점을 잃는다

S-Learner와 비교

AspectT-LearnerS-Learner
Models2 (separate)1 (combined)
Data per modelmm or nnm+nm + n
StructureCaptures different responsesAssumes similar responses
RiskNo data sharingMay ignore treatment effect

예시

시뮬레이션 설정:

  • μ0(x)=sin(x)\mu_0(x) = \sin(x) (복잡)
  • μ1(x)=cos(x)\mu_1(x) = \cos(x) (복잡, 다른 패턴)
  • τ(x)=cos(x)sin(x)\tau(x) = \cos(x) - \sin(x)

T-Learner:

  • 각 반응함수를 잘 포착한다 ✓
  • CATE 추정이 양호하다

S-Learner:

  • 두 패턴의 평균을 학습한다
  • 각 그룹의 고유 패턴을 놓친다

분산 분석

CATE 추정량(estimator)의 분산(variance)은 다음과 같다: Var(τ^T(x))=Var(μ^1(x))+Var(μ^0(x))\text{Var}(\hat{\tau}_T(x)) = \text{Var}(\hat{\mu}_1(x)) + \text{Var}(\hat{\mu}_0(x))

각 모델의 분산이 독립적으로 기여하므로, 데이터를 분할한 만큼 각 분산이 커진다.

관련 개념

구현

Python (econml):

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

t_learner = TLearner(models=RandomForestRegressor())
t_learner.fit(Y, T, X=X)
cate = t_learner.effect(X_test)

R:

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

참고 문헌

  • kunzelMetalearnersEstimatingHeterogeneous2019 - T-learner 분석 및 minimax optimality

연결 그래프