T-Learner
정의
T-Learner(Two Learner)는 처치군(treatment group)과 대조군(control group)에 대해 별도의 모델을 따로 학습해 CATE를 추정하는 Meta-learners다.
알고리즘:
-
대조군에서 를 추정한다:
-
처치군에서 를 추정한다:
-
CATE를 추정한다:
직관적 이해
핵심 아이디어:
두 그룹을 완전히 분리해 각각의 반응함수(response function)를 독립적으로 학습한다.
Control data: (X₀, Y₀) → μ̂₀(x)
Treatment data: (X₁, Y₁) → μ̂₁(x)
↓
CATE: τ̂(x) = μ̂₁(x) - μ̂₀(x)
장점:
- 각 그룹 고유의 반응 구조를 포착한다
- 와 이 매우 다를 때 적합하다
- 개념적으로 명확하다
단점:
- 데이터를 공유하지 않는다(각 모델이 절반의 데이터만 쓴다)
- CATE가 단순해도 수렴 속도(convergence rate)가 반응함수의 복잡성에 의존한다
- 그룹 크기가 불균형하면 비효율적이다
핵심 성질
데이터 비공유
- 각 모델이 해당 그룹의 데이터만 쓴다
- 대조군 개, 처치군 개
- 공통 패턴은 학습할 수 없다
수렴 속도가 반응함수에 의존
- 는 반응함수의 매끄러움(smoothness)이다
- CATE가 단순해도() 수렴 속도는 에 의존한다
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
활용
잘 맞는 상황
- 반응함수가 매우 다를 때: 와 의 구조가 서로 다른 경우다
- 그룹 크기가 균형적일 때: 각 모델이 충분한 데이터를 확보한다
- 처치효과(treatment effect)가 복잡할 때: 각 그룹의 복잡성을 따로 모델링한다
맞지 않는 상황
- CATE는 단순하지만 반응이 복잡할 때: 수렴 속도가 불필요하게 느려진다
- 그룹 크기가 불균형할 때: 작은 그룹의 추정이 부정확해진다
- 공통 패턴이 많을 때: 데이터 공유의 이점을 잃는다
S-Learner와 비교
| Aspect | T-Learner | S-Learner |
|---|---|---|
| Models | 2 (separate) | 1 (combined) |
| Data per model | or | |
| Structure | Captures different responses | Assumes similar responses |
| Risk | No data sharing | May ignore treatment effect |
예시
시뮬레이션 설정:
- (복잡)
- (복잡, 다른 패턴)
T-Learner:
- 각 반응함수를 잘 포착한다 ✓
- CATE 추정이 양호하다
S-Learner:
- 두 패턴의 평균을 학습한다
- 각 그룹의 고유 패턴을 놓친다
분산 분석
CATE 추정량(estimator)의 분산(variance)은 다음과 같다:
각 모델의 분산이 독립적으로 기여하므로, 데이터를 분할한 만큼 각 분산이 커진다.
관련 개념
- Meta-learners - 전체 framework
- S-Learner - 대안: 단일 모델
- X-Learner - T-learner의 개선
- CATE - 추정 대상
구현
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