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

如何使用innvestigate工具包针对表格数据实现LRP并可视化?

解决表格数据LRP可视化与innvestigate适配问题

一、可视化错误原因与解决方案

你遇到的TypeError: Invalid shape (10,) for image data是因为plt.imshow()仅支持2D(图像H×W)或3D(H×W×3/4)数组,而表格数据的LRP分析结果是1D特征相关性数组,直接使用imshow不符合要求。以下是两种适配方案:

方案1:条形图(更适合非技术用户)

条形图能直观展示每个特征的贡献大小与正负,比热力图更易理解:

import matplotlib.pyplot as plt
import numpy as np

# 替换为你的实际特征名称,无名称可用索引替代
feature_names = [f"Feature {i+1}" for i in range(analysis.shape[1])]
scores = analysis.squeeze()  # 获取单个样本的1D相关性分数

plt.figure(figsize=(10,6))
# 正贡献用红色,负贡献用蓝色
bars = plt.bar(feature_names, scores, color=np.where(scores>0, '#ff4444', '#33b5e5'))
plt.title("LRP特征贡献度(单样本)")
plt.xlabel("特征")
plt.ylabel("相关性分数")
plt.xticks(rotation=45)
# 为每个条形添加数值标签
for bar in bars:
    height = bar.get_height()
    plt.text(bar.get_x() + bar.get_width()/2., height,
             f'{height:.2f}', ha='center', va='bottom')
plt.tight_layout()
plt.show()

方案2:转2D数组实现热力图

若坚持使用热力图,可将1D数组转为1×N的2D矩阵:

import matplotlib.pyplot as plt

feature_names = [f"Feature {i+1}" for i in range(analysis.shape[1])]
heatmap = analysis.reshape(1, -1)  # 转为(1, 10)的2D数组

plt.figure(figsize=(10,2))
im = plt.imshow(heatmap, cmap="seismic", aspect='auto')
plt.xticks(range(len(feature_names)), feature_names, rotation=45)
plt.yticks([])  # 隐藏y轴(仅单个样本)
plt.title("LRP特征相关性热力图")
plt.colorbar(im, orientation='horizontal', pad=0.2)
plt.tight_layout()
plt.show()

二、innvestigate处理表格数据的注意事项

  1. 模型输入形状无需修改:innvestigate支持全连接网络的表格数据输入((batch_size, n_features)),不需要转成4D图像张量,你当前的分析器创建方式是正确的。
  2. LRP变体选择:若LRPZ效果不稳定,可尝试带ε的LRP变体,提升鲁棒性:
LRP_analyzer = innvestigate.analyzer.LRPEpsilon(model, epsilon=1e-2)
  1. 分析结果维度匹配:analyze()返回的结果形状与输入一致,输入(1, 10)则输出(1, 10),squeeze()后得到1D数组是正常现象,并非工具不支持表格数据。

内容的提问来源于stack exchange,提问作者user19534996

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 06:06:26