You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.29 05:17:51