Tae Hyun Kim (Lowell)

Causal Forest

3분 읽기 #causal-inference#hte#causal-forest

정의

Causal Forest는 Athey, Tibshirani, Wager (2019)가 제안한 일반화 랜덤 포레스트(GRF)를 인과추론에 응용한 방법으로, 처치효과(treatment effect)의 이질성이 최대가 되도록 데이터를 분할한다.

국지적 추정값은 다음과 같다. τ^(x)=iαi(x)(Yiμ^(Xi))(Tie^(Xi))iαi(x)(Tie^(Xi))2\hat{\tau}(x) = \frac{\sum_i \alpha_i(x) (Y_i - \hat{\mu}(X_i)) (T_i - \hat{e}(X_i))}{\sum_i \alpha_i(x) (T_i - \hat{e}(X_i))^2}

여기서 각 항은 다음을 뜻한다.

  • αi(x)\alpha_i(x): 트리 앙상블에서 계산된 가중치
  • μ^(Xi)\hat{\mu}(X_i): 결과 예측 모델
  • e^(Xi)\hat{e}(X_i): 성향점수(propensity score) 추정값

직관적 이해

일반적인 랜덤 포레스트는 결과를 잘 예측하도록 데이터를 분할한다. 반면 Causal Forest는 처치효과가 서로 다른 하위 집단을 찾도록 분할한다.

가격 책정에서 Causal Forest는 “어떤 고객 특성이 가격 민감도를 결정하는가?”라는 질문에 답한다.

핵심 성질

정직성(honesty)

정직한(honest) Causal Forest는 트리 구조를 결정할 때와 리프 안에서 효과를 추정할 때 서로 다른 데이터를 쓴다.

  1. 구조 결정: 데이터의 절반으로 분할 규칙을 학습한다.
  2. 효과 추정: 나머지 절반으로 리프 내 효과를 추정한다.

이렇게 데이터를 나누어 쓰면 유효한 신뢰 구간을 얻을 수 있다.

점근적 정규성

n(τ^(x)τ(x))dN(0,V(x))\sqrt{n}(\hat{\tau}(x) - \tau(x)) \xrightarrow{d} N(0, V(x))

점근적 정규성 덕분에 신뢰 구간과 가설 검정이 유효하다.

DML과의 결합

CausalForestDML은 Double Machine Learning과 결합한 방법으로, 다음 특성을 갖는다.

  • 보조 모델(결과·처치)에 유연한 ML을 쓴다.
  • 교차 적합(cross-fitting)으로 과적합(overfitting)을 막는다.
  • 연속 처치(가격)를 지원한다.

예시

코드 예시

from econml.dml import CausalForestDML
from sklearn.ensemble import GradientBoostingRegressor

forest_dml = CausalForestDML(
    model_y=GradientBoostingRegressor(n_estimators=200),
    model_t=GradientBoostingRegressor(n_estimators=200),
    discrete_treatment=False,  # 연속 가격
    n_estimators=1000,
    min_samples_leaf=20,
    honest=True
)

forest_dml.fit(
    Y=np.log1p(data['quantity']),  # 로그 수량
    T=np.log(data['price']),        # 로그 가격 → 탄력성
    X=data[heterogeneity_vars],
    W=data[confounders]
)

# 개인별 탄력성
individual_elasticities = forest_dml.effect(data[heterogeneity_vars])
lower, upper = forest_dml.effect_interval(data[heterogeneity_vars], alpha=0.05)

print(f"평균 탄력성: {individual_elasticities.mean():.3f}")
print(f"탄력성 범위: [{individual_elasticities.min():.3f}, {individual_elasticities.max():.3f}]")

세그먼트 발견

Causal Forest는 세그먼트(segment)를 자연스럽게 발견한다.

from sklearn.cluster import KMeans

data['elasticity'] = forest_dml.effect(data[heterogeneity_vars])

kmeans = KMeans(n_clusters=4, random_state=42)
data['segment'] = kmeans.fit_predict(
    np.column_stack([data['elasticity'], data[heterogeneity_vars]])
)

segment_profile = data.groupby('segment').agg({
    'elasticity': ['mean', 'std'],
    'income': 'mean',
    'age': 'mean'
})

변수 중요도

importance = forest_dml.feature_importances_
for feat, imp in zip(heterogeneity_vars, importance):
    print(f"{feat}: {imp:.3f}")

관련 개념

참고 문헌

  • Athey, S., Tibshirani, J., & Wager, S. (2019). “Generalized Random Forests.” Annals of Statistics.
  • Wager, S., & Athey, S. (2018). “Estimation and Inference of Heterogeneous Treatment Effects using Random Forests.”
  • Comprehensive Personalized Pricing Guide, Part III, §9

연결 그래프