Causal Forest
정의
Causal Forest는 Athey, Tibshirani, Wager (2019)가 제안한 일반화 랜덤 포레스트(GRF)를 인과추론에 응용한 방법으로, 처치효과(treatment effect)의 이질성이 최대가 되도록 데이터를 분할한다.
국지적 추정값은 다음과 같다.
여기서 각 항은 다음을 뜻한다.
- : 트리 앙상블에서 계산된 가중치
- : 결과 예측 모델
- : 성향점수(propensity score) 추정값
직관적 이해
일반적인 랜덤 포레스트는 결과를 잘 예측하도록 데이터를 분할한다. 반면 Causal Forest는 처치효과가 서로 다른 하위 집단을 찾도록 분할한다.
가격 책정에서 Causal Forest는 “어떤 고객 특성이 가격 민감도를 결정하는가?”라는 질문에 답한다.
핵심 성질
정직성(honesty)
정직한(honest) Causal Forest는 트리 구조를 결정할 때와 리프 안에서 효과를 추정할 때 서로 다른 데이터를 쓴다.
- 구조 결정: 데이터의 절반으로 분할 규칙을 학습한다.
- 효과 추정: 나머지 절반으로 리프 내 효과를 추정한다.
이렇게 데이터를 나누어 쓰면 유효한 신뢰 구간을 얻을 수 있다.
점근적 정규성
점근적 정규성 덕분에 신뢰 구간과 가설 검정이 유효하다.
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}")
관련 개념
- CATE - 추정 대상
- Double-Debiased ML - 이론적 기반
- Cross-fitting - 과적합 방지
- Meta-learners - 대안적 접근법
- Policy Trees - 해석 가능한 정책 학습
참고 문헌
- 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