BART因果推断入门指南:从理论到Python实战

1次阅读
没有评论

共计 2449 个字符,预计需要花费 7 分钟才能阅读完成。

image.webp

1. BART 在因果推断中的独特价值

因果推断是机器学习中一个重要的研究方向,它不同于传统的预测任务,而是旨在理解干预(treatment)对结果(outcome)的因果效应。在因果推断领域,传统的 Propensity Score Matching (PSM) 和 Difference in Differences (DID) 方法存在一些局限性:

BART 因果推断入门指南:从理论到 Python 实战

  • PSM 依赖于正确指定的倾向得分模型,模型误设会导致估计偏差
  • DID 需要平行趋势假设,这在许多实际场景中难以验证

相比之下,BART (Bayesian Additive Regression Trees) 作为一种非参数贝叶斯方法,具有以下优势:

  • 自动处理非线性关系和交互效应
  • 通过正则化避免过拟合
  • 提供不确定性量化
  • 对模型误设更鲁棒

2. BART 的核心数学原理

BART 模型可以表示为:

y = f(x) + ε = ∑_{k=1}^m g_k(x; T_k, M_k) + ε, ε ~ N(0, σ^2)

其中:

  • m 是树的数量(通常 200 左右)
  • g_k 表示第 k 棵树,由树结构 T_k 和叶节点参数 M_k 决定
  • 每个叶节点的参数 μ ~ N(0, σ_μ^2/m)

贝叶斯框架下,我们需要指定以下先验:

  1. 树结构先验:鼓励树保持浅层(通常深度≤5)
  2. 叶节点参数先验:正则化参数大小
  3. 噪声方差先验:通常使用逆伽马分布

3. Python 实现完整流程

3.1 数据预处理

import numpy as np
import pandas as pd
from sklearn.preprocessing import StandardScaler

# 加载数据
data = pd.read_csv('causal_data.csv')

# 处理分类变量
data = pd.get_dummies(data, columns=['category_var'])

# 标准化连续变量
cont_vars = ['age', 'income']
scaler = StandardScaler()
data[cont_vars] = scaler.fit_transform(data[cont_vars])

# 划分处理组和对照组
X = data.drop(['outcome', 'treatment'], axis=1)
y = data['outcome']
W = data['treatment']

3.2 模型训练

from bartpy.sklearnmodel import SklearnModel

# 初始化 BART 模型
bart = SklearnModel(
    n_trees=200,          # 树的数量
    n_chains=4,           # MCMC 链数
    n_samples=1000,       # 后验样本数
    n_burn=200,           # 预烧期样本数
    alpha=0.95,           # 树深度控制
    beta=2.0,             # 叶节点参数控制
    thin=0.1,             # 抽样间隔
)

# 拟合模型
bart.fit(X, y)

3.3 计算平均处理效应 (ATE)

# 反事实预测
X_treat = X.copy()
X_treat['treatment'] = 1
X_control = X.copy()
X_control['treatment'] = 0

# 预测潜在结果
y_treat = bart.predict(X_treat)
y_control = bart.predict(X_control)

# 计算 ATE
ate = np.mean(y_treat - y_control)
print(f"Estimated ATE: {ate:.3f}")

4. 结果可视化

4.1 特征重要性

import matplotlib.pyplot as plt

# 获取特征重要性
importance = bart.feature_importances()

# 绘制条形图
plt.figure(figsize=(10, 6))
plt.barh(X.columns, importance)
plt.title('Feature Importance')
plt.xlabel('Importance Score')
plt.tight_layout()
plt.show()

4.2 部分依赖图 (PDP)

from sklearn.inspection import PartialDependenceDisplay

# 绘制 age 变量的 PDP
fig, ax = plt.subplots(figsize=(10, 6))
PartialDependenceDisplay.from_estimator(bart, X, features=['age'],
    ax=ax, line_kw={'linewidth': 3}
)
plt.title('Partial Dependence Plot for Age')
plt.show()

5. 实战注意事项

5.1 样本量要求

  • 处理组和对照组样本量应尽可能平衡
  • 推荐每组至少 500 个样本
  • 协变量分布差异较大时考虑倾向得分加权

5.2 协变量平衡检查

from causalinference import CausalModel

# 计算标准化均值差异
causal = CausalModel(y, W, X)
causal.reset_stats()
causal.est_propensity()
causal.trim_s()
causal.calc_stats()

# 检查平衡
print(causal.summary_stats)

5.3 MCMC 收敛诊断

# 检查迹图
bart.plot_trace('sigma')

# 计算 R -hat 统计量
rhat = bart.get_rhat()
print(f"R-hat statistics: {rhat}")

6. 开放问题与未来方向

  1. 未观测混杂变量处理:
  2. 结合工具变量方法
  3. 引入敏感性分析

  4. 高维稀疏特征处理:

  5. 嵌入变量选择先验
  6. 结合深度学习表示

  7. 异质性处理效应建模:

  8. 发展基于 BART 的 CATE 估计
  9. 结合元学习框架

7. 结论

BART 为因果推断提供了一个强大的非参数框架,特别适合处理复杂的非线性关系。通过本文介绍的方法,研究者可以快速实现 BART 模型并应用于实际问题。未来随着计算方法的进步,BART 在因果推断中的应用前景将更加广阔。

完整代码示例可在 GitHub 获取:[链接]

正文完
 0
评论(没有评论)