PyTorch自编码器输入特征重要性获取及SHAP实现正确性问询
从PyTorch自编码器(表格数据)获取输入变量特征重要性的问题
我正在处理表格数据集(非图像数据),需要从训练好的PyTorch自编码器中获取各输入变量的特征重要性。
自编码器模型结构
class AE(torch.nn.Module): def __init__(self, input_size, hidden_layer, latent_layer): super().__init__() self.encoder = torch.nn.Sequential( torch.nn.Linear(input_size, hidden_layer), torch.nn.ReLU(), torch.nn.Linear(hidden_layer, latent_layer) ) self.decoder = torch.nn.Sequential( torch.nn.Linear(latent_layer, hidden_layer), torch.nn.ReLU(), torch.nn.Linear(hidden_layer, input_size) ) def forward(self, x): encoded = self.encoder(x) decoded = self.decoder(encoded) return decoded
获取训练好的模型
通过以下函数得到训练后的模型实例:
average_loss, model, train_losses, test_losses = fullAE(batch_size=128, input_size=genes_tensor.shape[1], learning_rate=0.0001, weight_decay=0, epochs=50, verbose=False, dataset=genes_tensor, betas_value=(0.9, 0.999), train_dataset=genes_tensor_train, test_dataset=genes_tensor_test)
模型初始化代码:
model = AE(input_size=input_size, hidden_layer=int(input_size * 0.75), latent_layer=int(input_size * 0.5)).to(device)
尝试用SHAP计算特征重要性
我用SHAP工具做了如下尝试:
e = shap.DeepExplainer(model, genes_tensor) shap_values = e.shap_values( genes_tensor ) shap.summary_plot(shap_values,genes_tensor,feature_names=features)
遇到的问题
- 不确定当前SHAP实现是否正确
- 计算速度极慢,哪怕只处理1个样本也耗时很久
- 模型是多输出结构,SHAP生成的summary_plot样式和常见的不一样
- 试过Captum工具,但它只支持单输出神经网络,GitHub上相关AE/VAE案例多针对图像场景,无法适配我的表格数据需求
请问我的SHAP实现是否正确?
解答
1. 你的SHAP实现是否正确?
方向是对的,但存在几个关键问题需要调整:
- 多输出处理:自编码器输出维度和输入一致(每个输入特征对应一个输出节点),所以
shap_values会是一个长度等于输入特征数的列表,每个元素对应一个输出维度的SHAP值(形状为(样本数, 特征数))。这也是你的summary_plot样式异常的原因——默认会展示所有输出维度的结果。 - 背景数据集选择错误:
DeepExplainer的第二个参数应该是小批量背景样本(参考分布),而不是全量数据集。用全量数据作为背景会导致SHAP需要计算海量条件期望,直接拖慢速度,这是你计算极慢的核心原因。
2. 修正与优化方案
(1)优化背景数据集
从训练集中随机采样小批量样本作为背景(比如100-200个),大幅降低计算量:
# 从训练集中随机抽取100个背景样本 background = genes_tensor_train[torch.randperm(genes_tensor_train.size(0))[:100]] e = shap.DeepExplainer(model, background)
(2)处理多输出的特征重要性
因为自编码器的目标是重构输入,你可以通过聚合SHAP值得到全局特征重要性:
- 方案一:计算每个输入特征对自身重构输出的SHAP值绝对值均值(最贴合自编码器的重构目标):
# 提取每个输入特征对应自身输出维度的SHAP值,计算绝对值均值 feature_importance = np.array([ np.abs(shap_vals[:, idx]).mean() for idx, shap_vals in enumerate(shap_values) ])
- 方案二:计算每个输入特征对所有输出维度的SHAP值绝对值均值(全局贡献):
feature_importance = np.array([ np.abs(shap_vals).mean() for shap_vals in shap_values ])
之后可以用matplotlib绘制条形图,或者调整SHAP的可视化方式:
# 绘制条形图展示特征重要性 shap.summary_plot(shap_values, genes_tensor, feature_names=features, plot_type="bar")
(3)额外加速技巧
- 对要解释的数据集采样,不用全量计算:
# 只解释100个样本,进一步降低计算量 sample_data = shap.sample(genes_tensor, 100) shap_values = e.shap_values(sample_data)
- 若GPU计算有问题,可将模型和数据转回CPU:
model = model.cpu() background = background.cpu() sample_data = sample_data.cpu()
3. 表格数据自编码器的替代特征重要性方法
除了SHAP,还有更适合表格场景的高效方法:
- 基于重构误差的特征重要性:打乱单个输入特征的值,观察重构误差的上升幅度,幅度越大特征越重要,简单高效:
def calc_recon_importance(model, data, feature_names): original_loss = torch.nn.MSELoss()(model(data), data).item() importance = [] for idx in range(data.size(1)): # 复制数据并打乱当前特征 perturbed_data = data.clone() perturbed_data[:, idx] = perturbed_data[torch.randperm(perturbed_data.size(0)), idx] perturbed_loss = torch.nn.MSELoss()(model(perturbed_data), perturbed_data).item() importance.append(perturbed_loss - original_loss) # 按重要性排序 sorted_indices = np.argsort(importance)[::-1] return [(feature_names[i], importance[i]) for i in sorted_indices] # 使用示例 importance_list = calc_recon_importance(model, genes_tensor_test, features)
- 基于编码器权重的特征重要性:计算编码器第一层线性层权重的绝对值均值,快速反映输入特征对隐藏层的贡献:
# 获取编码器第一层权重 encoder_weights = model.encoder[0].weight.data.cpu().numpy() # 计算每个输入特征的权重绝对值均值 feature_importance = np.abs(encoder_weights).mean(axis=0)
内容的提问来源于stack exchange,提问作者José Adrián Pardo Pérez
相关产品推荐
相关产品推荐

