如何使用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处理表格数据的注意事项
- 模型输入形状无需修改:innvestigate支持全连接网络的表格数据输入(
(batch_size, n_features)),不需要转成4D图像张量,你当前的分析器创建方式是正确的。 - LRP变体选择:若LRPZ效果不稳定,可尝试带ε的LRP变体,提升鲁棒性:
LRP_analyzer = innvestigate.analyzer.LRPEpsilon(model, epsilon=1e-2)
- 分析结果维度匹配:
analyze()返回的结果形状与输入一致,输入(1, 10)则输出(1, 10),squeeze()后得到1D数组是正常现象,并非工具不支持表格数据。
内容的提问来源于stack exchange,提问作者user19534996
相关产品推荐
相关产品推荐

