import matplotlib
if not hasattr(matplotlib.RcParams, "_get"):
    matplotlib.RcParams._get = dict.get

Interpretability in Causal Inference#

일반적인 지도학습(supervised learning)은 True Label (Y)가 있기에 모델 예측값을 활용하여 feature importance 등을 산출하여 변수 중요도를 판단할 수 있다. 그러나, 인과추론에서 개별효과 (ITE)는 관측되지 않기에 True Label이 존재하지 않는다. 심지어 CATE 추정을 위한 모델은 대체로 복잡하기 때문에 어떤 feature가 treatment 효과를 만든 것인지를 모델 내부에서 설명할 수 없게 된다.

이를 해결할 방법이 없을까?

!pip install econml causeinfer -q

Hillstorm Email DataSet#

T: 0 = No-Email (Control), 1 = Email => Binary Treatment 문제로 변형

X:

  • recency: 최근 방문일

  • history: 구매 히스토리 금액

  • mens / womens : 관심 상품 카테고리

  • newbie: 신규 고객인지

  • zip_code_*: 지역 더미

  • history_segment_*: 구매력 세그먼트

  • channel_multichannel / phone / web: 기존 구매 채널

Y: visit: 사이트 방문 여부 (0/1)

출처: K. Hillstrom. “The MineThatData E-Mail Analytics And Data Mining Challenge”. 2008. URL: https://blog.minethatdata.com/2008/03/minethatdata-e-mail-analytics-and-data.html.


결론적으로, Hillstrom Email Dataset에 대한 CATE 분석 결과, womens(여성 관심 카테고리)와 history(과거 구매 이력) 변수가 이메일 발송의 인과적 효과 이질성을 설명하는 데 가장 핵심적인 역할을 하는 것으로 나타난다.

from causeinfer.data.hillstrom import download_hillstrom, load_hillstrom
import pandas as pd
import numpy as np

# 데이터 다운로드
# download_hillstrom(data_path="data/causal_ml_data")
data = load_hillstrom(
    file_path="data/causal_ml_data",
    format_covariates=True,
    download_if_missing=True,
    normalize=True
)


X = data['features']
T = (data['treatment'] != 0).astype(int) # email을 받았는가 (Binary)
Y = data['response_visit']
X_df = pd.DataFrame(data['features'], columns=data['feature_names'])
X_df.info()
<class 'pandas.core.frame.DataFrame'>
RangeIndex: 64000 entries, 0 to 63999
Data columns (total 18 columns):
 #   Column                    Non-Null Count  Dtype 
---  ------                    --------------  ----- 
 0   recency                   64000 non-null  object
 1   history                   64000 non-null  object
 2   mens                      64000 non-null  object
 3   womens                    64000 non-null  object
 4   newbie                    64000 non-null  object
 5   zip_code_rural            64000 non-null  object
 6   zip_code_surburban        64000 non-null  object
 7   zip_code_urban            64000 non-null  object
 8   history_segment_0_100     64000 non-null  object
 9   history_segment_1000+     64000 non-null  object
 10  history_segment_100_200   64000 non-null  object
 11  history_segment_200_350   64000 non-null  object
 12  history_segment_350_500   64000 non-null  object
 13  history_segment_500_750   64000 non-null  object
 14  history_segment_750_1000  64000 non-null  object
 15  channel_multichannel      64000 non-null  object
 16  channel_phone             64000 non-null  object
 17  channel_web               64000 non-null  object
dtypes: object(18)
memory usage: 8.8+ MB
X_df
recency history mens womens newbie zip_code_rural zip_code_surburban zip_code_urban history_segment_0_100 history_segment_1000+ history_segment_100_200 history_segment_200_350 history_segment_350_500 history_segment_500_750 history_segment_750_1000 channel_multichannel channel_phone channel_web
0 1.207742 -0.389 1 0 0 False True False False False True False False False False False True False
1 0.067358 0.339611 1 1 1 True False False False False False True False False False False False True
2 0.352454 -0.239834 0 1 1 False True False False False True False False False False False False True
3 0.922646 1.693265 1 0 1 True False False False False False False False True False False False True
4 -1.073025 -0.768062 1 0 0 False False True True False False False False False False False False True
... ... ... ... ... ... ... ... ... ... ... ... ... ... ... ... ... ... ...
63995 1.207742 -0.533051 1 0 0 False False True False False True False False False False False False True
63996 -0.217738 -0.793163 0 1 1 False False True True False False False False False False False True False
63997 0.067358 -0.827986 1 0 1 False False True True False False False False False False False True False
63998 -1.358121 1.213523 1 0 1 False True False False False False False False True False True False False
63999 -1.358121 0.900748 0 1 0 False True False False False False False True False False False False True

