做机器学习模型解释的时候,SHAP 瀑布图几乎是绕不开的标配。不管是竞赛提分后的复盘、风控模型评审,还是写论文放实验对比图,一张干净的 SHAP 瀑布图往往比一大段文字更有说服力。不过用久了你会发现,SHAP 默认出图虽然能用,但细节上总差点意思——尤其是那条多余的坐标轴边框,放在 PPT 和论文里显得特别笨重,跟整体排版格格不入。
这篇文章专门聊一件事:怎么给 SHAP 值瀑布图去除边框,把出图质感从“能用”提升到“好看”。我会从 SHAP 可视化的底层逻辑讲起,把 matplotlib 控件的操作、不同版本 SHAP 的兼容写法、多组瀑布图的展示方案都过一遍,最后附带几个我实际踩过坑之后的排查清单。适合天天跟模型解释打交道的数据分析师、算法工程师,也适合刚接触 SHAP 想快速出图的新手。
1. 内容整体设计与思路拆解
1.1 先搞清楚 SHAP 瀑布图到底画了什么
SHAP(SHapley Additive exPlanations)的核心思想是把模型对某个样本的预测值,分解成“基线值 + 每个特征的贡献值”。瀑布图就是这种分解结果最直观的呈现方式:从底部的 base value 出发,每个特征像瀑布一样依次叠加,红色箭头表示把预测值往上推,蓝色箭头表示往下拉,最后到达最终的预测结果 f(x)。
这里有一个经常被忽略的点:SHAP 瀑布图本身并不是一个普通柱状图或折线图,它是由大量 matplotlib 线段、文本和刻度对象组合而成的复合图形。所以你去修改它的边框、字体、配色,本质上操作的还是 matplotlib 的 Axes 和 Figure 对象。理解了这一层,后续所有定制才有方向。很多新手拿着plt.gca()却改不动样式,就是因为没搞清当前操作的坐标轴到底是不是 SHAP 内部创建的那一个。
1.2 默认出图有哪些影响观感的细节
SHAP 默认的 waterfall 图,信息量没问题,但视觉上确实有几个容易被吐槽的点:
第一是四条坐标轴边框。SHAP 瀑布图的特征名称在左侧,数值刻度在底部,按理说左侧和底部的边框还有一点对齐作用,但顶部和右侧纯粹属于多余线条。在多图排版或者投到大屏上时,这四条线会把视觉焦点扯散。
第二是默认的标题和 caption。SHAP 内部的shap.plots.waterfall会自动生成一些辅助说明文字,比如 base value 对应的数值、样本编号之类。这些文字在论文里去水印时经常需要单独处理,或者直接手动改掉。
第三是字体和间距。默认字体在 Windows 和 Mac 上渲染效果不同,中文环境还容易出现方块字。保存出来的图片周围留白也偏大,如果直接插入 Word 或 Markdown,经常要靠手动裁剪。
1.3 为什么“去除边框”是定制瀑布图的第一步
因为边框是影响“干净感”最直接的变量。你去翻那些好看的模型解释图,几乎清一色都是无边框、极简风格。把边框去掉之后,图里剩下的就是特征贡献的方向和大小,信息传递更聚焦。
从操作顺序看,去除边框也是最容易上手的定制动作,只需要操作 matplotlib 的spines属性即可,不需要重写 SHAP 的绘制逻辑。它适合作为学习 SHAP 定制化的切入点:先用最小成本理解ax和spines的关系,再去扩展别的定制需求,比如调整颜色、字体、保存尺寸,就会顺手很多。
2. 核心细节解析与实操要点
2.1 必须掌握的三个关键参数:show=False、ax、spines
先说show=False。shap.plots.waterfall默认执行完绘制后会自动调用plt.show(),这会导致图形窗口弹出,而且后续代码无法再对图形进行修改。所以定制瀑布图的第一步就是显式传入show=False,把绘制和展示分离。
shap.plots.waterfall(shap_values[0], max_display=10, show=False)接着是ax参数。SHAP 0.42 之后的版本支持直接把外部创建的 Axes 传给 waterfall,这样就不需要去猜“当前活跃的坐标轴到底是哪一个”,尤其适合要在同一个 Figure 里绘制多个子图的场景。
fig, ax = plt.subplots(figsize=(10, 6)) shap.plots.waterfall(shap_values[0], max_display=10, show=False, ax=ax)最后是spines。matplotlib 中每个 Axes 都有四条边框线,分别叫 top、bottom、left、right。要去除边框,就是把这几个 spine 对象的可见性设为 False。
for spine in ax.spines.values(): spine.set_visible(False)这三步是去除边框的最小组合,缺一个都不稳定。
2.2 不同版本 SHAP 的兼容写法
SHAP 库版本迭代比较快,shap.plots.waterfall的行为在不同版本间有一些细节差异。我实测过 0.41、0.44、0.45 以及最新的 0.46,整体用法一致,但有两个地方要注意:
一是低版本可能不支持ax参数。如果你用的 SHAP 版本较老,传入ax会直接报TypeError。保险做法是先打印出函数签名确认一下。
import inspect print(inspect.signature(shap.plots.waterfall))二是部分版本中waterfall内部会创建新的figure,导致plt.gca()拿到的坐标轴不是瀑布图所在轴。所以拿到坐标轴的方式不要依赖plt.gca(),优先用函数返回值和手动传入ax。
你还可以用一个万能兼容写法:先调用 waterfall 不传ax,然后通过plt.gcf().axes拿到当前 Figure 里的坐标轴列表,从中筛选出数据轴。
shap.plots.waterfall(shap_values[0], max_display=10, show=False) axes_list = plt.gcf().axes不过说实话,最省心的还是升级到新版本,然后显式传ax。这样代码干净,也不容易因为版本差异出问题。
2.3 去除边框的完整最小实现
这里给出一份可以直接跑通的最小代码,基于随机森林回归模型:
import shap import matplotlib.pyplot as plt from sklearn.ensemble import RandomForestRegressor from sklearn.datasets import fetch_california_housing # 加载数据与建模 data = fetch_california_housing() X, y = data.data, data.target feature_names = data.feature_names model = RandomForestRegressor(n_estimators=100, random_state=42) model.fit(X, y) # 计算 SHAP 值 explainer = shap.TreeExplainer(model) shap_values = explainer(X[:100]) # 对前100个样本计算 # 绘制第一个样本并关闭默认显示 fig, ax = plt.subplots(figsize=(10, 6)) shap.plots.waterfall(shap_values[0], max_display=10, show=False, ax=ax) # 去除所有边框 for spine in ax.spines.values(): spine.set_visible(False) # 微调布局并显示 plt.tight_layout() plt.show()这份代码在 SHAP 0.45.0 上实测没问题。如果你只想保留底部的边框,把上面循环换成只对top、left、right操作即可:
ax.spines['top'].set_visible(False) ax.spines['left'].set_visible(False) ax.spines['right'].set_visible(False)2.4 保存图片时的额外参数
实际项目里,瀑布图最终要保存成图片插入报告或论文。保存时有两个参数非常关键:dpi控制分辨率,bbox_inches='tight'会自动裁剪掉多余留白。
fig.savefig('shap_waterfall.png', dpi=300, bbox_inches='tight', facecolor='white')这里有个容易忽略的坑:facecolor默认是白色,但如果你在 Jupyter Notebook 里设置了深色主题,保存出来的图可能带透明背景或深色背景。所以保存时最好显式指定facecolor='white',避免插入文档后背景不一致。
3. 实操过程与核心环节实现
3.1 从零到一的定制实操流程
我平时做 SHAP 瀑布图定制的时候,习惯按照下面这个流程走,你可以直接抄作业。
第一步,先跑通基础绘制。用最简单的代码把瀑布图画出来,确认计算逻辑没有问题。这一步不做任何定制,纯粹验证 SHAP 值计算正确。
shap_values = explainer(X[:100]) shap.plots.waterfall(shap_values[0], max_display=10, show=False) plt.show()第二步,把图拆成“可编辑对象”。也就是增加fig, ax = plt.subplots(),并把ax显式传给 waterfall。这一步的意义是把绘图层和展示层解耦。
第三步,做边框定制。用spines循环把边框去掉,顺便清理掉不需要的辅助文字。
第四步,调整字体与尺寸,把所有文本对象的大小统一。
plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'DejaVu Sans'] plt.rcParams['axes.unicode_minus'] = False第五步,指定 dpi 和 bbox 保存图片。
3.2 字体、颜色与高亮深入定制
去除边框只是万里长征第一步。实际出图时,我们经常还需要调整字体大小、改变特征条颜色、单独高亮某个特征。
字体大小可以通过遍历 Axes 里的文本对象来修改:
for text_obj in ax.texts: text_obj.set_fontsize(12)如果你只想调整 X 轴或者 Y 轴刻度字体,用ax.tick_params更精准:
ax.tick_params(axis='x', labelsize=12) ax.tick_params(axis='y', labelsize=12)对于颜色,SHAP 瀑布图默认红色表示正向贡献、蓝色表示负向贡献。这个配色在大多数场景下够用,但如果你的报告主色调是品牌色,可以修改 SHAP 内部使用的colors参数。可惜的是shap.plots.waterfall没有直接暴露颜色参数,需要修改源码或者通过matplotlib的 color cycle 间接调整。
比较实用的定制是highlight_index。这个参数可以突出显示某一个特征,常用于案例分析中强调关键变量:
shap.plots.waterfall( shap_values[0], max_display=10, show=False, ax=ax, highlight_index='MedInc' )highlight_index既支持特征名字符串,也支持整数索引。在风控模型里,我经常用它单独高亮“征信查询次数”这类核心变量,评审汇报时非常出效果。
3.3 瀑布图到底能不能显示多组数据
“瀑布图可以显示多组数据吗”这个问题我经常看到。直接说结论:SHAP 的 waterfall 图本质上是单样本解释图,它展示的是一个样本的特征贡献分解。但“多组数据”可以从几个层面去理解,不同层面有不同的解决方案。
要对比同一样本在不同模型下的贡献,可以把两个瀑布图并排放在同一个 Figure 里:
fig, axes = plt.subplots(1, 2, figsize=(16, 6)) for i, model_name in enumerate(['Model_A', 'Model_B']): # 假设已有对应模型的 shap_values shap.plots.waterfall( shap_values_a[0] if i == 0 else shap_values_b[0], max_display=8, show=False, ax=axes[i] ) axes[i].set_title(model_name, fontsize=14) for spine in axes[i].spines.values(): spine.set_visible(False) plt.tight_layout() plt.show()要对比同一个模型对多个样本的解释,也是同样的思路,循环画子图即可。不过要注意子图数量不能太多,一般 2 到 4 个比较合适,再多每个子图的信息密度就会下降。
如果你要的是“群体层面”的汇总,瀑布图其实不是最佳选择。SHAP 库提供了shap.plots.bar和shap.plots.beeswarm,分别用于展示全局特征重要度和特征影响分布。这两个图更适合回答“整体上哪些特征影响最大”这类问题。
shap.plots.bar(shap_values) shap.plots.beeswarm(shap_values)另外还有一个偏门做法:把多个样本的瀑布图叠加在一起。通过设置透明度,可以看到多条分解路径的重合关系。这种方式在探索性分析阶段用来找离群点挺有意思,但不适合正式汇报,因为颜色叠加后会比较混乱。
3.4 手动绘制瀑布图的备用方案
追求极致定制的时候,shap.plots.waterfall自带的渲染逻辑反而会成为限制。比如你想把箭头改成圆角、想在箭头旁边显示百分比、想彻底重排布局,直接改内置函数就很费劲。
我自己的做法是,遇到这种需求就直接绕过 SHAP 的绘图函数,用 matplotlib 手动绘制瀑布图。核心思路很简单:把 SHAP 值按大小排序,依次累加出每个特征的起点和终点,然后用ax.hlines和ax.plot画线。
import numpy as np def manual_waterfall(base_value, shap_values, feature_names, ax): order = np.argsort(np.abs(shap_values))[::-1] shap_values = shap_values[order] feature_names = feature_names[order] cumulative = base_value for i, (name, sv) in enumerate(zip(feature_names, shap_values)): start = cumulative end = cumulative + sv ax.plot([start, end], [i, i], linewidth=8, color='red' if sv > 0 else 'blue', alpha=0.8) cumulative = end ax.axvline(base_value, color='gray', linestyle='--') ax.set_yticks(range(len(feature_names))) ax.set_yticklabels(feature_names) for spine in ax.spines.values(): spine.set_visible(False)手动绘制的优势是每个元素都在你掌控之下,想加什么就加什么。代价是代码量增加,且需要自己处理排序、截断、标签防重叠等问题。我的建议是:普通场景用内置 waterfall,只有内置函数实在满足不了需求时再手动绘制。
4. 常见问题与排查技巧实录
4.1 怎么改都对但边框还在
这是定制瀑布图时最让人抓狂的问题。明明用了spines设置不可见,边框还是原样显示。我之前排查过几次,主要有两种原因。
第一种原因是改错了坐标轴。在部分 SHAP 版本中,shap.plots.waterfall内部会创建自己的 Figure,外部plt.gca()拿到的坐标轴和瀑布图所在坐标轴不是同一个。解决方法是显式创建 Figure 并传ax参数。
第二种原因是边框线的颜色和背景色接近,看起来像“还在”。比如你在浅色背景上设置了浅灰色边框,视觉上就像没去掉。这种情况可以检查ax.spines['top'].get_visible()的返回值。
4.2 顶部或底部出现多余说明文字
SHAP 瀑布图底部有时候会生成一行小字说明,类似 “base value …”,顶部还可能出现样本编号或引用信息。这些文字不属于边框,但它们的存在会让图显得不干净。
处理方式有两种。一种是通过ax.get_xlabel()和ax.get_title()找到对应文本,然后清空:
ax.set_xlabel('') ax.set_title('')另一种是直接遍历所有文本对象,把不需要的隐藏掉:
for text_obj in ax.texts: if 'base value' in text_obj.get_text(): text_obj.set_visible(False)还有一种更粗暴但有效的方式:在plt.show()之前调用plt.gcf().texts同样遍历一遍。具体用哪种看你的版本和具体残留文本情况。
4.3 中文乱码与负号显示异常
在中文操作系统上,SHAP 瀑布图如果包含中文特征名,容易出现方块字或者乱码。这是因为 matplotlib 默认字体是英文的 DejaVu Sans,不覆盖中文字符。
通用解决方案是在绘图前设置字体:
plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'Arial Unicode MS'] plt.rcParams['axes.unicode_minus'] = False第二个设置很关键,它确保负号显示为标准的短横线而不是汉字“-”。如果只设置字体不设置unicode_minus,坐标轴上的负号会变成一个小方块,非常难看。
4.4 保存图片后文字模糊或被截断
保存图片出现文字模糊,基本都是 dpi 不够。一般报告用 200 dpi,论文投稿建议 300 dpi 以上。文字被截断则是因为画布尺寸过小或bbox_inches没设置对。
我的经验是:尺寸别省,figsize至少 10 x 6,然后dpi=300,最后bbox_inches='tight'裁掉多余白边。这样出来的图既清晰又紧凑,插到文档里也不会被拉伸变形。
有一个坑要特别提一下:如果你在脚本里先调用了plt.tight_layout(),再调用plt.savefig(..., bbox_inches='tight'),有时候会出现标题和坐标轴标签被裁掉一半的情况。解决办法是二选一,不要两个同时用。
4.5 版本相关报错快速速查
| 报错信息 | 原因 | 解法 |
|---|---|---|
TypeError: waterfall() got an unexpected keyword argument 'ax' | SHAP 版本过低 | 升级pip install -U shap,或在旧版本中去掉 ax 参数 |
ValueError: Explanation needs values | 传入对象不是 Explanation | 确保shap_values[0]是Explanation对象,或先用shap.explainers计算 |
AttributeError: module 'shap' has no attribute 'plots' | SHAP 版本过旧 | 升级 SHAP,0.30 之后的版本才有完整 plots API |
| 图形窗口一闪而过 | show=True默认行为 | 统一使用show=False并手动调用plt.show() |
4.6 一个提升出图效率的小技巧
如果你需要做大量瀑布图的批量导出,比如每个样本一张图,推荐写一个统一的渲染函数,把边框去除、字体设置、颜色高亮、保存逻辑都封装进去。这样调用一次就是一张成品图,不用每次重复修改样式。
def export_waterfall(shap_value, save_path, max_display=10, highlight=None): fig, ax = plt.subplots(figsize=(10, 6)) shap.plots.waterfall( shap_value, max_display=max_display, show=False, ax=ax, highlight_index=highlight ) for spine in ax.spines.values(): spine.set_visible(False) plt.tight_layout() fig.savefig(save_path, dpi=300, bbox_inches='tight', facecolor='white') plt.close(fig)批量出几百张图的时候,记得在循环里调用plt.close(fig)释放内存。我之前有一次没关图,跑了 200 张图之后内存直接爆了,进程卡死,白白浪费了半小时。
我在实际项目里的习惯是:报告用的图统一走封装函数,保证风格一致;探索性分析直接用 Jupyter 里交互式调参,调满意了再固化到脚本里。这样既能快速迭代,又能保证最终交付物是高质量的。