TensorFlow实战:5步搞定因果推断模型TARNet(附完整代码)
如果你正在处理营销效果评估、药物疗效分析或者任何需要回答“如果...会怎样”的业务问题,那么因果推断就是你工具箱里不可或缺的利器。传统的机器学习模型擅长预测相关性,但要剥离出真正的因果效应,尤其是在高维、非随机化的观测数据中,常常力不从心。TARNet(Treatment-Agnostic Representation Network)作为深度因果推断领域的经典模型,提供了一种优雅的解决方案:它通过一个共享的表示层学习协变量特征,再针对不同干预(Treatment)分支进行预测,巧妙地平衡了偏差与方差。今天,我们就抛开复杂的理论推导,直接上手TensorFlow,用五个清晰的步骤,从零搭建一个可运行的TARNet模型,并解决实际工程中令人头疼的维度灾难和样本不平衡问题。
1. 环境准备与数据理解
在开始敲代码之前,确保你的工作环境已经就绪。我们将使用TensorFlow 2.x,这是目前的主流选择,其Keras API能让我们像搭积木一样构建模型。此外,一些数据处理和可视化的库也会用到。
pip install tensorflow==2.10.0 pandas numpy scikit-learn matplotlib seaborn
接下来,理解我们要处理的数据结构至关重要。因果推断的数据通常包含三部分:
- 协变量 (X): 描述样本特征的多维向量,例如用户的年龄、历史行为、设备信息等。
- 干预/处理 (T): 一个二元或多元的指示变量,表示样本接受了哪种处理(如:广告A=1,广告B=0;或用药=1,不用药=0)。
- 结果 (Y): 我们关心的观测结果,比如点击率、康复指标等。
核心挑战在于反事实的缺失:对于一个给定的用户,我们只能观察到他在一种处理下的结果,而无法知道如果给他另一种处理,结果会怎样。TARNet的目标,正是从观测数据中学习,去估计这个“缺失”的反事实结果。
一个典型的数据集可能长这样(以Pandas DataFrame为例):
| 用户ID | 年龄 (X1) | 收入 (X2) | 看到广告 (T) | 是否点击 (Y) |
|---|---|---|---|---|
| 1 | 25 | 50000 | 1 | 1 |
| 2 | 35 | 80000 | 0 | 0 |
| 3 | 25 | 45000 | 0 | 1 |
| ... | ... | ... | ... | ... |
注意:在实际业务数据中,处理组(T=1)和对照组(T=0)的样本分布很可能是不平衡的,例如高收入用户更可能看到高价广告。这种“选择偏差”会直接污染效应估计,是TARNet需要克服的关键。
2. 数据预处理与特征工程
拿到原始数据后,直接丢给模型往往效果不佳。我们需要进行一系列预处理,为模型训练打下坚实基础。这一步虽然繁琐,但很大程度上决定了模型的上限。
首先,处理缺失值和异常值。对于连续型协变量,常见的做法是用中位数或均值填充缺失值;对于类别型变量,可以单独设置一个“缺失”类别。异常值则可以根据业务逻辑或统计方法(如IQR法则)进行截断或转换。
import pandas as pd
import numpy as np
from sklearn.impute import SimpleImputer
from sklearn.preprocessing import StandardScaler, OneHotEncoder
from sklearn.compose import ColumnTransformer
# 假设df是我们的DataFrame
# 分离特征、处理和结果
X = df.drop([‘treatment‘, ‘outcome‘], axis=1)
T = df[‘treatment‘].values
Y = df[‘outcome‘].values
# 区分数值型和类别型特征
numeric_features = X.select_dtypes(include=[‘int64‘, ‘float64‘]).columns
categorical_features = X.select_dtypes(include=[‘object‘, ‘category‘]).columns
# 构建预处理管道
numeric_transformer = Pipeline(steps=[
(‘imputer‘, SimpleImputer(strategy=‘median‘)),
(‘scaler‘, StandardScaler())
])
categorical_transformer = Pipeline(steps=[
(‘imputer‘, SimpleImputer(strategy=‘constant‘, fill_value=‘missing‘)),
(‘onehot‘, OneHotEncoder(handle_unknown=‘ignore‘, sparse=False))
])
preprocessor = ColumnTransformer(
transformers=[
(‘num‘, numeric_transfor

&spm=1001.2101.3001.5002&articleId=148673742&d=1&t=3&u=a1d88e4beddf467dad45ef8548342273)

被折叠的 条评论
为什么被折叠?