64000 rows × 18 columns

import numpy as np
import pandas as pd

X_clean_df = X_df.copy()

for col in X_clean_df.columns:
    try:
        X_clean_df[col] = pd.to_numeric(X_clean_df[col])
    except (ValueError, TypeError):
        pass

    # 2️⃣ 여전히 object인 경우만 boolean 문자열 처리
    if X_clean_df[col].dtype == object:
        X_clean_df[col] = (
            X_clean_df[col]
            .astype(str)
            .str.lower()
            .map({'true': 1, 'false': 0})
        )

    # 3️⃣ 이진 변수는 int, 나머지는 float
    uniq = X_clean_df[col].dropna().unique()
    if set(uniq).issubset({0, 1}):
        X_clean_df[col] = X_clean_df[col].astype(int)
    else:
        X_clean_df[col] = X_clean_df[col].astype(float)

X_clean_df.info()
<class 'pandas.core.frame.DataFrame'>
RangeIndex: 64000 entries, 0 to 63999
Data columns (total 18 columns):
 #   Column                    Non-Null Count  Dtype  
---  ------                    --------------  -----  
 0   recency                   64000 non-null  float64
 1   history                   64000 non-null  float64
 2   mens                      64000 non-null  int64  
 3   womens                    64000 non-null  int64  
 4   newbie                    64000 non-null  int64  
 5   zip_code_rural            64000 non-null  int64  
 6   zip_code_surburban        64000 non-null  int64  
 7   zip_code_urban            64000 non-null  int64  
 8   history_segment_0_100     64000 non-null  int64  
 9   history_segment_1000+     64000 non-null  int64  
 10  history_segment_100_200   64000 non-null  int64  
 11  history_segment_200_350   64000 non-null  int64  
 12  history_segment_350_500   64000 non-null  int64  
 13  history_segment_500_750   64000 non-null  int64  
 14  history_segment_750_1000  64000 non-null  int64  
 15  channel_multichannel      64000 non-null  int64  
 16  channel_phone             64000 non-null  int64  
 17  channel_web               64000 non-null  int64  
dtypes: float64(2), int64(16)
memory usage: 8.8 MB
X_clean_df.head()
recency history mens womens newbie zip_code_rural zip_code_surburban zip_code_urban history_segment_0_100 history_segment_1000+ history_segment_100_200 history_segment_200_350 history_segment_350_500 history_segment_500_750 history_segment_750_1000 channel_multichannel channel_phone channel_web
0 1.207742 -0.389000 1 0 0 0 1 0 0 0 1 0 0 0 0 0 1 0
1 0.067358 0.339611 1 1 1 1 0 0 0 0 0 1 0 0 0 0 0 1
2 0.352454 -0.239834 0 1 1 0 1 0 0 0 1 0 0 0 0 0 0 1
3 0.922646 1.693265 1 0 1 1 0 0 0 0 0 0 0 1 0 0 0 1
4 -1.073025 -0.768062 1 0 0 0 0 1 1 0 0 0 0 0 0 0 0 1
# 요약통계
print(pd.Series(T, name='Treatment').value_counts(normalize=True), "\n")
print(pd.Series(Y, name='Response Visit').value_counts(normalize=True))
Treatment
1    0.667094
0    0.332906
Name: proportion, dtype: float64 

Response Visit
0    0.853219
1    0.146781
Name: proportion, dtype: float64
X = X_clean_df.values # 최종변환

Causal Tree#

Causal Tree는 트리 기반의 해석 가능한 모델이다.

Causal Tree는 Y (Outcome)이 아니라 Treatment Effect 이질성이 크게 나타나는 방향으로 데이터를 분할해 하위 집단별 평균 치료효과를 구조적으로 보여준다.

출처:

https://www.pywhy.org/EconML/spec/interpretability.html

이를 해결하기 위해 CATE 추정 후에 해석 가능한 surrogate 모델을 적용한다. 특히, econml은 CATE 예측값들을 새로운 label처럼 취급하여 다시 하나의 tree를 학습시키는 SingleTreeCateInterpreter를 제공한다. 또한, “누구에게 treatment를 적용해야 가장 이득이 되는가”라는 정책적 의사결정 문제를 해결하기 위한 방법으로 SingleTreePolicyInterpreter 역시 제공된다.

from econml.cate_interpreter import SingleTreeCateInterpreter, SingleTreePolicyInterpreter
from econml.dml import CausalForestDML
# 기준이 되는 CATE 예측 모델 (다른 모델을 사용해도 됨.)

est = CausalForestDML(
    n_estimators=500,
    min_samples_leaf=100,
    max_depth=10,
    random_state=42
)

est.fit(Y, T, X=X)
<econml.dml.causal_forest.CausalForestDML at 0x19d403a90>
# 각 leaf는 서로 다른 방식으로 개입에 반응하는(subgroup-level heterogeneity) 하위 집단을 나타냄.
# 이를 통해 왜 효과가 달라지는지를 설명할 수 있다. 

intrp = SingleTreeCateInterpreter(include_model_uncertainty=True, max_depth=2, min_samples_leaf=10)
intrp.interpret(est, X)
intrp.plot(feature_names=list(data['feature_names']))
../_images/62a6bd360e1943e6df8595414816f55d810c39e46f21008e7ef81a81bc182d19.png

Interpretation:

  • 루트 노드: 조건: womens == 0, 여성 카테고리에 관심이 없는 고객은 이메일을 받으면, 이메일을 받지 않은 유사 고객에 비해 +6.1%p 상승

  • womens == 0 그룹: 과거 구매력이 0.438 이하에 속하는지 여부

    • 속하지 않으면: 해당 특성을 가지지 않은 고객이 이메일을 받으면, 이메일을 받지 않은 유사 고객에 비해 +4.1%p 상승

    • 속하면: 해당 특성을 가진 고객이 이메일을 받으면, 이메일을 받지 않은 유사 고객에 비해 +5.3%p 상승

  • womens == 1 그룹: 과거 구매력이 1.107 이하에 속하는지 여부

    • 속하지 않으면: 해당 특성을 가지지 않은 고객이 이메일을 받으면, 이메일을 받지 않은 유사 고객에 비해 +7.3%p 상승

    • 속하면: 해당 특성을 가진 고객이 이메일을 받으면, 이메일을 받지 않은 유사 고객에 비해 +8.9%p 상승

결론: 이메일 마케팅 uplift의 가장 큰 결정 요인은 여성 카테고리 관심 여부(womens 변수)이며, 트리 마지막 leaf 기준으로 가장 uplift가 높은 그룹은 womens == 1 AND history > 1.107이다.

# Policy Interpreter는 데이터에서 처치 (이메일)를 받는 것이 유리한 사람과 불리한 사람을 분리하는 트리를 학습한다.
# 특히 정책의 '비용'을 설정할 수 있으며, 이를 통해 누구에게 treatment를 해야 이득인지 결정할 수 있다.

intrp = SingleTreePolicyInterpreter(
    max_depth=3,
    min_samples_leaf=500,
    min_impurity_decrease=1e-5,
    random_state=42
)
intrp.interpret(est, X,  sample_treatment_costs=0.06)
intrp.plot(feature_names=list(data['feature_names']))
../_images/96dfed5dda66f6c57e763afe918e9c920756a9cd92eb131052ab5963266d7c92.png

방문 확률이 +6%p 이상 올라가는 고객에게만 이메일을 보내는 게 경제적으로 의미 있다고 가정한다면:

  • 전혀 이메일을 보내지 않는 것보다 평균 방문 확률이 0.8%p 증가

  • “여성 카테고리에 관심 없는 고객”에게는 이메일을 보내면 오히려 방문률이 1.7%p 감소

  • “여성 카테고리에 관심 있는 고객”에게는 이메일을 보내면 오히려 방문률이 1.5%p 증가

=> 여성 카테고리에 관심 있는 고객에게 이메일을 보내는 것이 권장된다.

Feature Importance / Permutation importance#

  • Feature importance는 해당 모델 (여기서는 트리)에서 얼마나 자주 분할 기준으로 사용됐는가를 의미하며, Permutation importance는 섞었을 때 예측이 얼마나 망가지는가를 의미한다.

import pandas as pd
import numpy as np
import seaborn as sns
from sklearn.metrics import mean_squared_error
fi = est.feature_importances_

fi_df = pd.DataFrame({
    "feature": X_df.columns,
    "importance": fi
}).sort_values("importance", ascending=False)

sns.barplot(x='importance',y='feature',data=fi_df)
<Axes: xlabel='importance', ylabel='feature'>
../_images/cc899bf20476b88efec9531eda165cd632ca3aa80fc2b152fab73e58603f0279.png

[해석 결과]

  1. 여성 타겟 여부가 이메일 처치 효과의 이질성을 가장 강하게 결정한다.

  2. 과거 구매 이력은 이메일 효과를 구분하는 핵심 연속형 변수이며, 구매 이력이 많은 고객과 적은 고객 사이에서 효과 차이가 뚜렷하다.

tau_hat = est.effect(X)


def permutation_importance_cate(
    est, X, tau_ref, n_repeats=5, random_state=42
):
    rng = np.random.RandomState(random_state)
    importances = []

    for j in range(X.shape[1]):
        losses = []

        for _ in range(n_repeats):
            X_perm = X.copy()
            rng.shuffle(X_perm[:, j])

            tau_perm = est.effect(X_perm)
            loss = mean_squared_error(tau_ref, tau_perm)
            losses.append(loss)

        importances.append(np.mean(losses))

    return np.array(importances)
perm_imp = permutation_importance_cate(
    est,
    X,
    tau_hat,
    n_repeats=5
)

perm_df = pd.DataFrame({
    "feature": X_df.columns,
    "perm_importance": perm_imp
}).sort_values("perm_importance", ascending=False)

perm_df
feature perm_importance
3 womens 8.236541e-04
1 history 1.436563e-04
2 mens 1.056059e-04
0 recency 6.466865e-05
5 zip_code_rural 1.008519e-05
4 newbie 9.412965e-06
17 channel_web 5.177159e-06
7 zip_code_urban 4.063373e-06
6 zip_code_surburban 3.271713e-06
16 channel_phone 2.262342e-06
10 history_segment_100_200 1.711652e-06
11 history_segment_200_350 1.194691e-06
15 channel_multichannel 1.157469e-06
13 history_segment_500_750 8.163955e-07
14 history_segment_750_1000 6.557521e-07
12 history_segment_350_500 6.197446e-07
8 history_segment_0_100 3.680012e-09
9 history_segment_1000+ 1.436699e-10

[해석 결과]

  1. 여성 타겟 여부는 섞는 순간, CATE 예측이 가장 크게 무너진다.

  2. 과거 구매 이력은 성별 다음으로 중요한 효과 조절 변수이며, 고객의 반응 강도를 연속적으로 조절하는 역할을 한다.

  3. 남성 관련 관심 변수 역시 효과 차이를 만드는 데 기여한다.

SHAP#

위의 Feature Imporatnce나 Permutation importance과 달리, SHAP은 개인별로 원인을 설명할 수 있으며, 그리고 각 변수의 방향성에 대해 설명할 수 있다.


기존 SHAP 분석과 동일하나, 여기서 SHAP의 목표는 outcome 예측에 기여하는 변수가 아닌, 처치 반응성을 변화시키는 변수 (Effect Modifiers)라는 것에 주의하여 해석할 수 있다.

출처:

Individual#

import shap

# 묘사를 위해 100개의 샘플에서만 모델을 돌림.

ind = 0
X_bg = X[ind:ind+100]

shap_values = est.shap_values(X_bg)

exp = shap_values["Y0"]["T0"]
exp.feature_names = X_df.columns.tolist()
exp.data = X_bg

shap.plots.waterfall(exp[0])
../_images/b481174ffd10612e51f1bc459c3903f1a9e0aa986148c752091afbd8cc165ed8.png
  1. Average Treatment Effect (Base Value): 0.06

  2. Estimated CATE for this individual: 0.029

  3. 부정적으로 영향을 미친 요인

    • womens == 0 : 여성 고객이 아닌 경우

    • history: 구매 이력이 적어지는 경우

Global#

shap.plots.bar(exp)
../_images/4419f48e9f2d85eba9d184ccd336df7e2b8c5b2356d36fe7d42f6eb2545a47cf.png

[해석 결과]

  • 전체적으로 어떤 변수가 처치 효과 차이를 가장 많이 만들어내는가를 해석할 수 있는 Bar Plot

    • womens, mens 등의 성별 변수가 큰 영향을 미침.

# 묘사를 위해 100개의 샘플에서만 모델을 돌림.

np.random.seed(42)
idx = np.random.choice(X.shape[0], size=100, replace=False)
X_shap = X[idx]

shap_values = est.shap_values(X_shap)
exp = shap_values["Y0"]["T0"]

exp.data = X_shap
exp.feature_names = X_df.columns.tolist()

shap.plots.beeswarm(exp)
../_images/4ab0af9767d397c4cb4baf17301f0616584a134e868c71540eb4f8eb7fd698b4.png

[해석결과]

  1. 여성 타겟 제품 관심 고객일수록 처치 효과가 증가. (여성 타겟 제품 관심 고객이 아닌 경우 처치 효과가 감소)

  2. 구매 이력이 적은 고객일수록 처치 반응성이 크다. (다만 같은 세그먼트 내에서도 개인별 효과 차이가 존재한다.)

  3. 남성 타겟 제품 관심 고객일수록 이메일 처치 효과가 상대적으로 높다.